Repository navigation
Clip grad norm dtensor - #49189
Clip grad norm dtensor#49189michaelbenayoun wants to merge 4 commits into
Conversation
|
The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update. |
CI recapDashboard: View test results in Grafana |
qgallouedec
left a comment
There was a problem hiding this comment.
Thanks! Checked it against torch on the full grads (4 gloo ranks, grads on two sub-meshes of a 2-D mesh + plain tensors, norm_type 2/inf/0/1, max_norm 1/1e9/inf): all match!
It can be simpler though: the general path already handles the single-mesh case, so the fast path and the empty branch can go. This passes the same checks:
params_by_mesh = defaultdict(list)
for param in parameters:
if param.grad is not None:
params_by_mesh[param.grad.device_mesh if is_dtensor(param.grad) else None].append(param)
norms = [get_total_norm([p.grad for p in params], norm_type, foreach=foreach) for params in params_by_mesh.values()]
norms = torch.stack([n.full_tensor() if is_dtensor(n) else n for n in norms]) if norms else torch.zeros(1)
total_norm = norms.sum() if norm_type == 0 else torch.linalg.vector_norm(norms, norm_type)
# + error_if_nonfinite check
if max_norm != float("inf"):
for params in params_by_mesh.values():
clip_grads_with_norm_(params, max_norm, total_norm, foreach)
return total_normOnce huggingface/accelerate#4266 handles mixed meshes, could the Trainer just call accelerator.clip_grad_norm_ and drop this one?
| def setUp(self): | ||
| self.config = LlamaConfig( | ||
| vocab_size=16, | ||
| hidden_size=16, | ||
| intermediate_size=32, | ||
| num_hidden_layers=1, | ||
| num_attention_heads=4, | ||
| num_key_value_heads=4, | ||
| ) |
| parameter = torch.nn.Parameter(torch.zeros_like(gradient)) | ||
| parameter.grad = gradient | ||
| parameters.append(parameter) | ||
| expected_norm = torch.nn.utils.clip_grad_norm_(reference, max_norm, foreach=True) |
There was a problem hiding this comment.
worth parametrizing the test over norm_type=0 and inf since they have their own branch
What does this PR do?
One
clip_grad_norm_implementation that can support DTensors.We have one in the Trainer, we have another one in Accelerate, and plan to align there as well: huggingface/accelerate#4266.
The goal is to end up with one implementation.