[train] Fix hyper parallel tail accumulation loss scaling (#10705)

Co-authored-by: wcrzlh <weichaoran@huawei.com>
This commit is contained in:
Chaoran Wei
2026-07-30 17:29:59 +08:00
committed by GitHub
parent 9ce6b663e9
commit 1b47415a2f

View File

@@ -361,7 +361,12 @@ class HyperParallelTrainer(CustomSeq2SeqTrainer):
loss = loss.mean() loss = loss.mean()
if not getattr(self, "model_accepts_loss_kwargs", False) and getattr(self, "compute_loss_func", None) is None: if not getattr(self, "model_accepts_loss_kwargs", False) and getattr(self, "compute_loss_func", None) is None:
loss = loss / self.args.gradient_accumulation_steps accumulation_steps = getattr(
self,
"current_gradient_accumulation_steps",
self.args.gradient_accumulation_steps,
)
loss = loss / accumulation_steps
self.accelerator.backward(loss) self.accelerator.backward(loss)