lijn14
完成C部分内容
d572bbd
Raw
History Blame Contribute Delete
2.51 kB
"""
损失函数模块 — Person C 负责实现
功能要求:
1. LabelSmoothedCrossEntropyLoss: 带标签平滑的交叉熵损失
技术要点:
- 标签平滑 (Label Smoothing) 是 Transformer 训练的标准技巧
- smoothing=0.1 意味着将 10% 的概率均匀分配给非目标类
- 需要忽略 padding token 的损失
"""
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
class LabelSmoothedCrossEntropyLoss(nn.Module):
"""
带标签平滑的交叉熵损失。
标签平滑将真实标签的概率质量从 1.0 重新分配:
target token 概率 = 1 - smoothing
其他 token 概率 = smoothing / (V - 1)
使用 KL 散度实现: loss = KL(smooth_target || log_softmax(logits))
参考: "Rethinking the Inception Architecture for Computer Vision" (Szegedy et al.)
"""
def __init__(self, smoothing: float = 0.1, pad_id: int = 0):
super().__init__()
self.smoothing = smoothing
self.pad_id = pad_id
self.confidence = 1.0 - smoothing
def forward(self, logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:
"""
Args:
logits: [B, T, V] — 模型输出
targets: [B, T] — 目标 token ids (padding 位置应为 pad_id, 默认 -100)
Returns:
loss: scalar
"""
logits = logits.contiguous()
targets = targets.contiguous()
batch_size, seq_len, vocab_size = logits.size()
logits_flat = logits.view(-1, vocab_size)
targets_flat = targets.view(-1)
log_probs = F.log_softmax(logits_flat, dim=-1)
# Mask that is True for real (non-padding) tokens
non_pad_mask = targets_flat.ne(self.pad_id).float()
if self.smoothing > 0.0:
smooth_dist = torch.full_like(log_probs, self.smoothing / (vocab_size - 1))
# Clamp to avoid negative-index writes for ignore-index values (e.g. -100)
safe_targets = targets_flat.clamp(min=0)
smooth_dist.scatter_(1, safe_targets.unsqueeze(1), self.confidence)
nll_loss = -torch.sum(smooth_dist * log_probs, dim=-1)
nll_loss = nll_loss * non_pad_mask
else:
nll_loss = F.nll_loss(
log_probs, targets_flat,
ignore_index=self.pad_id, reduction="none",
)
num_tokens = non_pad_mask.sum().clamp(min=1)
loss = nll_loss.sum() / num_tokens
return loss