"""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