序列标注任务中CrossEntropyLoss交叉熵损失函数的正确使用
PyTorch序列标注任务CrossEntropyLoss正确使用方案
结论先行
你当前的写法完全符合PyTorch官方规范,不存在你担心的“判定为多维损失导致计算错误”的问题。
原理说明
PyTorch的CrossEntropyLoss原生支持高维分类场景,逻辑如下:
- 输入预测值要求形状为
(N, C, d1, d2, ..., dk),其中:N是batch sizeC是类别总数d1~dk是任意维度的额外空间/序列维度,没有数量限制
- 对应的标签要求形状为
(N, d1, d2, ..., dk),每个位置取值为0~C-1的类别索引 - 损失会默认对所有
d1~dk位置的样本做平均/求和(由reduction参数控制),完全匹配序列标注每个token单独分类的需求
你场景里的sequence_length就是上述的d1,属于官方支持的标准用法,不会出现计算错误。
两种等价的正确写法
写法1:维度置换(你当前在用的方案)
将类别维度置换到第2位,符合输入要求:
import torch.nn as nn loss_fn = nn.CrossEntropyLoss() # target_predictions shape: [batch_size, sequence_length, num_classes] # target shape: [batch_size, sequence_length] loss = loss_fn(target_predictions.permute(0, 2, 1), target)
写法2:维度展平
将batch维度和序列长度维度合并,转为最常见的2维输入格式:
loss_fn = nn.CrossEntropyLoss() # 展平预测值:[batch_size*sequence_length, num_classes] preds_flat = target_predictions.reshape(-1, number_of_classes) # 展平标签:[batch_size*sequence_length] target_flat = target.reshape(-1) loss = loss_fn(preds_flat, target_flat)
两种写法的计算结果完全一致,可以根据自己的代码习惯选择。
额外注意事项
如果你的序列存在填充(Padding)部分,需要忽略填充位置的损失,只需要在初始化损失函数时传入ignore_index参数即可,两种写法都可以自动生效:
# 例如填充位置的标签为-100时 loss_fn = nn.CrossEntropyLoss(ignore_index=-100)
内容的提问来源于stack exchange,提问作者Kenenbek Arzymatov
相关产品推荐
相关产品推荐

