JAM-20K-single / training_code /fast_execution.py
Recharge23's picture
Release separate task 4 and task 5 Franka-JAM 20K adapters
5e5ffd1 verified
Raw History Blame Contribute Delete
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