Skip to content

Clip grad norm dtensor - #49189

Open
michaelbenayoun wants to merge 4 commits into
huggingface:mainfrom
michaelbenayoun:clip_grad_norm_dtensor
Open

michaelbenayoun wants to merge 4 commits into
huggingface:mainfrom
michaelbenayoun:clip_grad_norm_dtensor

Conversation

@michaelbenayoun

@michaelbenayoun michaelbenayoun commented Sep 29, 2026 •

Copy link
Copy Markdown
Member

CPU CI GPU run-slow

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.

@HuggingFaceDocBuilderDev

Copy link
Copy Markdown

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.

@github-actions

Copy link
Copy Markdown
Contributor

CI recap

Dashboard: View test results in Grafana
Latest run: 36583605720:1
Result: success | Jobs: 16 | Tests: 195,342 | Failures: 0 | Duration: 14h 50m

@qgallouedec qgallouedec left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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_norm

Once huggingface/accelerate#4266 handles mixed meshes, could the Trainer just call accelerator.clip_grad_norm_ and drop this one?

Comment on lines +87 to +95
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,
)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

it's unused, no?

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)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

worth parametrizing the test over norm_type=0 and inf since they have their own branch

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants