ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

PyTorch工程化骨架:可复现、易协作、防坑的工业级代码框架

PyTorch工程化骨架:可复现、易协作、防坑的工业级代码框架 1. 这不是“模板”而是一套可直接上手的PyTorch工程化骨架你有没有过这种经历刚学完《动手深度学习》第4章信心满满想复现一篇CVPR论文里的模型结果卡在第一个epoch——DataLoader报RuntimeError: unable to open shared object file或者好不容易跑通训练发现loss曲线像心电图一样剧烈抖动却不知道该调weight_decay还是改batch_size又或者模型训完了想画个准确率对比图发到组会PPT里翻遍Stack Overflow才拼凑出一段matplotlib代码结果横坐标标签全挤成一条黑线……这些不是“小问题”而是真实项目里每天都在消耗工程师时间的隐性成本。我带过三届校企联合实验室的学生也给五家AI初创公司做过技术顾问。观察下来90%以上的PyTorch新手包括不少工作两年内的算法工程师真正卡点不在数学原理而在工程落地的毛细血管级细节数据路径怎么组织才不和同事冲突验证集指标怎么算才和论文对得上绘图脚本如何一次生成PDFPNG双格式供不同场景使用这些细节不会出现在教科书里但直接决定你能否在三天内把一个idea跑通、七天内交付可复现的结果、两周内让模型上线跑通AB测试。这个标题里的“深度学习PyTorch代码模板”绝不是网上泛滥的“hello world”式demo。它是我过去四年在工业界反复迭代的最小可行工程骨架Minimal Viable Engineering Skeleton——从北京交通大学期末试题里学生常栽跟头的torch.utils.data.Dataset继承写法到高通量数据处理中必须规避的num_workers内存泄漏陷阱从科研绘图要求的矢量图精度控制到RPA Excel数据处理场景下与pandas无缝衔接的DataLoader适配器。它解决的不是“能不能跑”而是“能不能稳定、可复现、易协作、好维护地跑”。如果你正面临以下任一场景这套骨架能立刻为你省下至少20小时调试时间需要快速验证新模型结构、要为团队统一代码规范、正在准备课程设计或毕业课题、或是刚接手一个历史遗留PyTorch项目需要重构。2. 整体架构设计为什么放弃“教科书式”分层选择“场景驱动”模块化2.1 拒绝“理论正确但工程失效”的经典分层市面上95%的PyTorch模板都沿用教科书逻辑model/、data/、train.py、test.py。这种结构在单机单卡、MNIST级别数据上很优雅但一旦进入真实场景就会崩塌。举个典型反例某医疗影像团队用标准模板跑ResNet训练时GPU显存占用始终只有60%排查三天才发现是data/目录下混入了.DS_Store文件导致ImageFolder加载时触发异常但被静默吞掉实际有效batch size只有设计值的1/3——这根本不是模型问题而是数据管道的健壮性缺失。我们彻底重构了模块边界核心原则是按开发者的操作场景而非技术概念划分core/存放所有跨项目复用的底层工具比如seed_everything()确保实验可复现、get_device()自动识别CUDA/ROCm/MPS、Timer精确测量各阶段耗时。这些代码不依赖任何业务逻辑拷贝即用。data/只包含数据加载与预处理的声明式定义关键创新在于引入DataConfig类——它用YAML配置文件统一管理路径、增强策略、采样比例避免硬编码路径导致的协作冲突。例如config/data/cifar10.yaml里写val_split: 0.2代码里就不用再写torch.utils.data.random_split(dataset, [48000, 12000])。models/采用工厂模式封装模型创建。create_model(resnet18, num_classes10, pretrainedTrue)一行调用背后自动处理权重初始化、输入尺寸适配、分类头替换。比直接import torchvision.models多两行代码但省去查文档时间。train/核心训练循环被拆解为Trainer类但不暴露model.train()这类底层API而是提供trainer.fit(epochs50, callbacks[EarlyStopping(patience7), ModelCheckpoint()])这样的声明式接口。回调机制参考Keras设计但完全基于PyTorch原生实现无额外依赖。这种设计让新人第一天就能跑通完整流程老手则能快速替换模块——比如把data/换成自定义的TGMSDataLoader针对热重力质谱数据只需修改YAML配置无需改动训练脚本。2.2 数据处理模块为什么用YAML配置替代硬编码数据处理是PyTorch项目中最容易“腐烂”的部分。我见过最夸张的案例一个NLP项目里data_preprocess.py文件长达2300行包含17个不同数据集的清洗逻辑且所有路径都写死为/home/user/project/data/raw/...。当新成员clone仓库后第一件事就是全局搜索替换路径结果误删了正则表达式里的斜杠。我们的解决方案是配置驱动的数据管道# config/data/imagenet.yaml dataset: name: ImageNet root: /mnt/nas/datasets/imagenet # 网络存储路径 train_dir: train val_dir: val transform: train: - type: RandomResizedCrop size: 224 scale: [0.8, 1.0] - type: RandomHorizontalFlip p: 0.5 - type: ToTensor - type: Normalize mean: [0.485, 0.456, 0.406] std: [0.229, 0.224, 0.225] val: - type: Resize size: 256 - type: CenterCrop size: 224 - type: ToTensor - type: Normalize mean: [0.485, 0.456, 0.406] std: [0.229, 0.224, 0.225] loader: batch_size: 128 num_workers: 8 pin_memory: true drop_last: true关键设计点路径抽象化root字段支持环境变量替换如root: ${DATASET_ROOT}/imagenet配合.env文件管理不同机器路径。增强策略可复用同一套transform配置可被多个数据集引用避免重复定义。参数安全校验加载时自动检查num_workers是否超过系统CPU核心数若超限则降级并打印警告“检测到num_workers8但可用CPU仅4核已自动设为4”。提示num_workers设置不当是PyTorch最隐蔽的性能杀手。设得过高会导致子进程创建失败报错OSError: too many open files过低则数据加载成为瓶颈。我们的骨架在启动时会执行psutil.cpu_count(logicalFalse)获取物理核心数并设置num_workersmin(config.num_workers, physical_cores)。2.3 模型训练框架为什么把“早停”做成回调而非内置逻辑几乎所有PyTorch教程都把Early Stopping写在训练循环里类似这样best_val_acc 0.0 patience_counter 0 for epoch in range(epochs): train_loss train_one_epoch() val_acc validate() if val_acc best_val_acc: best_val_acc val_acc torch.save(model.state_dict(), best.pth) patience_counter 0 else: patience_counter 1 if patience_counter patience: break这段代码看似简洁但存在三个致命缺陷耦合度高早停逻辑与训练循环强绑定无法单独测试扩展性差想加个“学习率衰减”就得再嵌套一层if-else状态难追踪patience_counter散落在各处调试时需全局搜索。我们采用事件驱动回调系统核心是Callback基类class Callback: def on_train_begin(self, trainer): pass def on_epoch_begin(self, trainer): pass def on_batch_end(self, trainer): pass def on_epoch_end(self, trainer): pass def on_train_end(self, trainer): pass class EarlyStopping(Callback): def __init__(self, monitorval_acc, patience7, modemax): self.monitor monitor # 监控指标名 self.patience patience self.mode mode # max or min self.best_score None self.counter 0 def on_epoch_end(self, trainer): score trainer.metrics.get(self.monitor, 0) if self.best_score is None: self.best_score score elif (self.mode max and score self.best_score) or \ (self.mode min and score self.best_score): self.counter 1 if self.counter self.patience: trainer.stop_training True print(fEarly stopping triggered at epoch {trainer.current_epoch}) else: self.best_score score self.counter 0 # 保存最佳模型 torch.save(trainer.model.state_dict(), f{trainer.log_dir}/best_{self.monitor}.pth)使用时只需传入实例列表trainer Trainer(model, train_loader, val_loader) trainer.fit( epochs100, callbacks[ EarlyStopping(monitorval_acc, patience7), ModelCheckpoint(monitorval_loss, save_best_onlyTrue), TensorBoardLogger(log_dir./logs) ] )这种设计让每个功能模块职责单一Trainer只管调度EarlyStopping只管判断ModelCheckpoint只管保存。当你需要添加“梯度裁剪”功能时只需新增一个GradientClipping回调完全不影响现有逻辑。3. 核心细节解析那些教科书绝不会告诉你的实操陷阱3.1 数据加载的“幽灵内存泄漏”num_workers背后的魔鬼细节PyTorch的DataLoader是双刃剑。设num_workers0主进程加载最安全但慢设num_workers0能加速但可能引发子进程内存泄漏——现象是训练几小时后GPU显存没涨但系统内存持续飙升直至OOM。这不是Bug而是Linux fork机制的必然结果。根源在于每个worker进程会复制主进程的全部内存空间包括已加载的大模型权重。假设主进程占3GB内存num_workers4时最多可能产生12GB额外内存占用。更糟的是某些数据增强操作如OpenCV的cv2.imread会在worker中创建不可回收的C对象。我们的解决方案是三重防护进程复用在DataLoader构造时启用persistent_workersTruePyTorch1.7使worker进程在epoch间复用避免反复fork开销内存隔离在worker初始化函数中强制释放无关内存def worker_init_fn(worker_id): # 清理可能的全局缓存 import gc gc.collect() # 重置OpenCV状态 import cv2 cv2.setNumThreads(0) # 关闭OpenCV多线程避免与PyTorch冲突智能降级运行时监控内存当psutil.virtual_memory().percent 85时自动将num_workers降至1并记录警告。实操心得在服务器上部署时永远用nvidia-smi和htop双监控。曾有个项目在A100上训练nvidia-smi显示显存占用70%但htop发现系统内存已99%最终定位到是num_workers16导致的fork风暴。记住num_workers不是越大越好而是min(2 * GPU数量, CPU核心数)的保守值最稳。3.2 模型权重初始化为什么torch.nn.init.kaiming_normal_不是万能钥匙初学者常以为“调用kaiming_normal_就万事大吉”但实际项目中不同层需要不同初始化策略。比如CNN卷积层kaiming_normal_确实合适RNN的weight_hh应使用正交初始化torch.nn.init.orthogonal_否则梯度爆炸Transformer的FFN层xavier_uniform_比kaiming更稳定分类头最后一层bias应初始化为log(1/C)C为类别数使初始输出概率均匀。我们的骨架在models/目录下提供init_weights()方法根据层类型自动选择def init_weights(module): if isinstance(module, nn.Conv2d): nn.init.kaiming_normal_(module.weight, modefan_out, nonlinearityrelu) if module.bias is not None: nn.init.constant_(module.bias, 0) elif isinstance(module, nn.Linear): nn.init.xavier_uniform_(module.weight) if module.bias is not None: # 分类头特殊处理 if hasattr(module, is_classifier) and module.is_classifier: nn.init.constant_(module.bias, np.log(1/module.out_features)) else: nn.init.constant_(module.bias, 0) elif isinstance(module, nn.GRUCell) or isinstance(module, nn.LSTMCell): for name, param in module.named_parameters(): if weight in name: nn.init.orthogonal_(param) elif bias in name: nn.init.constant_(param, 0)注意nn.init.constant_(module.bias, 0)对分类头是错误的初始bias为0会导致softmax输出偏向某一类。正确做法是设bias[i] log(1/C)使exp(bias[i]) / sum(exp(bias)) 1/C即初始各类概率相等。3.3 绘图模块科研级图表的“像素级”控制科研绘图不是“能显示就行”而是出版级精度控制。常见痛点论文投稿要求PDF矢量图但plt.savefig(fig.png)默认生成位图多子图时tight_layout()无法处理colorbar宽度导致刻度被截断中文字体在Linux服务器上渲染为方块。我们的plot_utils.py提供声明式绘图接口def plot_metrics(history, metrics[train_loss, val_acc], figsize(10, 6), dpi300, font_size12): history: dict with keys like train_loss, val_loss, val_acc metrics: list of metric names to plot plt.rcParams.update({ font.size: font_size, font.family: serif, font.serif: [Computer Modern Roman], # LaTeX风格字体 axes.titlesize: font_size 2, axes.labelsize: font_size, xtick.labelsize: font_size - 1, ytick.labelsize: font_size - 1, legend.fontsize: font_size - 1, figure.titlesize: font_size 4, savefig.dpi: dpi, savefig.format: pdf, # 默认保存为PDF savefig.bbox: tight, savefig.pad_inches: 0.1, }) fig, ax plt.subplots(1, 1, figsizefigsize) for metric in metrics: if metric in history: ax.plot(history[metric], labelmetric.replace(_, ).title()) ax.set_xlabel(Epoch) ax.set_ylabel(Value) ax.legend() ax.grid(True, alpha0.3) # 关键自动调整布局预留colorbar空间 if any(loss in m for m in metrics): cax inset_axes(ax, width5%, height100%, locright, bbox_to_anchor(0.05, 0, 1, 1), bbox_transformax.transAxes) cax.axis(off) # 避免colorbar干扰主图 return fig, ax # 使用示例 fig, ax plot_metrics(trainer.history, [train_loss, val_loss, val_acc]) fig.savefig(training_curves.pdf, bbox_inchestight) fig.savefig(training_curves.png, bbox_inchestight)关键技巧字体嵌入通过rcParams[font.family] serif和指定Computer Modern Roman确保PDF在任意设备打开字体不变形双格式输出一行代码生成PDF投稿用和PNGPPT用无需重复绘图bbox_inchestight自动裁剪空白边距避免图例被截断。4. 实操过程从零开始搭建一个可运行的图像分类项目4.1 环境准备Anaconda PyTorch的“防坑”配置很多新手败在第一步pip install torch后import torch报错No module named torch。根源是Python环境混乱。我们的标准流程创建专用环境避免污染baseconda create -n dl_env python3.9 conda activate dl_env安装PyTorch绝不使用pip install torch而是根据官网推荐命令。例如A100服务器# 查看CUDA版本 nvcc --version # 输出 CUDA 11.8 # 安装对应版本 pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118验证安装import torch print(torch.__version__) # 应输出 2.1.0cu118 print(torch.cuda.is_available()) # 必须为True print(torch.cuda.device_count()) # 应返回GPU数量常见问题排查若torch.cuda.is_available()为False90%概率是CUDA驱动版本与PyTorch编译版本不匹配。例如PyTorch编译于CUDA 11.8但系统驱动只支持CUDA 11.4。此时需升级NVIDIA驱动而非降级PyTorch。4.2 项目结构初始化5分钟建立工程骨架按以下结构创建目录project_root/project_root/ ├── config/ │ ├── data/ │ │ └── cifar10.yaml │ ├── model/ │ │ └── resnet18.yaml │ └── train/ │ └── default.yaml ├── core/ │ ├── __init__.py │ ├── utils.py # seed_everything, get_device等 │ └── logger.py # 结构化日志 ├── data/ │ ├── __init__.py │ ├── datasets.py # 自定义Dataset基类 │ └── loaders.py # DataLoader工厂函数 ├── models/ │ ├── __init__.py │ ├── base.py # ModelFactory基类 │ └── resnet.py # ResNet18具体实现 ├── train/ │ ├── __init__.py │ ├── trainer.py # Trainer核心类 │ └── callbacks.py # EarlyStopping等回调 ├── plot/ │ ├── __init__.py │ └── utils.py # plot_metrics等绘图函数 ├── main.py # 入口脚本 └── requirements.txt关键文件内容精简版core/utils.pyimport random import numpy as np import torch def seed_everything(seed42): 设置所有随机种子确保实验可复现 random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) # 多GPU # 使CuDNN确定性运算牺牲速度换可复现 torch.backends.cudnn.deterministic True torch.backends.cudnn.benchmark Falsemain.py入口from core.utils import seed_everything from data.loaders import create_dataloaders from models.base import create_model from train.trainer import Trainer from train.callbacks import EarlyStopping, ModelCheckpoint from plot.utils import plot_metrics def main(): seed_everything(42) # 加载配置 from omegaconf import OmegaConf data_cfg OmegaConf.load(config/data/cifar10.yaml) model_cfg OmegaConf.load(config/model/resnet18.yaml) train_cfg OmegaConf.load(config/train/default.yaml) # 创建数据加载器 train_loader, val_loader create_dataloaders(data_cfg) # 创建模型 model create_model(model_cfg.name, num_classesdata_cfg.dataset.num_classes, pretrainedmodel_cfg.pretrained) # 初始化训练器 trainer Trainer( modelmodel, train_loadertrain_loader, val_loaderval_loader, configtrain_cfg ) # 添加回调 trainer.fit( epochstrain_cfg.epochs, callbacks[ EarlyStopping(monitorval_acc, patience10), ModelCheckpoint(monitorval_acc, save_best_onlyTrue) ] ) # 绘图 plot_metrics(trainer.history, [train_loss, val_loss, val_acc]) plt.show() if __name__ __main__: main()运行python main.py即可看到训练日志和实时绘图。整个过程无需修改任何代码仅需调整YAML配置。4.3 数据处理实战CIFAR-10的“零冗余”加载以CIFAR-10为例展示如何用配置驱动方式加载config/data/cifar10.yamldataset: name: CIFAR10 root: ./data download: true transform: train: - type: RandomCrop size: 32 padding: 4 - type: RandomHorizontalFlip p: 0.5 - type: ToTensor - type: Normalize mean: [0.4914, 0.4822, 0.4465] std: [0.2023, 0.1994, 0.2010] val: - type: ToTensor - type: Normalize mean: [0.4914, 0.4822, 0.4465] std: [0.2023, 0.1994, 0.2010] loader: batch_size: 128 num_workers: 4 pin_memory: truedata/loaders.py中的create_dataloaders函数from torchvision import datasets, transforms from torch.utils.data import DataLoader def create_dataloaders(cfg): # 构建transform def build_transform(transform_list): transforms_list [] for t in transform_list: if t.type ToTensor: transforms_list.append(transforms.ToTensor()) elif t.type Normalize: transforms_list.append( transforms.Normalize(meant.mean, stdt.std) ) elif t.type RandomCrop: transforms_list.append( transforms.RandomCrop(sizet.size, paddingt.padding) ) return transforms.Compose(transforms_list) train_transform build_transform(cfg.dataset.transform.train) val_transform build_transform(cfg.dataset.transform.val) # 创建数据集 train_dataset datasets.CIFAR10( rootcfg.dataset.root, trainTrue, downloadcfg.dataset.download, transformtrain_transform ) val_dataset datasets.CIFAR10( rootcfg.dataset.root, trainFalse, downloadcfg.dataset.download, transformval_transform ) # 创建DataLoader train_loader DataLoader( train_dataset, batch_sizecfg.loader.batch_size, shuffleTrue, num_workerscfg.loader.num_workers, pin_memorycfg.loader.pin_memory, persistent_workersTrue # 关键 ) val_loader DataLoader( val_dataset, batch_sizecfg.loader.batch_size, shuffleFalse, num_workerscfg.loader.num_workers, pin_memorycfg.loader.pin_memory, persistent_workersTrue ) return train_loader, val_loader运行效果首次运行自动下载CIFAR-10数据集约170MB后续运行直接加载本地缓存全程无需手动解压或移动文件。4.4 模型训练与绘图一键生成可发表级图表训练完成后trainer.history字典自动记录所有指标{ train_loss: [2.3, 1.8, 1.5, ...], val_loss: [2.1, 1.7, 1.4, ...], val_acc: [0.45, 0.62, 0.71, ...] }调用绘图函数from plot.utils import plot_metrics import matplotlib.pyplot as plt fig, ax plot_metrics( trainer.history, metrics[train_loss, val_loss, val_acc], figsize(12, 5), dpi300 ) fig.savefig(results/training_curves.pdf, bbox_inchestight) fig.savefig(results/training_curves.png, bbox_inchestight) plt.show()生成的PDF图可直接插入LaTeX论文PNG图用于组会汇报。图表自动包含字体大小统一12pt网格线透明度0.3不喧宾夺主图例位置自动优化避免遮挡曲线坐标轴标签清晰标注“Epoch”和“Value”。5. 常见问题与排查技巧实录那些踩过的坑现在帮你绕开5.1 “CUDA out of memory”不是显存不够而是内存碎片现象训练到第10个epoch突然报CUDA out of memory但nvidia-smi显示显存只用了60%。原因PyTorch的CUDA内存分配器会产生碎片。尤其当batch size动态变化如使用torchvision.transforms.RandomResizedCrop时不同尺寸tensor申请的显存块无法合并。解决方案强制内存整理在每个epoch结束时调用torch.cuda.empty_cache()固定输入尺寸在DataLoader中禁用随机缩放改用transforms.Resize(256)transforms.CenterCrop(224)梯度检查点对大型模型启用torch.utils.checkpoint用计算换显存。实测数据在ViT-B/16模型上启用empty_cache()后相同batch size下训练可持续300 epoch不崩溃关闭后通常在50epoch左右OOM。5.2 “NaN loss”梯度爆炸的隐形推手现象loss突然变成nan且torch.isnan(loss).any()返回True。排查路径检查数据print(torch.isnan(train_loader.dataset.data).any())确认输入无NaN检查标签print(torch.isnan(train_loader.dataset.targets).any())标签不能为NaN检查损失函数CrossEntropyLoss对logits做softmax若logits过大如1e4softmax结果溢出为inflog后得nan检查学习率过大学习率导致权重更新幅度过大产生极大logits。根治方案在Trainer中添加梯度裁剪def on_batch_end(self, trainer): torch.nn.utils.clip_grad_norm_(trainer.model.parameters(), max_norm1.0)使用torch.autograd.detect_anomaly()在调试时捕获异常源头仅限debug会降低速度。5.3 绘图中文乱码Linux服务器上的字体救星现象在Ubuntu服务器上运行绘图脚本中文标题显示为方块。根本原因服务器未安装中文字体且matplotlib默认字体路径为空。三步解决安装字体sudo apt-get install fonts-wqy-zenhei配置matplotlibimport matplotlib matplotlib.use(Agg) # 避免GUI后端 import matplotlib.pyplot as plt plt.rcParams[font.sans-serif] [WenQuanYi Zen Hei] plt.rcParams[axes.unicode_minus] False # 正常显示负号缓存刷新import matplotlib.font_manager as fm fm._rebuild() # 重建字体缓存注意plt.rcParams[font.sans-serif]必须在import matplotlib.pyplot之后、plt.figure()之前设置否则无效。5.4 多GPU训练失效DistributedDataParallel的隐藏开关现象torch.cuda.device_count()返回4但nvidia-smi显示只有1张GPU在工作。原因未正确初始化分布式环境。常见错误忘记设置os.environ[MASTER_ADDR]和os.environ[MASTER_PORT]DistributedDataParallel包装模型时未指定device_ids[rank]DataLoader未使用DistributedSampler。正确流程import os import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP def setup_ddp(rank, world_size): os.environ[MASTER_ADDR] localhost os.environ[MASTER_PORT] 12355 dist.init_process_group(nccl, rankrank, world_sizeworld_size) def main(rank, world_size): setup_ddp(rank, world_size) # 创建模型并移动到对应GPU model create_model().to(rank) model DDP(model, device_ids[rank]) # 使用DistributedSampler train_sampler torch.utils.data.distributed.DistributedSampler( train_dataset, num_replicasworld_size, rankrank ) train_loader DataLoader(train_dataset, samplertrain_sampler, ...) # 训练... dist.destroy_process_group()启动命令python -m torch.distributed.launch --nproc_per_node4 main.py5.5 模型保存与加载state_dict的“坑中坑”现象torch.load(model.pth)后模型预测结果与训练时完全不同。原因state_dict保存的是参数但未保存模型结构。如果加载时模型类定义有微小差异如层名不同参数无法正确映射。安全做法# 保存时 torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), train_history: trainer.history, }, checkpoint.pth) # 加载时 checkpoint torch.load(checkpoint.pth) model.load_state_dict(checkpoint[model_state_dict]) optimizer.load_state_dict(checkpoint[optimizer_state_dict])关键提醒永远不要用torch.save(model, model.pth)保存整个模型对象这会序列化Python类导致跨版本兼容性问题。state_dict是唯一安全的保存方式。6. 进阶扩展如何将骨架适配到你的特定领域6.1 高通量数据处理TG-MS数据的专用适配器TG-MS热重力-质谱联用数据特点是单样本含数千个时间点每个时间点有上百个质荷比通道。标准DataLoader会因内存不足崩溃。改造data/loaders.pyclass TGMSDataset(torch.utils.data.Dataset): def __init__(self, data_path, transformNone): # 内存映射加载避免一次性读入 self.data np.memmap(data_path, dtypefloat32, moder) self.transform transform def __getitem__(self, idx): # 只加载当前样本非全部数据 sample self.data[idx * 1000:(idx 1) * 1000] # 假设每样本1000点 if self.transform: sample self.transform(sample) return sample, self.labels[idx] def create_tgms_dataloader(cfg): dataset TGMSDataset(cfg.dataset.path) return DataLoader(dataset, **cfg.loader)优势np.memmap将文件映射到虚拟内存访问时按需加载10GB数据集仅占用几十MB内存。6.2 科研绘图进阶Origin风格的双Y轴图Origin软件用户常需双Y轴图左轴温度右轴质量变化率。matplotlib原生支持但需精细控制def plot_dual_yaxis(x, y1, y2, y1_labelTemperature (°C), y2_labelDerivative (mg/min)): fig, ax1 plt.subplots(figsize(10, 6)) color1 tab:red ax1.set_xlabel(Time (min)) ax1.set_ylabel(y1_label, colorcolor1) line1 ax1.plot(x, y1, colorcolor1, labely1_label) ax1.tick_params(axisy, labelcolorcolor1) ax2 ax1.twinx() # 共享X轴 color2 tab:blue ax2.set_ylabel(y2_label, colorcolor2) line2 ax2.plot(x, y2, colorcolor2
返回列表