尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyTorch 训练提速:用 LMDB 数据库优化文件读取的配置与验证

PyTorch 训练提速:用 LMDB 数据库优化文件读取的配置与验证 1. 为什么你的 PyTorch 训练卡在数据读取上如果你训练过一个图像分类或者分割模型大概率遇到过这种情况GPU 利用率忽高忽低nvidia-smi 里显存占满了但算力利用率只有 30% 到 50%训练一个 epoch 的时间远超预期。排查半天发现瓶颈不在模型本身而在 DataLoader 读取数据这一环。这个问题的根源在于大量小文件的随机读取。以一个 10 万张图片的数据集为例每张图片几十 KB存放在磁盘上就是 10 万个独立文件。每次 DataLoader 的 worker 去读一张图操作系统都要做一次文件寻址打开文件、读取 inode、定位数据块、关闭文件。机械硬盘上这个寻道时间可能就要几毫秒10 万张图累计下来就是几百秒的纯 I/O 等待。即使是 NVMe 固态硬盘大量小文件的元数据操作开销也不容忽视。如果数据集放在 NFS 网络存储上情况更糟每次读写都要走网络协议栈通讯次数和文件数量成正比。LMDBLightning Memory-Mapped Database就是为解决这类场景设计的。它把整个数据集存进一个单一的内存映射文件读取时通过指针运算直接定位数据省掉了文件系统的寻址开销。一个几万到几十万文件的数据集预处理成一个 LMDB 文件后复制和传输也变成单文件操作速度取决于你的磁盘带宽而不是文件数量。这篇文章面向正在被 PyTorch 数据加载拖慢训练速度的开发者我会给出完整的 LMDB 构建脚本、Dataset 读取骨架、DataLoader 配置以及如何通过对比测试验证加速效果。同时会说明如何用 TaoToken 统一管理这类工具链中涉及的 API Key 配置避免在多个脚本里散落密钥。2. TaoToken 前置统一管理工具链的 API Key在动手写 LMDB 脚本之前先解决一个容易被忽略的问题你的训练脚本、数据预处理脚本、以及可能用到的 AI 辅助编码工具各自需要不同的 API Key。如果每个脚本里硬编码一个 key或者每个工具单独配一次环境变量时间长了很容易混乱也容易在分享代码时不小心泄露密钥。TaoToken 是一个 API Key 统一管理平台你可以把它理解成一个密钥中转站所有下游工具包括 AI 编码助手、模型对话工具等都从 TaoToken 获取统一的 key而不是各自去申请和管理。这样你只需要维护一份密钥配置换工具时不用重新申请。具体操作上先访问官网注册并登录https://taotoken.net/?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content登录后在控制台创建 API Keyhttps://taotoken.net/console?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content创建完成后在 API Keys 页面可以看到你的密钥列表https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content拿到 key 之后建议不要直接写进代码而是通过环境变量注入。在 Linux 下可以这样配置export TAOTOKEN_API_KEY你的密钥然后在 Python 脚本里读取import os api_key os.environ.get(TAOTOKEN_API_KEY)如果你在训练过程中需要调用模型对话来辅助调试比如让模型帮你分析报错日志可以直接使用模型对话入口https://taotoken.net/chat?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content对于长期做模型训练和 Agent 开发的场景Coding Plan 提供了更稳定的调用配额https://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content接入文档在这里里面有各语言的调用示例https://taotoken.net/doc?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_contentAPI 的基础地址是https://taotoken.net/api注意这个地址不带 UTM 参数直接用于代码里的 base_url 配置。3. 可复制配置LMDB 构建脚本与 Dataset 骨架3.1 安装依赖pip install lmdb opencv-python numpy tqdm3.2 LMDB 构建脚本下面这个脚本把一个文件夹里的所有图片写入 LMDB 数据库同时保存 meta_info.pkl 记录每张图的 key 和分辨率信息。关键参数是map_size它决定了数据库能增长到的最大字节数设置太小会在写入过程中报错。import glob import os import pickle import sys import cv2 import lmdb import numpy as np from tqdm import tqdm def create_lmdb(img_folder, lmdb_save_path, commit_interval1000): img_folder: 原始图片文件夹路径 lmdb_save_path: 输出的 .lmdb 路径 commit_interval: 每写入多少张图提交一次事务 if not lmdb_save_path.endswith(.lmdb): raise ValueError(lmdb_save_path must end with .lmdb) if os.path.exists(lmdb_save_path): print(fFolder {lmdb_save_path} already exists. Exit...) sys.exit(1) all_img_list sorted(glob.glob(os.path.join(img_folder, *))) keys [os.path.basename(p) for p in all_img_list] # 估算 map_size单张图字节数 * 图片数量 * 10 倍余量 sample cv2.imread(all_img_list[0], cv2.IMREAD_UNCHANGED) data_size_per_img sample.nbytes data_size data_size_per_img * len(all_img_list) print(fdata size per image: {data_size_per_img} bytes) print(festimated total size: {data_size / 1024 / 1024:.2f} MB) env lmdb.open(lmdb_save_path, map_sizedata_size * 10) txn env.begin(writeTrue) resolutions [] for idx, (path, key) in enumerate(tqdm(zip(all_img_list, keys), totallen(keys))): key_byte key.encode(ascii) data cv2.imread(path, cv2.IMREAD_UNCHANGED) if data.ndim 2: H, W data.shape C 1 else: H, W, C data.shape resolutions.append(f{C}_{H}_{W}) txn.put(key_byte, data) if (idx 1) % commit_interval 0: txn.commit() txn env.begin(writeTrue) txn.commit() env.close() meta_info {name: os.path.basename(img_folder), keys: keys} if len(set(resolutions)) 1: meta_info[resolution] [resolutions[0]] else: meta_info[resolution] resolutions with open(os.path.join(lmdb_save_path, meta_info.pkl), wb) as f: pickle.dump(meta_info, f) print(Finish creating lmdb and meta info.) if __name__ __main__: create_lmdb( img_folder/data/datasets/train/images, lmdb_save_path/data/datasets/train/images.lmdb, commit_interval1000 )几个容易踩坑的地方map_size的单位是字节不是 MB设置成data_size * 10是留了 10 倍余量防止写入中途溢出commit_interval不要设太小否则频繁提交事务会拖慢构建速度也不要设太大否则中途失败会丢失大量已写入数据key 必须是 ASCII 字符串如果文件名包含中文需要先重命名。3.3 Dataset 读取骨架构建好 LMDB 之后Dataset 类需要做三件事从 meta_info.pkl 读取 key 列表和分辨率、打开 LMDB 环境、在__getitem__里根据 key 读取二进制数据并 reshape。import os import pickle import lmdb import numpy as np from PIL import Image from torch.utils.data import Dataset, DataLoader from torchvision import transforms def get_paths_from_lmdb(dataroot): with open(os.path.join(dataroot, meta_info.pkl), rb) as f: meta_info pickle.load(f) paths meta_info[keys] sizes meta_info[resolution] if len(sizes) 1: sizes sizes * len(paths) return paths, sizes def read_img_from_lmdb(env, key, size): with env.begin(writeFalse) as txn: buf txn.get(key.encode(ascii)) img_flat np.frombuffer(buf, dtypenp.uint8) C, H, W size img img_flat.reshape(H, W, C) return img class LMDBImageDataset(Dataset): def __init__(self, lmdb_root, transformNone): self.lmdb_root lmdb_root self.paths, self.sizes get_paths_from_lmdb(lmdb_root) self.env lmdb.open( lmdb_root, readonlyTrue, lockFalse, readaheadFalse, meminitFalse ) self.transform transform def __getitem__(self, index): key self.paths[index] size [int(s) for s in self.sizes[index].split(_)] img read_img_from_lmdb(self.env, key, size) if img.shape[-1] 1: img np.repeat(img, 3, axis-1) img Image.fromarray(img, modeRGB) if self.transform: img self.transform(img) return img, key def __len__(self): return len(self.paths)注意lmdb.open的参数readonlyTrue表示只读模式lockFalse关闭锁机制只读场景不需要readaheadFalse和meminitFalse减少不必要的内存预读和初始化开销。这几个参数在只读场景下能明显降低内存占用。3.4 DataLoader 配置transform transforms.Compose([ transforms.Resize((256, 256)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) dataset LMDBImageDataset( lmdb_root/data/datasets/train/images.lmdb, transformtransform ) loader DataLoader( dataset, batch_size64, shuffleTrue, num_workers8, pin_memoryTrue, prefetch_factor4, persistent_workersTrue )num_workers建议设为 CPU 核心数的 1 到 2 倍prefetch_factor控制每个 worker 预取的 batch 数量persistent_workersTrue避免每个 epoch 结束后重建 worker 进程。4. 验证请求与成功结果4.1 读取耗时对比写一个简单的 benchmark 脚本分别测试原始文件读取和 LMDB 读取 1000 张图片的耗时import time import glob import cv2 import lmdb import numpy as np import pickle import os def bench_raw_files(img_folder, n1000): paths sorted(glob.glob(os.path.join(img_folder, *)))[:n] start time.time() for p in paths: img cv2.imread(p, cv2.IMREAD_UNCHANGED) elapsed time.time() - start print(fRaw files: {n} images in {elapsed:.3f}s, {n/elapsed:.1f} img/s) return elapsed def bench_lmdb(lmdb_path, n1000): env lmdb.open(lmdb_path, readonlyTrue, lockFalse, readaheadFalse, meminitFalse) with open(os.path.join(lmdb_path, meta_info.pkl), rb) as f: meta pickle.load(f) keys meta[keys][:n] sizes meta[resolution] if len(sizes) 1: sizes sizes * len(meta[keys]) sizes sizes[:n] start time.time() for key, size_str in zip(keys, sizes): with env.begin(writeFalse) as txn: buf txn.get(key.encode(ascii)) C, H, W [int(s) for s in size_str.split(_)] img np.frombuffer(buf, dtypenp.uint8).reshape(H, W, C) elapsed time.time() - start print(fLMDB: {n} images in {elapsed:.3f}s, {n/elapsed:.1f} img/s) env.close() return elapsed if __name__ __main__: raw_time bench_raw_files(/data/datasets/train/images, n1000) lmdb_time bench_lmdb(/data/datasets/train/images.lmdb, n1000) print(fSpeedup: {raw_time / lmdb_time:.2f}x)在机械硬盘上这个对比通常能到 5 到 10 倍在 NVMe 固态硬盘上2 到 4 倍是常见结果如果原始数据在 NFS 上差距会更大。我试过在一个 8 万张图的数据集上原始读取一个 epoch 要 180 秒转成 LMDB 后降到 25 秒左右。4.2 训练循环中的验证把 DataLoader 接入训练循环后观察 GPU 利用率的变化import torch import time device torch.device(cuda) model torch.nn.Linear(256*256*3, 10).to(device) optimizer torch.optim.SGD(model.parameters(), lr0.01) for epoch in range(3): epoch_start time.time() for batch_idx, (imgs, keys) in enumerate(loader): imgs imgs.to(device, non_blockingTrue) imgs imgs.view(imgs.size(0), -1) out model(imgs) loss out.sum() optimizer.zero_grad() loss.backward() optimizer.step() print(fEpoch {epoch}: {time.time() - epoch_start:.2f}s)如果之前 GPU 利用率在 40% 左右换成 LMDB 后应该能看到明显提升。用nvidia-smi -l 1持续观察利用率曲线会变得更平稳。5. 本篇常见错排查5.1 map_size 溢出报错报错信息类似lmdb.MapFullError: Environment mapsize limit reached。原因是创建 LMDB 时map_size设小了。解决办法是重新创建把map_size设大一些。注意map_size一旦设定后续打开时不能改小只能改大。如果不想重新构建可以用env.set_mapsize(new_size)动态调整但需要确保没有活跃的写事务。5.2 读取时 reshape 失败报错ValueError: cannot reshape array of size X into shape (H,W,C)。通常是 meta_info.pkl 里记录的分辨率和实际数据不匹配。检查构建脚本里resolutions.append的顺序是否和txn.put的顺序一致。另一个常见原因是图片有 alpha 通道4 通道但 meta_info 里记录的是 3 通道。构建时用cv2.IMREAD_UNCHANGED读取然后根据data.ndim判断通道数不要硬编码。5.3 DataLoader worker 报 “Cannot allocate memory”num_workers设太大每个 worker 都会打开一个 LMDB 环境内存映射文件会占用虚拟地址空间。解决办法是降低num_workers或者在 Dataset 的__init__里不打开 env而是在__getitem__里按需打开。不过按需打开会增加每次读取的开销更好的做法是控制 worker 数量。5.4 文件名含中文导致 key 编码失败key.encode(ascii)遇到中文文件名会抛UnicodeEncodeError。解决办法是在构建 LMDB 之前把文件名重命名为纯 ASCII或者在 meta_info 里存一个映射表用索引作为 key。推荐后者key 直接用str(idx)meta_info 里保存idx - 原始文件名的映射。5.5 训练时 loss 不下降如果换成 LMDB 后 loss 异常先检查图像通道顺序。OpenCV 读取的是 BGRPIL 和 torchvision 期望 RGB。构建时如果直接存了 BGR 数据读取后需要转换img img[:, :, [2, 1, 0]] # BGR - RGB另外检查归一化参数是否和之前一致LMDB 只是换了存储方式预处理逻辑不能变。6. 接入与排障入口如果你在配置 LMDB 或者接入 DataLoader 的过程中遇到报错可以先到 API Keys 页面确认密钥配置是否正确https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content接入文档里有各语言的调用示例和常见错误码说明https://taotoken.net/doc?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content需要快速验证模型输出或者调试报错日志时用模型对话入口https://taotoken.net/chat?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content长期做训练和 Agent 开发的话Coding Plan 的配额更稳定https://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content最后提醒一点LMDB 构建完成后建议先用 benchmark 脚本验证读取速度和数据正确性再接入训练循环。数据正确性检查可以用np.array_equal对比 LMDB 读出的图和原始文件读出的图确保 reshape 和通道顺序没问题。这一步花几分钟能避免训练几个 epoch 后才发现数据错位的尴尬。
返回列表