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

PyTorch实现图像A+B神经网络遇IndexError问题求助

问题解决:图像A+B神经网络的IndexError修复与优化

错误原因分析

你遇到的IndexError: index 4 is out of bounds for dimension 0 with size 4核心问题出在模型与训练代码的两处错误:

  1. 模型维度逻辑错误:输入input1/input2的维度是[batch_size, 3072](batch_size=4),a[i]表示取第i个样本,你通过range(3072)循环会尝试访问索引4到3071的样本,但批次内仅包含4个样本(索引0-3),直接触发越界错误;且列表推导式生成的是Python列表,不是PyTorch张量,无法传入全连接层计算。
  2. 训练代码冗余操作:trainloader输出的im1/im2/mid本身就是张量,无需重复执行torch.from_numpy(np.array(...))转换,这会浪费计算资源且可能破坏设备一致性。

修复后的模型代码

import torch.nn as nn

class myNet(nn.Module):
    def __init__(self):
        super(myNet, self).__init__()
        
        self.fc1 = nn.Linear(3072, 3072)  
        self.fc2 = nn.Linear(3072, 3072)  
        self.fc3 = nn.Linear(3072, 3072)  

    def forward(self, input1, input2):
        a = self.fc1(input1)
        b = self.fc2(input2)
        # 直接用张量逐元素相加,自动保留batch维度,无需手动循环
        combined = a + b
        out = self.fc3(combined)
        return out

优化后的训练代码

num_epochs = 10
losses = []
batch_size = 4

for epoch in range(num_epochs):
    for i, (im1, im2, mid) in enumerate(trainloader):
        # 归一化到[0,1]并转换为浮点型
        im1 = im1.float() / 255.0
        im2 = im2.float() / 255.0
        mid = mid.float() / 255.0
        
        # 移至目标计算设备
        im1 = im1.to(device)
        im2 = im2.to(device)
        mid = mid.to(device)
       
        # 前向传播
        outputs = model(im1, im2)
        
        # 计算损失
        loss = criterion(outputs, mid)
        losses.append(loss.item())

        # 反向传播与参数更新
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
        
        # 打印训练日志
        if (i + 1) % 50 == 0:
            print('Epoch [%2d/%2d], Step [%3d/%3d], Loss: %.4f'
                  % (epoch + 1, num_epochs, i + 1, len(trainloader), loss.item()))

额外优化建议

  • 替换为卷积层:全连接层会丢弃图像的空间结构,对于图像任务,使用nn.Conv2d卷积层更合理,既能保留像素间的空间关联,又能减少参数量。
  • 添加激活函数:在全连接层后加入nn.ReLU()等激活函数,增加模型的非线性表达能力,避免模型退化为简单线性映射。
  • 监控验证损失:加入验证集循环,定期计算验证损失,及时发现过拟合问题。
  • 简化张量操作:可使用链式调用im1 = im1.float().div_(255).to(device),让代码更简洁高效。

内容的提问来源于stack exchange,提问作者Ali.A

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 07:34:59