深度学习必备技能:够用就行的数据Dataset、DataLoader
一、为什么需要使用Dataset与DataLoader处理数据?
1.深度学习数据规模巨大,把所有数据一次性读入内存,这很可能导致内存溢出(OOM)
Dataset为数据集建立索引,加载部分内容,
DataLoader决定加载多少内容
2.提升训练效率:
自动批处理:DataLoader 可以轻松地将多个独立样本“堆叠”成一个批次。
多进程加速:I/O速度显著低于CPU速度的情况下,开启多个子进程,让 CPU 在 GPU 计算当前批次的同时,提前去硬盘加载并预处理下一批数据。GPU 处于满负荷工作状态,而不是干等着数据。
3.代码模块化:解耦数据逻辑与训练逻辑
二、数据流
数据->Dataset->DataLoader->模型
1.数据->Dataset
Dataset 从磁盘中读取单个数据,返回Python 原生的数值、列表、NumPy 数组或 PIL Image,并不是张量
- len(self):返回数据集总大小。
- getitem(self, idx):根据索引 idx 返回一个样本(通常是一个元组 (data, label))
2.Dataset->Dataloader
DataLoader是一个数据加载器,通过分配系统资源,指定batch大小来使用Dataset从磁盘中读取数据,以确保系统资源得到充分利用
DataLoader(dataset, batch_size=32, shuffle=True, num_workers=4, collate_fn=None)
参数解析:
- dataset指定数据集
- shuffle 生成一个索引迭代器
- num_workers 使用多个cpu进程,并行调用 dataset[idx] 获取每一个独立样本。
- collate_fn 将样本列表(长度为 batch_size)合并成一个“批次”数据,并转换为tensor的函数。
- 流程:
- 采样:DataLoader 根据 batch_size 和 shuffle 等参数,决定要取哪些样本的索引,。
- 获取单个样本:对于每个索引 i,调用 dataset[i](即 getitem 方法),得到一个样本(如 (encoded_sequence, metadata) 元组)。
- 收集 batch:将获取的 batch_size 个样本,将它们放入一个 Python 列表 batch 中,列表长度 = batch_size,每个元素就是一次 getitem 返回的结果
- 调用 collate_fn:将刚刚得到的 batch 列表作为唯一参数传递给 collate_fn(batch)。
- 返回 batch 数据:collate_fn 返回处理后的结果(通常是张量或张量的元组/字典),这就是你每次迭代拿到的 batch。
3.Dataloader->模型
使用for循环逐个读取dataloader打包好的batch数据
for batch_data, batch_labels in dataloader:
三、一般骨架
主要分两个部分:
1.建立dataset对象,实现三个函数 :
- init:初始化,不加载具体数据
- len:获取总数据量
- getitem:加载单个数据
2.数据加载部分-位于dataset外部
- get_dataloader:确定batch的规则,是否使用子集,如何打包batch
- collate_fn:确定打包数据的规则
#Dataset.py
from torch.utils.data import Dataset, DataLoader, Subset
class MyDataset(Dataset):
def __init__(self,path,transform):
super().__init__()
'''
根据路径加载文件目录
'''
def __len__(self):
'''
获取数据集总大小
返回数据集总大小
'''
return
def __getitem__(self,idx):
'''
根据索引获取其中一项
返回单个数据
'''
return
def my_collate_fn(batch):
"""
自定义批量收集函数,将序列和元数据打包,并对序列进行padding
Args:
batch: 包含序列和元数据的元组列表
Returns:
包含打包后的序列和元数据的元组
"""
data,label= zip(*batch) #直接使用迭代器
#或
data = list(zip(*batch)) #getitem只有单个对象
"""
数据格式转换为torch
数据处理如padding
"""
return processed_data
def get_dataloader(data_dir, batch_size=1, shuffle=False, transform=None, subset=None):
"""
创建数据加载器
Args:
data_dir: 数据目录
batch_size: 批大小
shuffle: 是否打乱数据
transform: 数据转换函数
subset: 数据集子集大小,如果为None则使用全部数据
Returns:
DataLoader实例
"""
dataset = MyDataset(data_dir, transform)
# 如果指定了subset,创建数据集子集
if subset is not None and subset > 0:
subset_indices = list(range(min(subset, len(dataset))))
dataset = Subset(dataset, subset_indices)
print(f"Using subset of dataset with {len(dataset)} samples")
return DataLoader(dataset, batch_size=batch_size, shuffle=shuffle, collate_fn=my_collate_fn,num_workers=1)
四、Dataset详解
getitem的返回值可以是一个或者多个,如果只有一个,那么返回值的类型是变量本身,如果是多个,则为元组,这是由于python会将多个返回值打包为一个元组,好处:
- 便利性:避免手动构造元组,代码更简洁。
- 解包方便:调用方可以直接写 img, lbl = dataset[0] 进行解包。
- 一致性:与函数参数中的 *args 等机制保持统一。
五、collate_fn详解
1.默认collate_fn
当你创建一个 DataLoader 且不指定 collate_fn 时,PyTorch 会使用默认的 default_collate 函数。它的逻辑很简单:
- 假设每个样本是一个等长的张量、数值、字符串或元组/列表/字典。
- 将同类型的多个样本沿第 0 维(batch 维)堆叠起来,形成一个更大的张量或嵌套结构。
例如:有 N 个形状为 [C, H, W] 的图像张量,堆叠后得到形状 [N, C, H, W]
2.迭代器
迭代: 重复地从一个数据源中取出元素的过程,如for i in [1,2,3]:
- 迭代器协议: 一个对象要成为迭代器,必须实现两个方法:
- iter(self):返回自身(return self)
- next(self):返回下一个可用元素,如果没有元素了,抛出 StopIteration
- 可迭代对象: 实现了 iter() 方法的对象,该方法返回一个迭代器
- 迭代器: 实现了 iter() 和 next() 方法的对象,next() 每次返回下一个元素,无元素时抛出 StopIteration
- 迭代器特性:
- 惰性求值(Lazy Evaluation): 迭代器不会在创建时就把所有元素计算出来,而是每次调用 next() 才产生一个元素。这对处理无限序列或大文件非常友好。
- **一次性消耗:**迭代器一旦遍历完,就“耗尽”了,再次遍历不会得到任何元素(会立即抛出 StopIteration)。如果需要重新遍历,必须重新创建迭代器。
- for循环本质–自动调用迭代器的iter与next方法伪代码
it = iter(iterable) # 获得迭代器
while True:
try:
x = next(it) # 不断取下一个
# 循环体
except StopIteration:
break
- list(iterable)调用迭代器伪代码
def list(iterable):
# 1. 获取可迭代对象的迭代器
it = iter(iterable)
# 2. 准备一个空列表用于存放结果
result = []
# 3. 不断调用 next(it) 直到 StopIteration
while True:
try:
value = next(it)
except StopIteration:
break
result.append(value)
# 4. 返回填充好的列表
return result
- tuple(iterable)调用迭代器伪代码
#伪代码
def tuple(iterable):
it = iter(iterable)
result = [] # 同样先用列表收集
while True:
try:
value = next(it)
except StopIteration:
break
result.append(value)
return tuple(result) # 几乎一样,最后转为元组
3.zip迭代器
zip() 是 Python 内置函数,它接受多个可迭代对象作为参数,返回一个迭代器。这个迭代器生成元组,每个元组包含来自每个输入可迭代对象的第 i 个元素。
zip 迭代器的行为
- 惰性:zip(a, b) 不会立即生成所有元组,只在遍历时逐个生成。
- 当最短的输入可迭代对象耗尽时停止迭代,且不会报错。
- 它遵循标准迭代器协议:可以传给 next(),可以用 for 循环,可以转换为列表/元组。
- 转至效果
matrix = [[1, 2, 3],
[4, 5, 6],
[7, 8, 9]]
transposed = zip(*matrix) # 等价于 zip([1,2,3], [4,5,6], [7,8,9])
print(list(transposed)) ''' [(1,4,7),
(2,5,8),
(3,6,9)]'''
zip迭代器的内存优势
- 当处理多个列表时,list(zip(a, b)),会一次性生成所有元组,占用 O(N) 内存,zip 迭代器内部不存储数据,只存储对输入可迭代对象的引用,每次 next() 时从每个输入中取下一个元素。
- 功能上:本框架中可以不使用zip(),因为送进来的batch是一个list
- 风格上:推荐使用,因为它让代码更清晰、更 Pythonic,轻松实现转至。
- 性能上:zip(*batch) 只产生迭代器,不额外分配大列表(与手动列表推导相比,差异很小,可忽略)。
4.自定义collate_fn
当你的样本不是单纯的数值张量,或者样本大小不一(如变长序列、不同尺寸的图像),或者包含无法自动堆叠的对象,你需要主动实现一个collate_fn,当样本是字典、自定义对象、或需要动态 pad ,默认 default_collate 会尝试 torch.stack,遇到非 Tensor 或形状不一致就会报错。
- 输入: 一个列表,长度为batch_size,列表中的每个元素是__getitem__的返回对象,如果__getitem__返回多个对象,每一项的类型是一个元组,元组内是每个对象
- 解包: 使用迭代器zip(*batch)将每个元组的第i个元素放在一起
- zip迭代器: 不存储数据,而是按需生成数据
- __getitem__只有一个返回值:a=list(zip(*batch))会遍历迭代,得到一个XX
- __getitem__有多个返回值:a=zip(*batch)不会遍历,只得到一个迭代器,类型为
- __getitem__有多个返回值时,x1,x2,x3 = zip(*batch)将遍历迭代器
- 由于zip迭代器存在因此读取数据的时候坐标的下标会发生互换,即转至
*x = list(zip(batch)) batch[0][1]–>> x[1][0]
| batch | 输出 | |
|---|---|---|
| [(a1,b1,c1), (a2, b2,c2), (a3,b3,c3)] |
----zip(*batch) ---- 》 | [a1,a2,a3] = x1 [b1,b2,b3] = x2 [c1,c2,c3] = x3 |
- 类型转换(可选): 模型训练的时候通常会使用tensor类型,需要将数据转换为torch.tensor
- **torch.Tensor.dtype:**表示张量中每个元素的数据类型
-
- 常见类型:
| 类别 | dtype | 说明 |
|---|---|---|
| 浮点 类型 |
torch.float32 或 torch.float torch.float64 或 torch.doubletorch.float16torch.bfloat16 |
单精度浮点数(32位)默认类型 双精度浮点数(64位) 半精度浮点数(16位,常用于混合精度训练) Brain 浮点格式(16位,指数与 float32 相同,精度较低) |
| 整数类型 | torch.int32 或 torch.inttorch.int64 或 torch.longtorch.int16 或 torch.shorttorch.int8torch.uint8 |
有符号整数(32位) 有符号长整数(64位),常用于索引 有符号短整数(16位) 有符号8位整数 无符号8位整数(常用于图像像素 0-255) |
| 布尔类型 | torch.bool |
布尔值 True / False |
| 其他专用类型 | torch.complex64torch.complex128 |
单精度复数(两个 float32) 双精度复数(两个 float64) |
- padding: 将所有序列补全到长度一致
- pad_sequence(sequences, batch_first=True, padding_value=0) – pytorch自带,支持多维
- ** 输入:sequences 必须是 Tensor 列表(每个 Tensor 形状任意,但通常是 1D 或 2D),且所有 Tensor 的维度相同(除序列长度维外)。
- 输出:直接返回一个堆叠好的张量,形状为 (batch, max_len, …)(若 batch_first=True)或 (max_len, batch, …)(若 batch_first=False)。
- 性能:底层用 C 实现,比手动循环快得多,尤其在大 batch 时。
- 灵活性:通过 padding_value 指定填充值,但不支持左填充或复杂填充模式(需自己写)。
- 额外功能:自动按输入顺序填充,不改变序列相对顺序。
- pad_sequence(sequences, batch_first=True, padding_value=0) – pytorch自带,支持多维
from torch.nn.utils.rnn import pad_sequence
padded_seqs = pad_sequence(sequences, batch_first=True, padding_value=0)
- 自定义padding(一维为例)
- torch.from_numpy(ndarray)
将 NumPy 数组 (ndarray) 转换为 PyTorch 张量 (Tensor)。- 共享内存:转换后的张量与原始 NumPy 数组共享同一块数据内存,修改其中一个会影响另一个。
- 输入:NumPy 数组ndarray (numpy.ndarray)。
- 输出:返回一个 torch.Tensor,与输入的 NumPy 数组具有相同的数据类型和形状。
- 位置:返回的张量默认位于 CPU 内存。若需移到 GPU,需接着调用 .cuda() 或 .to(device)
- 数据类型映射:np.int32 → torch.int32,np.float64 → torch.float64,np.bool → torch.bool 等。
当 NumPy 数组是只读或非连续内存时,行为可能受限,通常仍能工作但会复制数据。
- torch.zeros(*size, out=None, dtype=None, layout=torch.strided, device=None,requires_grad=False)
创建一个全零的张量,形状由参数指定。可以指定数据类型(dtype)、设备(device)等。- *size (int 或 tuple of ints): 张量的形状。例如 3 或 (2, 3)。
- dtype (torch.dtype, 可选): 张量的数据类型,默认为 torch.float32。
- device (torch.device, 可选): 张量所在的设备(CPU/GPU),默认为当前默认设备。
- requires_grad (bool, 可选): 是否需要计算梯度,默认为 False。
- out (Tensor, 可选): 输出张量(通常不手动使用)。
- layout (torch.layout, 可选): 内存布局,默认为 torch.strided。
- 不指定 dtype 时默认为 torch.float32,与 torch.tensor(0.0) 一致。
- 创建大尺寸张量时,全零填充是高效的(不涉及随机数生成)。
- 配合 requires_grad=True 可用于需要梯度的参数初始化(但通常用随机初始化)。
- torch.cat(tensors, dim=0, *, out=None)
将多个张量沿着指定的现有维度进行拼接(concatenate)。不会增加新维度,只是将多个张量在某个维度上堆叠起来。- tensors (sequence of Tensors): 需要拼接的张量序列(列表或元组)。所有张量在非拼接维度上的形状必须相同。
- dim (int, 可选): 沿着哪个维度进行拼接。默认为 0(第一维,通常指 batch 维或序列长度维)。
- out (Tensor, 可选): 输出张量(通常不手动使用)。
- torch.from_numpy(ndarray)
#获取最长序列长度
max_len = max(len(seq) for seq in sequences)
padded_sequences = []
for seq in sequences:
pad_len = max_len - len(seq)
seq_tensor = torch.from_numpy(seq) # 假设 seq 是 numpy 数组,转换为tensor
pad_tensor = torch.zeros(pad_len, dtype=seq_tensor.dtype)
padded_seq = torch.cat([seq_tensor, pad_tensor], dim=0)
padded_sequences.append(padded_seq)
- 输出: 任何你想要的数据结构,但通常是一个或多个张量(如 (batch_imgs, batch_labels)),长度与形状一致,可以直接送给模型。
六、DataLoader详解
torch.utils.data.DataLoader 是 PyTorch 中数据加载的核心组件。它把一个 Dataset 对象变成一个可迭代的批次生成器。
1.核心功能
| 功能 | 说明 |
|---|---|
| 批次迭代 | 将多个样本组合成 batch,供模型一次处理 |
| 自动打乱 | 每个 epoch 可以随机重排样本顺序,防止模型记住顺序 |
| 多进程加载 | 通过 num_workers 参数开启多个子进程,并行加载数据,掩盖 I/O 延迟 |
| 自定义批处理 | 通过 collate_fn 自定义如何把多个样本拼接成一个 batch(例如处理变长序列) |
| 自动内存管理 | 支持 pin_memory 将数据锁页,加速 CPU→GPU 传输 |
2. 常用参数
| 参数 | 含义 |
|---|---|
| dataset | 一个 Dataset 实例(或 Subset 包装的实例) |
| batch_size | 每个 batch 的样本数 |
| shuffle | 每个 epoch 开始时是否打乱索引顺序 |
| collate_fn函数 | 接收一个 batch 列表(长度为 batch_size),返回处理后的 batch 数据 |
| num_workers | 子进程数量(0 表示主进程加载。与CPU数保持一致) |
| drop_last | 如果样本总数不能被 batch_size 整除,是否丢弃最后一个不完整的 batch |
| pin_memory | 是否将张量复制到锁页内存(Page-Locked Memory / Pinned Memory),加快 GPU 传输,默认False |
num_workser 在 Windows 上必须把训练循环放在 if name == ‘main’: 保护下,否则会引发 RuntimeError 或无限递归。
pin_memory锁页内存(Page-Locked Memory / Pinned Memory)
- 锁页内存指的是被操作系统锁定,物理地址始终固定,且不会被换出到磁盘的内存区域。pin_memory=True 让 DataLoader 分配“不可换出”的内存,使得 CPU → GPU 的数据传输跳过中间缓冲区,直接通过 DMA 进行,从而加速。
- 加速原理: CPU 的可分页内存物理地址不固定,且可能被换出,将 CPU 上的数据复制到 GPU 时,需要先将数据从可分页内存临时复制到一个内部的锁页缓冲区(操作系统自动完成)。再从该锁页缓冲区通过 DMA(直接内存访问)传输到 GPU。增加了一次内存拷贝。
- 加速数据传输:当你在训练循环中调用 .to(‘cuda’) 或 .cuda() 时,如果数据一开始就在锁页内存中,则可以直接通过 DMA 传输到 GPU,省去了中间的临时拷贝,因此传输速度更快(通常能快 2~3 倍,具体取决于数据量)
- 异步传输:锁页内存支持异步拷贝(non_blocking=True),可以与 CPU 上的数据预处理并行执行,进一步提高吞吐量。
| 优点 | 缺点 / 限制 |
|---|---|
| 显著加速 CPU → GPU 传输 | 分配/释放锁页内存比普通内存慢,且可能失败(内存碎片化) |
| 配合 non_blocking=True 可实现传输与计算重叠 | 占用更多物理内存(因为不能被换出), 大量使用pin_memory 可能降低系统整体性能 |
| 对大批量、频繁传输的场景收益明显 | 仅对张量有效;如果样本是自定义对象、PIL Image 或字符串,pin_memory 无法作用(但不会报错),实际使用中样本经过default_collate 或自定义 collate_fn 返回的张量后再放入锁页内存。如果自定义 collate_fn 返回的不是张量,则 pin_memory 无效且不报错。 |
- 使用时机:
- 当你的 GPU 训练数据流是瓶颈(即 GPU 经常等待数据),且内存足够。
- 内存紧张(比如只有 8GB),或数据集很小、传输开销微不足道。
- 在多进程数据加载(num_workers>0)时,pin_memory=True 效果更好,因为每个 worker 会独立分配锁页内存。
- 如果使用了 pin_memory=True,在将数据移到 GPU 时加上 non_blocking=True,并在下一次迭代前用 torch.cuda.synchronize() 或依赖依赖关系隐式同步(更常见的是直接 loss.backward() 会自动等待传输完成)。
3.Subset
torch.utils.data.Subset Subset 是 PyTorch 提供的一个数据集包装器(wrapper),用于从原始数据集中选取一部分样本,形成一个子集视图(不复制数据,只存储原数据集的引用和一个索引列表,额外存储一个整数列表(indices),内存开销可以忽略不计)
- 实现原理
- dataset:原始数据集(任何继承自 torch.utils.data.Dataset 的类)
- indices:整数列表或序列,指定要保留的样本索引
- 返回值Subset对象,而非原数据集对象,不具备原数据集的方法,但可以作为参数输入原数据集其他对象的方法中处理数据
python
class Subset(Dataset):
def __init__(self, dataset, indices):
self.dataset = dataset
self.indices = indices
def __getitem__(self, idx):
# 关键:先查 indices 表,再访问原数据集
return self.dataset[self.indices[idx]]
def __len__(self):
return len(self.indices)
- 应用
| 场景 | 示例 |
|---|---|
| 划分训练/验证/测试集 | train_subset = Subset(dataset, train_indices) |
| 快速创建小规模实验 | tiny_set = Subset(dataset, range(100)) |
| 交叉验证 | 每次选择不同的 indices 作为验证集 |
| 类别不平衡采样 | 提取所有正样本索引 + 部分负样本索引,组成平衡子集 |
| 调试/可视化 | 只加载前几十个样本,加快迭代 |
4.random_split
torch.utils.data.random_split是PyTorch 提供的一个数据集划分工具,用于将数据集随机切分成多个非重叠的子集(例如训练集、验证集、测试集)。它内部使用 Subset 来实现,因此同样不复制数据,只是返回基于原始数据集的视图。
- torch.utils.data.random_split(dataset, lengths, generator=torch.default_generator)
- dataset (Dataset):要划分的原始数据集。
- lengths (List[int]):每个子集的长度(样本数量)。长度之和必须等于 len(dataset)。
- generator (Generator, 可选):随机数生成器,用于控制随机打乱的过程(可复现性)
- 返回的就是一个 Subset 对象列表
- 实现原理
def random_split(dataset, lengths, generator=default_generator):
# 1. 随机打乱所有索引
indices = torch.randperm(sum(lengths), generator=generator).tolist()
# 2. 按 lengths 切分索引
splits = []
start = 0
for length in lengths:
splits.append(Subset(dataset, indices[start:start+length]))
start += length
return splits
5.get_dataloader()
为了方便使用DataLoader,引入自定义工厂函数get_dataloader(),它封装了创建 PyTorch DataLoader() 的常见流程,让训练代码更简洁、可配置。
- get_dataloader(data_dir, batch_size=1, shuffle=False, transform=None, subset=None)
| 参数 | 类型 | 默认值 | 含义 |
|---|---|---|---|
| data_dir | str | 必填 数据文件夹路径,会传给 MyDataset 去扫描文件 | |
| batch_size | int | 1 | 每个批次包含的样本数量 |
| shuffle | bool | False | 每个 epoch 开始时是否打乱数据顺序 |
| transform | callable or None | None | 数据预处理函数(如 torchvision.transforms.ToTensor()),传给MyDataset |
| subset | int or None | None | 如果指定了正整数,只使用数据集的前 subset 个样本(常用于快速测试或小规模实验) |
| 1.2 内部流程 |
def get_dataloader(data_dir, batch_size=1, shuffle=False, transform=None, subset=None):
"""
创建数据加载器
Args:
data_dir: 数据目录
batch_size: 批大小
shuffle: 是否打乱数据
transform: 数据转换函数
subset: 数据集子集大小,如果为None则使用全部数据
Returns:
DataLoader实例
"""
dataset = MyDataset(data_dir, transform)
# 如果指定了subset,创建数据集子集
if subset is not None and subset > 0:
from torch.utils.data import Subset
subset_indices = list(range(min(subset, len(dataset))))
dataset = Subset(dataset, subset_indices)
print(f"Using subset of dataset with {len(dataset)} samples")
return DataLoader(dataset, batch_size=batch_size, shuffle=shuffle, collate_fn=my_collate_fn)
七、数据集调用
import torch
from Dataset import get_dataloader
data_dir = "/data" #数据路径
if __name__ == "__main__"::
dataset = get_dataloader(data_dir, batch_size=1, shuffle=False, transform=None, subset=None,num_workers=8)
for batch,(data,lable) in emulate(dataset):
print(batch,data,lable)
八、其他细节
1.python的模块调用
from torch.utils.data import Dataset
from Dataset import get_dataloader
两个Dataset为什么编译器不会误认?
在 Python 里,from A import B 这个句式,真正拿到你当前代码“名字牌”的,只有 B。A 仅仅是一个“用来指路的地址”,指完路它就隐身了。
from torch.utils.data import Dataset
Python 的动作:
- Python 跑去官方的 torch 库里,翻箱倒柜找到了 Dataset 类(Class)。
- 发名字牌: Python 给它发了名字牌,贴在你的当前脚本里。
- 现在的字典状态: 你的代码认识了Dataset (指向 PyTorch 官方的 Dataset 类)
from DataLoader import get_dataloader
Python 的动作:
- 在当前目录下找一个叫 DataLoader.py 的文件(文件即模块 Module),把它打开,从里面抓出一个叫 get_dataloader 的函数(Function)。
- 发名字牌: Python 只给被抓出来的函数发了名字牌。它绝对不会给文件本身(DataLoader.py)发名字牌!
- 现在的字典状态: 你的代码又新认识了get_dataloader (指向你自己写的那个工厂函数)
在你的当前脚本运行完这两行 import 之后,它脑子里真正记住的变量名只有这三个:
Dataset和get_dataloader,没有冲突 - DataLoader.py,它只是在第二行代码执行的瞬间,充当了一次“门牌号”。Python 敲开这个门,拿走了 get_dataloader,然后就把这个门牌号抛在脑后了。它根本不会把“门牌号”当成一个变量存起来。
冲突报错情况:
1.作死写法 1(导入整个文件):
from torch.utils.data import Dataset
import Dataset
- import Dataset 会把你的那个文件本身作为一个模块(变量)拉进代码里。这时候,你的代码字典里就会有两个一模一样的名字牌 DataLoader:一个是 PyTorch 的类,一个是你的文件。Python 会陷入混乱(通常后导入的会覆盖先导入的),当你后面想调用官方的 DataLoader(dataset…) 时,程序直接报错,因为它以为你想调用那个文件。
作死写法 2(函数名撞车):
- 假设你在 DataLoader.py 里,而是自己手搓了一个类,名字也叫 Dataset。
from torch.utils.data import Dataset
from Dataset import Dataset # 完蛋了!
这样写,后一句导入的你自己写的 DataLoader,会直接把第一句 PyTorch 官方的 DataLoader 覆盖掉(覆盖效应)。后面你再用 PyTorch 的功能时就会疯狂报错。
2.训练效能优化
- 了解具体的硬件信息
nproc
sysctl -n hw.ncpu(Mac) # 查看CPU核心数量
nvidia-smi # 可以查看GPU的内存使用率(静态)
nvidia-smi -l 1
watch -n 1 nvidia-smi # 可以查看GPU的内存使用率(动态)
ctrl + C #退出watch
- 训练前期可以观察GPU信息以提升训练效率
如果GPU内存空间巨大,可以提升batch_size
如果可改为“如果 GPU 利用率出现波动甚至周期性归零,很有可能是CPU提供的数据跟不上,因此可以适当提升num_workers,num_workers<=CPU核心数,如果过大会导致内存占用过高与CPU调度负担,一般从核心数的一半开始测试。num_workser 在 Windows 上必须把训练循环放在 if name == ‘main’: 保护下,否则会引发 RuntimeError 或无限递归。
更多推荐



所有评论(0)