From 2e5e4170820e25cffe2e7ab8ef11608ade0c6bfc Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Sun, 20 Sep 2026 19:44:37 +0800 Subject: [PATCH] fix: detach probability weights in DFT loss --- src/twinkle/loss/chunked_cross_entropy.py | 6 +++--- src/twinkle/loss/cross_entropy.py | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/src/twinkle/loss/chunked_cross_entropy.py b/src/twinkle/loss/chunked_cross_entropy.py index 1e829a69..ad7ffef2 100644 --- a/src/twinkle/loss/chunked_cross_entropy.py +++ b/src/twinkle/loss/chunked_cross_entropy.py @@ -48,7 +48,7 @@ def forward(ctx, logits, labels, chunk_size, ignore_index, reduction, dft): logps = F.log_softmax( logits_chunk, dim=-1).gather(-1, labels_chunk.clamp(min=0).unsqueeze(-1)).squeeze(-1) - per_token = -logps * logps.exp() if dft else -logps + per_token = -logps * logps.exp().detach() if dft else -logps total_loss = total_loss + (per_token * mask).sum() total_count = total_count + mask.sum() @@ -86,7 +86,7 @@ def backward(ctx, grad_output): logps = F.log_softmax( logits_chunk, dim=-1).gather(-1, labels_chunk.clamp(min=0).unsqueeze(-1)).squeeze(-1) - per_token = -logps * logps.exp() if dft else -logps + per_token = -logps * logps.exp().detach() if dft else -logps loss_chunk = (per_token * mask).sum() grad_chunk = torch.autograd.grad(loss_chunk, logits_chunk, retain_graph=False)[0] @@ -168,7 +168,7 @@ def __call__(self, inputs, outputs, **kwargs): def _loss_from_logps(self, labels, logps): mask = (labels != self.ignore_index).float() - per_token = -logps * logps.exp() if self.dft else -logps + per_token = -logps * logps.exp().detach() if self.dft else -logps if self.reduction == 'mean': return LossOutput(loss=(per_token * mask).sum() / mask.sum().clamp(min=1), num_tokens=0) return LossOutput(loss=(per_token * mask).sum(), num_tokens=mask.sum().clamp(min=1)) diff --git a/src/twinkle/loss/cross_entropy.py b/src/twinkle/loss/cross_entropy.py index 8d3627c4..ed949f6f 100644 --- a/src/twinkle/loss/cross_entropy.py +++ b/src/twinkle/loss/cross_entropy.py @@ -39,7 +39,7 @@ def __call__(self, inputs, outputs, **kwargs): mask = (labels != self.ignore_index).float() # DFT: -p·log(p) instead of -log(p) - per_token = -logps * logps.exp() if self.dft else -logps + per_token = -logps * logps.exp().detach() if self.dft else -logps if self.reduction != 'sum': return LossOutput(loss=(per_token * mask).sum() / mask.sum().clamp(min=1), num_tokens=0)