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

使用ViT训练二分类图像模型时遇ValueError数组均分错误

问题排查:ViT二分类训练报错ValueError: array split does not result in an equal division

问题背景

使用ViT进行图像二分类(类别0代表false、1代表true),设置batch size=32、epochs=3,训练时触发报错:

ValueError: array split does not result in an equal division

报错代码行:x = np.split(np.squeeze(np.array(x)), BATCH_SIZE)

训练脚本如下:

import torch.utils.data as data
from torch.autograd import Variable
import numpy as np
train_loader = data.DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,  num_workers=2)
test_loader  = data.DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=2) 

# Train the model
for epoch in range(EPOCHS):        
  for step, (x, y) in enumerate(train_loader):
    # Change input array into list with each batch being one element
    x = np.split(np.squeeze(np.array(x)), BATCH_SIZE)
    # Remove unecessary dimension
    for index, array in enumerate(x):
      x[index] = np.squeeze(array)
    # Apply feature extractor, stack back into 1 tensor and then convert to tensor
    x = torch.tensor(np.stack(feature_extractor(x)['pixel_values'], axis=0))
    # Send to GPU if available
    x  = x.to(device)
    y = y.to(device)
    b_x = Variable(x)   # batch x (image)
    b_y = Variable(y)   # batch y (target)
    # Feed through model
    output = model(b_x, None)
    loss = output[0]
    # Calculate loss
    if loss is None: 
      loss = loss_func(output, b_y)   
      optimizer.zero_grad()           
      loss.backward()                 
      optimizer.step()

    if step % 50 == 0:
      # Get the next batch for testing purposes
      test = next(iter(test_loader))
      test_x = test[0]
      # Reshape and get feature matrices as needed
      test_x = np.split(np.squeeze(np.array(test_x)), BATCH_SIZE)
      for index, array in enumerate(test_x):
        test_x[index] = np.squeeze(array)
      test_x = torch.tensor(np.stack(feature_extractor(test_x)['pixel_values'], axis=0))
      # Send to appropirate computing device
      test_x = test_x.to(device)
      test_y = test[1].to(device)
      # Get output (+ respective class) and compare to target
      test_output, loss = model(test_x, test_y)
      test_output = test_output.argmax(1)
      # Calculate Accuracy
      accuracy = (test_output == test_y).sum().item() / BATCH_SIZE
      print('Epoch: ', epoch, '| train loss: %.4f' % loss, '| test accuracy: %.2f' % accuracy)

错误原因

  1. 批量数不匹配:当数据集总样本数无法被batch size整除时,PyTorch DataLoader的最后一批会返回剩余的不足批量的样本(比如总样本100,batch size32,最后一批仅4个样本)。此时用固定值BATCH_SIZE做np.split,会因数组长度无法被均分报错。
  2. 维度错误:np.squeeze会错误移除批量维度,比如原本形状为(32, C, H, W)的张量,squeeze后变成(C, H, W),彻底破坏了批量结构,导致拆分逻辑完全失效。

解决方案

1. 修复批量处理逻辑

直接根据当前批次的实际样本数拆分,而非固定BATCH_SIZE,同时保留样本维度:

修改训练循环内的处理代码:

for step, (x, y) in enumerate(train_loader):
    # 转numpy并按实际样本数拆分
    x_np = x.cpu().numpy()
    x_list = [x_np[i] for i in range(len(x_np))]
    # 应用特征提取器
    pixel_values = feature_extractor(x_list)['pixel_values']
    x = torch.tensor(np.stack(pixel_values, axis=0))
    
    # 后续设备转移、模型输入逻辑不变
    x = x.to(device)
    y = y.to(device)
    b_x = Variable(x)
    b_y = Variable(y)
    # ... 其余训练代码

2. 同步修复测试代码

测试部分的批量处理逻辑和训练端保持一致:

if step % 50 == 0:
    test_x, test_y = next(iter(test_loader))
    # 按实际样本数拆分
    test_x_np = test_x.cpu().numpy()
    test_x_list = [test_x_np[i] for i in range(len(test_x_np))]
    pixel_values = feature_extractor(test_x_list)['pixel_values']
    test_x = torch.tensor(np.stack(pixel_values, axis=0))
    
    test_x = test_x.to(device)
    test_y = test_y.to(device)
    # ... 其余测试代码

3. 修正准确率计算

因为最后一批样本数可能不是32,计算准确率时要用实际样本数:

accuracy = (test_output == test_y).sum().item() / len(test_y)

可选:强制批量大小一致(不推荐)

如果必须保证每个批次都是32样本,可以给DataLoader添加drop_last=True参数,丢弃最后一批不足批量的样本(会损失部分数据):

train_loader = data.DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=2, drop_last=True)
test_loader  = data.DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=2, drop_last=True)

内容的提问来源于stack exchange,提问作者Rachele Franceschini

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 23:40:28