Skip to content

Commit 9949642

Browse files
committed
optimize clip_grad_norm_ function
Optimize clip_grad_norm_ function by removing .item() calls to reduce wait time for the device on the host.
1 parent 9d2660d commit 9949642

1 file changed

Lines changed: 27 additions & 18 deletions

File tree

deepspeed/runtime/utils.py

Lines changed: 27 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -363,44 +363,53 @@ def clip_grad_norm_(parameters, max_norm, norm_type=2, mpu=None):
363363
if isinstance(parameters, torch.Tensor):
364364
parameters = [parameters]
365365
parameters = list(filter(lambda p: p.grad is not None, parameters))
366-
max_norm = float(max_norm)
367366
norm_type = float(norm_type)
367+
all_norms = []
368368
if norm_type == inf:
369-
total_norm = max(p.grad.data.abs().max() for p in parameters)
370-
total_norm_cuda = get_accelerator().FloatTensor([float(total_norm)])
369+
for p in parameters:
370+
all_norms.append(p.grad.data.abs().max().float())
371+
total_norm = torch.stack(all_norms).max()
372+
origin_device = total_norm.device.type
373+
total_norm = total_norm.to(get_accelerator().device_name())
371374
# Take max across all GPUs.
372375
if mpu is not None:
373-
dist.all_reduce(total_norm_cuda, op=dist.ReduceOp.MAX, group=mpu.get_model_parallel_group())
374-
total_norm = total_norm_cuda[0].item()
376+
dist.all_reduce(total_norm, op=dist.ReduceOp.MAX, group=mpu.get_model_parallel_group())
375377
else:
376378
total_norm = 0
377379
for p in parameters:
378380
if mpu is not None:
379381
if (mpu.get_model_parallel_rank() == 0) or is_model_parallel_parameter(p):
380-
param_norm = p.grad.data.norm(norm_type)
381-
total_norm += param_norm.item()**norm_type
382+
param_norm = p.grad.data.detach().float().norm(norm_type)
383+
all_norms.append(param_norm)
382384
else:
383-
param_norm = p.grad.data.float().norm(norm_type)
384-
total_norm += param_norm.item()**norm_type
385-
385+
param_norm = p.grad.data.detach().float().norm(norm_type)
386+
all_norms.append(param_norm)
387+
if len(all_norms) > 0:
388+
total_norm = torch.stack(all_norms).square().sum().float()
389+
else:
390+
total_norm = torch.FloatTensor([0.0]).to(parameters[0].device)
391+
origin_device = total_norm.device.type
392+
total_norm = total_norm.to(get_accelerator().device_name())
386393
# Sum across all model parallel GPUs.
387-
total_norm_cuda = get_accelerator().FloatTensor([float(total_norm)])
388394
if mpu is not None:
389-
dist.all_reduce(total_norm_cuda, op=dist.ReduceOp.SUM, group=mpu.get_model_parallel_group())
390-
total_norm = total_norm_cuda[0].item()**(1. / norm_type)
395+
dist.all_reduce(total_norm, op=dist.ReduceOp.SUM, group=mpu.get_model_parallel_group())
396+
total_norm = total_norm.pow(1. / norm_type)
391397

392398
# Need to average total_norm across different GPUs due to the presence of moe params
393399
pg = groups._get_data_parallel_group()
394400
scaled_norm = total_norm * 1.0 / float(dist.get_world_size(group=pg))
401+
scaled_norm_tensor = scaled_norm
395402

396-
scaled_norm_tensor = get_accelerator().FloatTensor([float(scaled_norm)])
397403
dist.all_reduce(scaled_norm_tensor, group=pg)
398-
total_norm = scaled_norm_tensor.item()
404+
total_norm = scaled_norm_tensor
405+
total_norm = total_norm.to(origin_device)
399406

407+
max_norm = torch.tensor([float(max_norm)], device=parameters[0].device)
400408
clip_coef = max_norm / (total_norm + 1e-6)
401-
if clip_coef < 1:
402-
for p in parameters:
403-
p.grad.data.mul_(clip_coef)
409+
tmp_tensor = torch.tensor([1.0], device=parameters[0].device)
410+
clip_coef = torch.max(tmp_tensor, clip_coef)
411+
for p in parameters:
412+
p.grad.data.mul_(clip_coef)
404413
return total_norm
405414

406415

0 commit comments

Comments
 (0)