Download training_code/fast_execution.py from Recharge23/JAM-20K-single: direct link, hf CLI and curl.
- Browser
- Download file 1.25 kB
-
https://huggingface.co/Recharge23/JAM-20K-single/resolve/main/training_code/fast_execution.py
- Command line
-
hf download hf://Recharge23/JAM-20K-single/training_code/fast_execution.py
-
curl -L -o fast_execution.py https://huggingface.co/Recharge23/JAM-20K-single/resolve/main/training_code/fast_execution.py
1.25 kB
| """Execution-only optimizations; model tensors and inference contract are unchanged.""" | |
| import torch | |
| import torch.distributed as dist | |
| def install_fast_execution(model, precision): | |
| assert precision in ('fp32', 'tf32') | |
| torch.backends.cuda.matmul.allow_tf32 = precision == 'tf32' | |
| model._mot_driver.mot_checkpoint_mixed_attn = False | |
| model.set_training_runtime(use_gradient_checkpointing=False, | |
| use_gradient_checkpointing_offload=False, max_timestep_boundary=1., min_timestep_boundary=0.) | |
| def sum_gradients_fast(parameters, world): | |
| """SUM global-accumulation-scaled FP32 gradients in one collective.""" | |
| assert parameters and all(p.dtype == torch.float32 for p in parameters) | |
| flat = torch.empty(sum(p.numel() for p in parameters), device=parameters[0].device, dtype=torch.float32) | |
| views = [] | |
| offset = 0 | |
| for p in parameters: | |
| views.append(flat[offset:offset+p.numel()].view_as(p)) | |
| offset += p.numel() | |
| def reduce(): | |
| for p, view in zip(parameters, views): | |
| if p.grad is None: view.zero_() | |
| else: view.copy_(p.grad) | |
| if world > 1: dist.all_reduce(flat, op=dist.ReduceOp.SUM) | |
| for p, view in zip(parameters, views): p.grad = view | |
| return reduce | |