[v1] fix grad norm and lr log (#10640)

Co-authored-by: mhh111 <mahonghao1@huawei.com>
This commit is contained in:
HelloWorldBeginner
2026-07-17 22:50:13 +08:00
committed by GitHub
parent 5f653cb96a
commit ef2d8f9da6
3 changed files with 31 additions and 8 deletions

View File

@@ -31,6 +31,7 @@ from abc import abstractmethod
import torch
import torch.nn.functional as F
from torch.distributed.tensor import DTensor
from ..accelerator.helper import ReduceOp
from ..accelerator.interface import Dim, DistributedInterface
@@ -279,12 +280,21 @@ class BaseTrainer:
# deepspeed: engine.step() already ran inside backward at the sync boundary
grad_norm = self._deepspeed_engine.get_grad_norm()
else:
grad_norm = torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.args.max_grad_norm).item()
if self.args.dist_config and self.args.dist_config.get("cp_size", 1) > 1:
grad_norm = grad_norm**2
grad_norm = DistributedInterface().all_reduce(grad_norm, op=ReduceOp.SUM, dim=Dim.CP)
grad_norm = grad_norm**0.5
# FSDP2 shards params/grads across the fsdp mesh, so clip_grad_norm_ returns a
# per-rank local shard norm (global / sqrt(shard_size)): reported grad_norm then
# scales as 1/sqrt(dp_size) and the clip coefficient is applied per-shard. Reduce
# to the true global norm first, then clip with it.
grads = [p.grad for p in self.model.parameters() if p.grad is not None]
total_norm = torch.nn.utils.get_total_norm(grads)
if isinstance(total_norm, DTensor):
# full_tensor all-reduces across the fsdp mesh (spans CP under default
# mp_shard=world); a separate CP reduce would over-count by sqrt(cp_size).
total_norm = total_norm.full_tensor()
# pass a Tensor: clip_grads_with_norm_ clamps max_norm / (total_norm + 1e-6).
torch.nn.utils.clip_grads_with_norm_(
self.model.parameters(), self.args.max_grad_norm, total_norm
)
grad_norm = total_norm.item()
if not torch.isfinite(torch.tensor(grad_norm)): # type: ignore # pyright: ignore [reportUnknownReturnType]
logger.warning_rank0(f"Gradient norm is not finite: {grad_norm}")

View File

@@ -90,7 +90,10 @@ def apply_sequence_parallel(model, model_args):
set_ulysses_sequence_parallel_group(DistributedInterface().get_group(Dim.CP))
try:
num_attention_heads, num_key_value_heads = model.config.num_attention_heads, model.config.num_attention_heads
num_attention_heads, num_key_value_heads = (
model.config.num_attention_heads,
model.config.num_key_value_heads,
)
except AttributeError:
num_attention_heads, num_key_value_heads = (
model.config.text_config.num_attention_heads,

View File

@@ -54,7 +54,17 @@ class LoggingCallback(TrainerCallback):
# Human-readable output to stdout
display_logs = {**logs, "step": state.global_step, "total_steps": state.num_training_steps}
parts = ", ".join(f"{k}: {v:.4f}" if isinstance(v, float) else f"{k}: {v}" for k, v in display_logs.items())
def _fmt(k: str, v) -> str:
if not isinstance(v, float):
return f"{k}: {v}"
# learning_rate is often < 1e-4 (e.g. 1e-5); :.4f would print "0.0000".
# Use :.4g so small values show as "1e-05" while 1e-4 still shows "0.0001".
if k == "learning_rate":
return f"{k}: {v:.4g}"
return f"{k}: {v:.4f}"
parts = ", ".join(_fmt(k, v) for k, v in display_logs.items())
logger.info_rank0(parts)
# Append to JSONL log file in output_dir