Skip to content

Commit 03d2880

Browse files
author
Mostofa Patwary
committed
Merge branch 'l2_grad_clip_fix' into 'master'
Reverting l2 grad optimization See merge request ADLR/megatron-lm!74
2 parents 3c709cb + d218f9c commit 03d2880

1 file changed

Lines changed: 15 additions & 10 deletions

File tree

‎megatron/mpu/grads.py‎

Lines changed: 15 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)