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