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

如何让one-hot独热编码输入与非独热目标值适配,解决PyTorch二分类损失报错?

问题核心原因
  • 直接触发报错的原因是nn.BCELoss调用方式错误:PyTorch中nn模块下的损失函数都是类,需要先实例化再传入预测值、标签计算损失,你直接将两个张量传入nn.BCELoss的构造函数,触发了类型校验的逻辑错误。
  • 除此之外代码还存在两处逻辑错误会导致后续训练失败:
    1. 输入张量维度不匹配:你的输入x_train形状为[样本数, 队员总数, 角色数],是三维张量,但全连接层nn.Linear要求输入最后一维为特征维度,需要先将每个样本的所有角色one-hot向量展平为一维特征。
    2. Sigmoid激活重复调用:你在layer2的最后已经加了nn.Sigmoid()层,前向传播的返回值又套了一层torch.sigmoid(),相当于对输出做了两次sigmoid压缩,会导致梯度消失、模型无法收敛。
可行修正方案

按照以下步骤调整代码即可正常运行:

  1. 先将输入特征展平,调整全连接层输入维度
  2. 修正损失函数的调用逻辑
  3. 删除重复的Sigmoid激活
  4. 补充训练循环的反向传播、参数更新逻辑

修正后的完整代码示例:

import torch
import torch.nn as nn
import torch.optim as optim

# 示例参数
num_characters = 4
team_member_per_side = 2
total_member = team_member_per_side * 2

# 输入输出数据示例
x_data = [ [[0,0,1,0], [0,1,0,0], [1,0,0,0],[0,1,0,0]], [[1,0,0,0], [0,0,0,1], [0,1,0,0], [0,0,1,0]] ]
y_data = [[0], [1]] # team blue win: 1, lose : 0

# 处理输入:将每个样本的所有角色one-hot展平为一维
x_train = torch.FloatTensor(x_data).flatten(start_dim=1) # 形状变为 [样本数, total_member * num_characters]
y_train = torch.FloatTensor(y_data)

class BinaryClassifier(nn.Module):
    def __init__(self):
        super(BinaryClassifier, self).__init__()
        self.layer1 = nn.Sequential(
            # 输入维度调整为展平后的特征长度
            nn.Linear(in_features=total_member * num_characters, out_features=10, bias=True),
            nn.ReLU(), 
            )
        self.layer2 = nn.Sequential(
            nn.Linear(in_features=10, out_features=1, bias=True),
            nn.Sigmoid(), # 这里已经做了sigmoid,后续不需要重复调用
            )
    
    def forward(self, x):
        x = self.layer1(x) 
        x = self.layer2(x)
        return x # 删掉额外的sigmoid

model = BinaryClassifier()
optimizer = optim.SGD(model.parameters(), lr=0.01) # 建议调低学习率,原设置的1过大容易不收敛
# 先实例化损失函数
criterion = nn.BCELoss()

nb_epochs = 1000
for epoch in range(nb_epochs + 1):
    optimizer.zero_grad() # 梯度清零
    hypothesis = model(x_train)
    # 调用实例化后的损失函数计算loss
    cost = criterion(hypothesis, y_train)
    # 反向传播+参数更新
    cost.backward()
    optimizer.step()

    if epoch % 100 == 0:
        print(f'Epoch {epoch:4d}/{nb_epochs} Cost: {cost.item():.6f}')

内容的提问来源于stack exchange,提问作者Ingyu Seo

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.23 16:54:02