跳转到内容
新建笔记

PyTorch:set_epoch 属性错误与分布式采样

AttributeError: 'PoseDataset' object has no attribute 'set_epoch' 表示代码在错误对象上调用了方法。原记录还关联到 LaneGCN 的 RandomSampler 同名报错;这两种对象都不能仅凭“参与训练”就假定有 set_epoch。

对象职责
Dataset提供样本及访问方式;自定义数据集是否实现额外方法要看代码
Sampler决定访问样本索引的次序
DistributedSampler为多个分布式进程划分索引,支持按 epoch 更新随机打乱
DataLoader组合数据集、采样器、批处理和加载进程

先打印类型并查找报错调用处:

print(type(train_loader.dataset).__name__)
print(type(train_loader.sampler).__name__)

若使用标准 DistributedSampler,应在每轮创建 DataLoader 迭代器之前调用采样器的方法:

from torch.utils.data.distributed import DistributedSampler
for epoch in range(num_epochs):
sampler = train_loader.sampler
if isinstance(sampler, DistributedSampler):
sampler.set_epoch(epoch)
for batch in train_loader:
train_step(batch)

这里的 train_loader、num_epochs 和 train_step 是训练程序已经定义的对象。单进程普通 RandomSampler 不需要照抄这一调用;自定义 sampler 或 batch_sampler 则按其接口处理。

为什么不能只把报错行删掉

跳转到“为什么不能只把报错行删掉”

删除不适用的 dataset.set_epoch(epoch) 能消除属性错误,但若训练确实使用分布式随机采样,还应把调用放到正确 sampler 上,否则各 epoch 的打乱次序可能重复。反过来,给所有 Dataset 强行增加空方法也会掩盖配置错误。

验证时同时检查训练是否继续、实际使用哪种 sampler,以及不同 epoch 的索引顺序;不要把“没有报错”当作采样逻辑正确的全部证据。

参考:PyTorch 数据加载文档。历史问题线索:LaneGCN issue 26,它描述具体项目情境,不构成适用于所有训练代码的删除规则。