@@ -32,16 +32,21 @@ def l2_grad_clipper(parameters, max_norm):
3232 """Efficient L2 norm gradient clipping."""
3333
3434 overflow_buf = torch .zeros (1 , dtype = torch .int , device = 'cuda' )
35+ # Make sure we have an iterable.
3536 if isinstance (parameters , torch .Tensor ):
3637 parameters = [parameters ]
38+ # Filter parameters with gradients.
39+ parameters_with_grads = list (filter (
40+ lambda p : p .grad is not None , parameters ))
41+ # Filter parameters for norm calculations.
3742 mp_rank_is_zero = (get_model_parallel_rank () == 0 )
38- parameters = list (filter (lambda p : ( p . grad is not None ) and
39- ( p .model_parallel or mp_rank_is_zero ),
40- parameters ))
43+ parameters_for_norm = list (filter (
44+ lambda p : p .model_parallel or mp_rank_is_zero , parameters_with_grads ))
45+ # Calculate L2 norm.
4146 norm , _ = multi_tensor_applier (
4247 amp_C .multi_tensor_l2norm ,
4348 overflow_buf ,
44- [parameters ],
49+ [parameters_for_norm ],
4550 False # no per-parameter norm
4651 )
4752 # Sum across all model parallel GPUs.
@@ -50,10 +55,10 @@ def l2_grad_clipper(parameters, max_norm):
5055 op = torch .distributed .ReduceOp .SUM ,
5156 group = get_model_parallel_group ())
5257 total_norm = norm_2 .item () ** 0.5
53-
54- clip_coef = max_norm / (total_norm + 1e -6 )
55- grads = [p .grad for p in parameters ]
56- if clip_coef < 1 :
58+ # Scale to get max_norm.
59+ clip_coef = float ( max_norm ) / (total_norm + 1.0e -6 )
60+ grads = [p .grad for p in parameters_with_grads ]
61+ if clip_coef < 1.0 :
5762 multi_tensor_applier (
5863 amp_C .multi_tensor_scale ,
5964 overflow_buf ,
@@ -96,8 +101,8 @@ def clip_grad_norm(parameters, max_norm, norm_type=2):
96101 if clip_coef < 1 :
97102 for p in parameters :
98103 p .grad .data .mul_ (clip_coef )
99- elif norm_type == 2 :
100- total_norm = l2_grad_clipper (parameters , max_norm )
104+ # elif norm_type == 2:
105+ # total_norm = l2_grad_clipper(parameters, max_norm)
101106
102107 else :
103108 total_norm = 0
0 commit comments