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

基于PyTorch的药物活性多分类(一对一)训练IndexError问题求助

问题:One-vs-One二分类器训练时的IndexError问题

我用PyTorch处理药物活性的三分类问题(类别0、1、2),采用一对一(one vs. one)策略构建三个二分类器:0 vs. 1、1 vs. 2、2 vs. 0。但训练第二个分类器(1 vs. 2)时,出现错误:

IndexError: Target 2 is out of bounds.

有没有无需重新分配标签的解决方法?


我的GIN模型代码

class GIN1(torch.nn.Module): 
  def __init__(self, h):
     super(GIN1, self).__init__()
     dim_h_conv = h
     dim_h_fc = dim_h_conv*5

     # Convolutional layers
     self.conv1 = GINConv(Sequential(Linear(14, dim_h_conv),
                                     BatchNorm1d(dim_h_conv), ReLU(),
                                     Linear(dim_h_conv, dim_h_conv), ReLU()))
     self.conv2 = GINConv(Sequential(Linear(dim_h_conv, dim_h_conv),
                                     BatchNorm1d(dim_h_conv), ReLU(),
                                     Linear(dim_h_conv, dim_h_conv), ReLU()))
     self.conv3 = GINConv(Sequential(Linear(dim_h_conv, dim_h_conv),
                                     BatchNorm1d(dim_h_conv), ReLU(),
                                     Linear(dim_h_conv, dim_h_conv), ReLU()))
     self.conv4 = GINConv(Sequential(Linear(dim_h_conv, dim_h_conv),
                                     BatchNorm1d(dim_h_conv), ReLU(),
                                     Linear(dim_h_conv, dim_h_conv), ReLU()))
     self.conv5 = GINConv(Sequential(Linear(dim_h_conv, dim_h_conv),
                                     BatchNorm1d(dim_h_conv), ReLU(),
                                     Linear(dim_h_conv, dim_h_conv), ReLU()))
 
     # Fully connected layers
     self.lin1 = Linear(dim_h_fc, dim_h_fc)
     self.lin2 = Linear(dim_h_fc, 2)

     self.initialize_w()

  def forward(self, x, edge_index, batch):
    h1 = self.conv1(x, edge_index)
    h2 = self.conv2(h1, edge_index)
    h3 = self.conv3(h2, edge_index)
    h4 = self.conv4(h3, edge_index)
    h5 = self.conv5(h4, edge_index)

    # Graph level readout
    h1 = global_add_pool(h1, batch)
    h2 = global_add_pool(h2, batch)
    h3 = global_add_pool(h3, batch)
    h4 = global_add_pool(h4, batch)
    h5 = global_add_pool(h5, batch)

    # Concatenate graph embeddings
    h = torch.cat((h1, h2, h3, h4, h5), dim=1)

    # Classifier
    h = self.lin1(h)
    h = h.relu()
    h = F.dropout(h, p=hp_gin1['p'], training=self.training)
    h = self.lin2(h)
    h = F.log_softmax(h, dim=1)

    return h

  def initialize_w(self):
    for m in self.modules():

      if isinstance(m, Linear):
        torch.nn.init.kaiming_normal_(m.weight, mode='fan_in', nonlinearity='relu')
        torch.nn.init.constant_(m.bias, 0)

      if isinstance(m, BatchNorm1d):
        torch.nn.init.constant_(m.weight, 1)
        torch.nn.init.constant_(m.bias, 0)

训练循环代码

gin2 = GIN2(h=hp_gin2['h']) #40
optimizer = torch.optim.Adam(gin2.parameters(), lr=hp_gin2['lr']) 
criterion = torch.nn.CrossEntropyLoss()

def train(train_loader):
  gin2.train()
  loss_all = 0
  for data in train_loader:
    output = gin2(data.x, data.edge_index, data.batch)
    loss = criterion(output, data.y)

    l2_lambda = hp_gin2['lambda']
    l2_norm = sum(p.pow(2.0).sum()
                 for p in gin2.parameters())
    loss = loss + l2_lambda * l2_norm

    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    loss_all += loss.item() * data.num_graphs
  return loss_all / len(train_loader.dataset)

def test_loss(loader):
  total_loss_val = 0
  with torch.no_grad():
   for data in loader:
     output = gin2(data.x, data.edge_index, data.batch)
     batch_loss = criterion(output, data.y)
  
     total_loss_val += batch_loss.item() * data.num_graphs
  return total_loss_val / len(loader.dataset)

def test(loader):
  gin2.eval()
  correct = 0
  for data in loader:
    output = gin2(data.x, data.edge_index, data.batch)

    accuracy = Accuracy(average='macro', num_classes=2)
    acc = accuracy(output, data.y)

  return acc

解决方法

核心原因

二分类器输出维度为2(对应类别索引0和1),但训练1 vs. 2时,样本标签包含值为2的类别。CrossEntropyLoss会将标签当作类别索引处理,要求索引必须小于输出的类别数(2),因此标签2超出范围导致报错。

无需重设标签的两种方案

方案1:临时映射标签适配CrossEntropyLoss

在计算损失前,临时将标签转换为0/1的二分类索引,不修改原始数据集标签:

# 训练1 vs. 2时,在计算loss前添加以下代码
target = torch.where(data.y == 1, torch.tensor(0, device=data.y.device), torch.tensor(1, device=data.y.device))
loss = criterion(output, target)

该代码会把原始标签1转为0,标签2转为1,完全适配二分类器的输出维度,且不会改动原始数据的标签值。

方案2:改用BCEWithLogitsLoss(推荐)

修改模型输出逻辑,配合二元交叉熵损失,彻底规避标签索引范围问题:

  1. 修改模型forward方法:移除最后一层的log_softmax,直接输出原始logits
# 替换模型forward中的最后两行
h = self.lin2(h)
# 去掉 h = F.log_softmax(h, dim=1)
return h
  1. 训练时动态生成二元标签:
# 训练1 vs. 2时,调整损失计算逻辑
criterion = torch.nn.BCEWithLogitsLoss()
# 将标签转为二元浮点数:类别1对应0,类别2对应1
target = (data.y == 2).float()
loss = criterion(output.squeeze(), target)

这种方式不依赖类别索引的连续性,灵活性更高,完全避免了标签范围错误。


内容的提问来源于stack exchange,提问作者Gianmarco Luchetti

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 10:50:24