Spaces:
Running on Zero
Running on Zero
Download mamba_ssm/utils/torch.py from voidful/BlueMagpie-TTS-Demo: direct link, hf CLI and curl.
- Browser
- Download file 676 Bytes
-
https://huggingface.co/spaces/voidful/BlueMagpie-TTS-Demo/resolve/44b3f8d43a273e20e4c406f7aa48b94fbacf9f38/mamba_ssm/utils/torch.py
- Command line
-
hf download hf://spaces/voidful/BlueMagpie-TTS-Demo@44b3f8d43a273e20e4c406f7aa48b94fbacf9f38/mamba_ssm/utils/torch.py
-
curl -L -o torch.py https://huggingface.co/spaces/voidful/BlueMagpie-TTS-Demo/resolve/44b3f8d43a273e20e4c406f7aa48b94fbacf9f38/mamba_ssm/utils/torch.py
676 Bytes
| import torch | |
| from functools import partial | |
| from typing import Callable | |
| def custom_amp_decorator(dec: Callable, cuda_amp_deprecated: bool): | |
| def decorator(*args, **kwargs): | |
| if cuda_amp_deprecated: | |
| kwargs["device_type"] = "cuda" | |
| return dec(*args, **kwargs) | |
| return decorator | |
| if hasattr(torch.amp, "custom_fwd"): # type: ignore[attr-defined] | |
| deprecated = True | |
| from torch.amp import custom_fwd, custom_bwd # type: ignore[attr-defined] | |
| else: | |
| deprecated = False | |
| from torch.cuda.amp import custom_fwd, custom_bwd | |
| custom_fwd = custom_amp_decorator(custom_fwd, deprecated) | |
| custom_bwd = custom_amp_decorator(custom_bwd, deprecated) | |