一、为什么需要使用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.float32torch.float
torch.float64torch.double
torch.float16
torch.bfloat16
单精度浮点数(32位)默认类型
双精度浮点数(64位)
半精度浮点数(16位,常用于混合精度训练)
Brain 浮点格式(16位,指数与 float32 相同,精度较低)
整数类型 torch.int32torch.int
torch.int64torch.long
torch.int16torch.short
torch.int8
torch.uint8
有符号整数(32位)
有符号长整数(64位),常用于索引
有符号短整数(16位)
有符号8位整数
无符号8位整数(常用于图像像素 0-255)
布尔类型 torch.bool 布尔值 True / False
其他专用类型 torch.complex64
torch.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 指定填充值,但不支持左填充或复杂填充模式(需自己写)。
      • 额外功能:自动按输入顺序填充,不改变序列相对顺序。
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, 可选): 输出张量(通常不手动使用)。
#获取最长序列长度
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_workserWindows 上必须把训练循环放在 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_workserWindows 上必须把训练循环放在 if name == ‘main’: 保护下,否则会引发 RuntimeError 或无限递归。
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐