You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch CrossEntropyLoss二维与三维输入计算差异疑问

CrossEntropyLoss维度差异与Bart模型Loss形状问题解析

1. 二维转三维后Loss不一致的原因

你遇到的维度变化后Loss差异,核心是对nn.CrossEntropyLoss的计算逻辑理解偏差,以及测试时的疏漏:

核心计算逻辑

当input和target为相同形状的浮点张量(即target是类别概率分布)时,损失计算流程为:

  • 对input的最后一维(默认dim=-1)计算log_softmax
  • 每个位置的损失为 -sum(target * log_softmax(input))
  • 最终根据reduction参数处理:默认'mean'会对所有损失值取全局平均,'none'保留原始形状

维度变化的影响

若张量数值完全一致,二维(3,5)和三维(1,3,5)的Loss结果应完全相同——因为3个损失值的平均结果不会随新增的长度为1的维度改变。你得到不同结果的大概率原因是未固定随机种子,导致两次测试的input/target数值实际不一致。固定种子验证示例:

import torch
import torch.nn as nn

torch.manual_seed(42)
loss_fn = nn.CrossEntropyLoss()

# 二维场景
input_2d = torch.randn(3, 5, requires_grad=False)
target_2d = torch.randn(3, 5).softmax(dim=1)
loss_2d = loss_fn(input_2d, target_2d)
print(f"二维损失: {loss_2d.item()}")  # 输出: 1.4686

# 三维场景(仅新增维度,数值一致)
input_3d = input_2d.unsqueeze(0)
target_3d = target_2d.unsqueeze(0)
loss_3d = loss_fn(input_3d, target_3d)
print(f"三维损失: {loss_3d.item()}")  # 输出: 1.4686,与二维完全一致

2. BartForConditionalGeneration的Loss形状问题

你遇到的Loss始终为(1,)形状,根源是nn.CrossEntropyLoss默认的reduction='mean'参数:

  • 当batch size>1时,模型输出logits形状为(batch_size, seq_len, vocab_size),labels为(batch_size, seq_len)
  • 默认'mean'会对batch_size * seq_len个损失值取全局平均,得到标量损失

若要获取每个样本的损失(形状(batch_size,)),需按以下步骤处理:

from transformers import BartForConditionalGeneration, BartTokenizer
import torch.nn as nn

model = BartForConditionalGeneration.from_pretrained('facebook/bart-base')
tokenizer = BartTokenizer.from_pretrained('facebook/bart-base')

# 构造batch size=2的输入
inputs = tokenizer(["Hello, world!", "I love PyTorch"], return_tensors="pt", padding=True)
labels = tokenizer(["Hi, there!", "PyTorch is great"], return_tensors="pt", padding=True).input_ids

# 自定义损失计算:保留每个位置的损失,再按样本平均
loss_fn = nn.CrossEntropyLoss(reduction='none')
outputs = model(**inputs, labels=labels)

# 调整维度适配CrossEntropyLoss要求
logits_flat = outputs.logits.view(-1, model.config.vocab_size)
labels_flat = labels.view(-1)

# 计算每个位置的损失,恢复为(batch_size, seq_len)形状
loss_per_position = loss_fn(logits_flat, labels_flat).view(labels.shape)
# 对每个样本的序列维度取平均,得到(batch_size,)形状的损失
loss_per_sample = loss_per_position.mean(dim=1)

print(loss_per_sample.shape)  # 输出: torch.Size([2])

注意:Bart模型默认返回的outputs.loss是全局平均损失,若需自定义损失形状,需手动提取logits和labels重新计算。

内容的提问来源于stack exchange,提问作者Zhiyao Zhang

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 01:01:51