__getitem__ 必须返回张量,因为 DataLoader 多进程下 default_collate 仅支持 tensor、标量、字典和列表;返回 np.ndarray 或 list 在遇到不规则尺寸或非数值类型时会报 TypeError。

为什么 __getitem__ 必须返回张量,而不是 NumPy 数组或 Python 列表
PyTorch 的 DataLoader 在多进程加载时会自动调用 torch.utils.data.default_collate,它只识别 torch.Tensor、Python 标量、字典和列表(含张量)。若 __getitem__ 返回 np.ndarray 或 list,遇到含不规则尺寸(如不同长度的文本 token)或非数值类型(如字符串路径)时,default_collate 会直接报错:TypeError: default_collate: batch must contain tensors, numpy arrays, numbers, dicts or lists。
实操建议:
立即学习“Python免费学习笔记(深入)”;
- 在
__getitem__内部就完成torch.tensor()转换,尤其是图像用torch.from_numpy(img_array)+.float(),标签用torch.tensor(label).long() - 若需保留原始路径或 ID 做 debug,可额外存为字典字段(如
{"image": img_tensor, "label": label_tensor, "path": self.img_paths[idx]}),但确保 collate 不会尝试堆叠"path" - 避免在
__getitem__中做耗时操作(如实时解码视频帧),应提前缓存或用torchvision.io.read_image替代 PIL 打开再转 tensor
如何安全重写 __len__:别数文件夹,要数实际可用样本
__len__ 返回的必须是 __getitem__ 可合法索引的数量。常见错误是直接 len(os.listdir("imgs/")),但目录里可能混着 .DS_Store、损坏图片、或标注缺失的样本,导致 DataLoader 迭代到某 idx 时在 __getitem__ 报错(如 FileNotFoundError 或 cv2.error),训练中途崩溃。
实操建议:
立即学习“Python免费学习笔记(深入)”;
- 初始化时就构建有效样本索引列表:
self.valid_indices = [i for i, p in enumerate(self.img_paths) if self._is_valid(p)],然后__len__返回len(self.valid_indices) -
_is_valid应轻量:检查文件是否存在、后缀是否为[".jpg", ".png"]、用cv2.imread(p) is not None粗筛(不用全解码) - 若数据集动态变化(如边训练边加新样本),不要缓存
valid_indices,改用每次__len__重新扫描——但注意这会拖慢DataLoader初始化,仅限小数据集
要不要重写 __init__?关键在路径与预处理逻辑分离
自定义 Dataset 的 __init__ 不是必须重写,但必须明确三件事:数据源在哪、怎么读、是否需要预处理。很多人把路径拼接、CSV 解析、transform 初始化全塞进 __init__,结果一换环境(Windows/macOS 路径分隔符)、一升级库(PIL.Image.open 对中文路径失败),整个 Dataset 就挂。
实操建议:
立即学习“Python免费学习笔记(深入)”;
- 路径传参强制用
pathlib.Path:接收root: Path,内部统一用root / "images"拼接,自动兼容系统差异 - CSV 或 JSON 元数据解析放在
__init__里,但用pandas.read_csv(..., nrows=10)先试读,确认列名和 dtype 正确再全读 - transform(如
torchvision.transforms.Resize)作为参数传入,**不要在__init__里硬编码**;否则无法复用同一 Dataset 实例做 train/val 不同增强
调试 __getitem__ 的最快方法:绕过 DataLoader 直接索引
当模型训练时报错 “index out of range” 或输出全是 NaN,问题大概率出在 __getitem__ 的某次调用。但 DataLoader 多进程 + shuffle 会让错误难以复现。
实操建议:
立即学习“Python免费学习笔记(深入)”;
- 写一行测试代码:
ds = MyDataset(root=Path("data")); print(ds[0]); print(type(ds[0]["image"]), ds[0]["image"].shape)—— 确保单次调用能稳定返回正确结构 - 在
__getitem__开头加assert idx = len {len(self)}",避免静默截断 - 对图像类任务,临时加
plt.imshow(ds[0]["image"].permute(1,2,0))看是否是黑图/白图(常因归一化参数错或通道顺序反)
最易被忽略的是 transform 的 inplace 参数:比如 transforms.RandomHorizontalFlip(p=1.0) 后接 transforms.ToTensor(),若前者修改了原 PIL 图像而后者又依赖其 mode,可能触发 TypeError:pic should be PIL Image or ndarray —— 这类链式副作用,只有单步索引才能快速定位。


















