PyTorch训练循环无法自动停止:指定300样本后仍持续运行
问题根源及修复方案
1. DataLoader初始化错误
你创建train_loader时,传入的是CustomImageDataset类而非实例,正确做法是传入已初始化好的Data实例:
# 错误写法 train_loader=DataLoader(CustomImageDataset,batch_size, shuffle = False, drop_last= True) # 正确写法 train_loader=DataLoader(Data,batch_size, shuffle = False, drop_last= True)
2. 训练循环参数传递错误
执行训练时,你把Data(Dataset实例)直接传给train_loop的dataloader参数,应该传入train_loader:
# 错误写法 train_loop(Data,model,loss_fn,optimizer,batch_size) # 正确写法 train_loop(train_loader,model,loss_fn,optimizer,batch_size)
3. CustomImageDataset的__getitem__逻辑错误
Dataset的设计逻辑是每个__getitem__返回单个样本,DataLoader会自动打包成batch。你当前的写法让每个__getitem__返回batch_size个样本,会导致:
- 索引
idx+x不断超出设定的300条数据范围,甚至越界到CSV的42000条之后 - DataLoader会持续从Dataset取数据,实际处理样本数远超预期
修复后的Dataset:
class CustomImageDataset(Dataset): def __init__(self, img_dir): self.img_dir = img_dir self.data=pd.read_csv(self.img_dir) def __len__(self): return 300 # 仅取前300个样本 def __getitem__(self, idx): # 读取单个样本标签 label = self.data.at[idx,'label'] label = label.item() # 初始化单个图像张量 image = torch.zeros(1,28,28) for i in range(784): z = int(i%28) y = int(i/28) # 修正原错误的索引计算 column = 'pixel' + str(i) image[0,z,y] = self.data.at[idx, column] return image, label
4. 训练循环样本数计算优化
原train_loop中size可通过dataloader.dataset.__len__()动态获取,避免硬编码:
def train_loop(dataloader, model, loss_fn, optimizer): size = dataloader.dataset.__len__() for batch, (X, y) in enumerate(dataloader): pred = model(X) loss = loss_fn(pred, y) optimizer.zero_grad() loss.backward() optimizer.step() if batch % 100 == 0: loss_val, current = loss.item(), (batch + 1) * len(X) print(f"loss: {loss_val:>7f} [{current:>5d}/{size:>5d}]")
5. 其他细节修正
- 初始化Dataset时无需传入
batches参数,batch逻辑由DataLoader处理 - 原代码中
y = int((i-x)/28)是错误的,修正为y = int(i/28),确保像素位置映射正确
内容的提问来源于stack exchange,提问作者tridentifer
相关产品推荐
相关产品推荐

