PyTorch深度学习实战避坑指南:从环境配置到模型部署的常见问题与解决方案
2026/9/3 21:11:08 网站建设 项目流程

1. 项目概述:从“炼丹”到“避坑”的必经之路

搞深度学习,尤其是用PyTorch,这几年下来感觉就像在“炼丹”。模型是丹方,数据是药材,GPU是炉火,而我们这些工程师,就是守着炉子看火候的道童。一开始总是信心满满,照着论文把模型搭起来,数据喂进去,然后满怀期待地等着Loss曲线优雅下降,验证集指标一路飙升。但现实往往是,你盯着那个波动剧烈、甚至纹丝不动的Loss值,或者遇到一个莫名其妙的RuntimeError,一耗就是大半天。这些就是“坑”,它们不写在任何官方教程的显眼处,却遍布在实际项目的每一个角落。

今天聊的“深度学习中的踩过的一些坑”,不是什么高深的理论突破,而是那些在真实项目推进中,一旦遇到就能让你进度卡壳、头皮发麻的实战问题。它们涉及环境配置、数据加载、分布式训练、内存管理等多个方面,很多错误信息看起来云里雾里,但背后的原因往往很简单。这篇文章的目的,就是把我个人和团队在多个CV、NLP项目里,用PyTorch框架时反复遇到的那些典型“坑”梳理出来,不仅告诉你“坑”在哪,更重点分析“为什么”会这样,以及“怎么”干净利落地填上它。无论你是刚入门的新手,还是有一定经验的研究员,希望这些从真实教训里总结出的经验,能帮你少走些弯路,把更多精力花在算法创新和模型调优上。

2. 环境配置与依赖管理的“隐形雷区”

环境问题堪称深度学习项目的“第零大坑”。模型代码明明一样,在A机器上跑得飞快,到B机器上就各种报错,很多时候根源就在环境。

2.1 CUDA、cuDNN与PyTorch版本的“三角关系”

这是最经典,也最容易出问题的地方。PyTorch的GPU版本依赖于特定版本的CUDA驱动和cuDNN库,这三者必须严格匹配。

常见坑点:直接从PyTorch官网用pip install torch torchvision命令安装,默认可能会安装最新的版本,而你的服务器CUDA驱动版本可能较旧。例如,你的驱动只支持CUDA 11.7,却安装了需要CUDA 12.1的PyTorch,运行时就会报错:CUDA error: no kernel image is available for execution on the device

避坑实操

  1. 首先检查驱动版本:在命令行执行nvidia-smi,右上角会显示CUDA Version,这个是你的驱动最高能支持的CUDA运行时版本,不代表已安装。
  2. 确定已安装的CUDA Toolkit版本:执行nvcc --version。如果未安装,则需要先安装。务必确保这个版本 ≤nvidia-smi显示的版本。
  3. 去PyTorch官网获取精确安装命令:访问 pytorch.org ,使用其提供的安装命令生成器。选择你的PyTorch版本、操作系统、包管理工具(Conda或Pip)、CUDA版本。例如,对于CUDA 11.7,命令可能类似:
    # Conda conda install pytorch==1.13.1 torchvision==0.14.1 torchaudio==0.13.1 cudatoolkit=11.7 -c pytorch -c conda-forge # Pip pip install torch==1.13.1+cu117 torchvision==0.14.1+cu117 torchaudio==0.13.1 --extra-index-url https://download.pytorch.org/whl/cu117
    关键点cu117这样的后缀就是为特定CUDA版本编译的包。

注意:使用Conda安装时,cudatoolkit包是Conda环境内独立的CUDA运行时,可能与系统全局安装的CUDA不同。Conda环境会优先使用自带的。这有时能解决系统环境混乱的问题,但也可能因版本不一致引发新问题。

2.2 Conda虚拟环境:隔离与复现的生命线

强烈建议为每个项目创建独立的Conda虚拟环境。这能避免包版本冲突,也是项目可复现性的基础。

常见坑点

  • 环境未导出或导出文件不完整:只导出了conda env export > environment.yaml,但这个YAML文件包含了通过pip安装的包吗?在旧版本Conda中,可能不包含。
  • 跨平台/跨硬件复现失败:在Linux上导出的环境,在Windows上安装时,某些包可能没有对应平台的版本。

避坑实操

  1. 创建环境时指定Python版本conda create -n my_project python=3.9
  2. 使用--from-history标志导出最小环境conda env export --from-history > environment.yaml。这只会导出你显式安装的包,依赖项由在新机器上解决,兼容性更好。
  3. 分离Conda和Pip包:更稳健的做法是,在环境YAML文件中手动管理,或使用两个文件:conda_env.yaml(仅Conda包)和requirements.txt(仅Pip包)。安装时先conda env create -f conda_env.yaml,再进入环境pip install -r requirements.txt
  4. 对于绝对复现:使用conda env export --no-builds > environment.yml导出完整环境(包含构建号),但这可移植性较差,仅用于相同系统环境的克隆。

2.3 第三方库的版本“刺头”

一些常见的库,如opencv-pythonpillownumpy,如果版本不匹配,会导致从图像解码到张量计算各种稀奇古怪的错误。

案例:一个图像预处理管道中,使用了PIL.Imagetorchvision.transforms。升级pillow到某个版本后,ToTensor()转换突然报错,原因是新版本PIL返回的图像模式与旧版本有细微差别。

避坑实操

  • 在项目初期就固定所有核心依赖的版本。在requirements.txtsetup.py中使用==指定版本号。
  • 定期更新依赖,但要在隔离环境中测试,确认无误后再更新项目主环境。

3. 数据加载管道(DataLoader)的“性能陷阱”与“逻辑暗坑”

DataLoader是PyTorch训练流程的“炊事班”,负责把原始数据“做熟”(预处理)并源源不断地送到模型“嘴边”。这里一旦出问题,直接影响训练效率和正确性。

3.1num_workers设置不当:不是越多越好

DataLoadernum_workers参数指定了用于数据加载的子进程数,旨在通过并行预加载数据来避免GPU等待CPU(数据预处理),从而提升GPU利用率。

常见坑点

  1. 盲目设大:在Windows或某些Linux环境下,num_workers设置过大(如超过CPU核心数),会导致进程创建、切换开销巨大,反而使数据加载变慢,甚至因为内存消耗过多而崩溃。
  2. 设为0:在主进程加载数据,GPU在大部分时间处于空闲状态,严重拖慢训练速度。
  3. 内存泄漏:如果自定义的Dataset在每个__getitem__中打开了文件句柄、数据库连接等资源但未正确关闭,在多进程模式下,这些资源会因进程复制而泄漏,最终耗尽系统资源。

避坑实操

  • 经验值:通常设置为CPU逻辑核心数的2到4倍,但需要实测。可以从4开始,逐步增加,观察GPU利用率(使用nvidia-smi查看Volatile GPU-Util)和训练速度,找到一个稳定且高效的平衡点。
  • 监控工具:在Linux下,可以用htop命令观察worker进程的CPU和内存占用是否正常。
  • 资源管理:在自定义Dataset__init__中打开共享资源(如打开一个HDF5文件),在__getitem__中只读取。或者使用torch.multiprocessing的共享内存机制。确保没有在__getitem__内部进行频繁的IO连接操作。

3.2Sampler与分布式训练(DDP)的冲突

当你使用DistributedDataParallel进行多卡训练时,数据分配逻辑变得复杂。每张卡(每个进程)应该看到数据的一个子集,且彼此不重叠,这是通过DistributedSamplerBatchSampler实现的。

常见坑点

  • 忘记使用DistributedSampler:在DDP模式下,如果仍然使用默认的RandomSampler,那么每个进程都会加载全部数据,导致数据重复,模型无法正确学习,且浪费资源。
  • Epoch长度计算错误DistributedSampler会在每个epoch对数据索引进行分区和打乱。但需要注意的是,默认情况下,每个epoch的迭代次数(steps)是ceil(len(dataset) / world_size) / batch_size_per_gpu。这意味着,如果数据集总长度不能被world_size * batch_size整除,最后一个batch可能包含重复数据以“凑齐”一个batch。这可能会轻微影响训练,尤其是当验证集指标在每个epoch末计算时。
  • 验证集的数据重复:在验证时,通常不需要分布式采样,因为评估是在每个进程上独立进行然后聚合的。如果验证DataLoader也错误地使用了DistributedSampler,会导致验证指标计算错误。

避坑实操

import torch.distributed as dist from torch.utils.data import DataLoader, DistributedSampler # 初始化进程组 dist.init_process_group(backend='nccl') local_rank = dist.get_rank() world_size = dist.get_world_size() train_dataset = MyDataset(...) train_sampler = DistributedSampler(train_dataset, shuffle=True) # 注意:在DataLoader中设置 shuffle=False,因为Sampler已经处理了打乱 train_loader = DataLoader(train_dataset, batch_size=per_gpu_batch, sampler=train_sampler, num_workers=4, pin_memory=True) # 验证集不需要DistributedSampler,或者使用一个不shuffle的DistributedSampler以保证每个进程看到相同的验证集顺序(如果做同步评估) val_sampler = DistributedSampler(val_dataset, shuffle=False) if ddp_mode else None val_loader = DataLoader(val_dataset, batch_size=per_gpu_batch, sampler=val_sampler, num_workers=2, pin_memory=True) # 在每个epoch开始前,调用sampler的set_epoch,确保不同epoch的数据打乱不同 for epoch in range(num_epochs): train_sampler.set_epoch(epoch) # 这行很重要! for batch in train_loader: # training...

3.3pin_memory的误解与正确使用

pin_memory=True可以将数据加载到主机(CPU)的页锁定内存中。这种内存不会被操作系统交换到磁盘,并且可以通过DMA(直接内存访问)更快地传输到GPU,从而加速DataLoader的数据传输。

常见坑点

  • 认为用了就一定能加速pin_memory的加速效果在数据量小、传输频繁时最明显。如果数据预处理本身是瓶颈(CPU耗时远大于数据传输耗时),或者batch size极大导致传输时间占比很小,那么开启pin_memory的收益可能不明显,反而会占用更多的主机内存。
  • 内存不足:页锁定内存是稀缺资源。如果数据集很大,且num_workers较多,每个worker都pin住一个batch的数据,可能导致主机内存耗尽,触发OOM(Out-Of-Memory)。

避坑实操

  • 默认开启:对于大多数CV、NLP任务,只要主机内存充足(例如,空闲内存远大于num_workers * batch_size * sample_size),建议默认设置pin_memory=True
  • 监控内存:在训练时,使用htopfree -h命令监控主机内存使用情况。如果发现内存使用持续增长直至占满,可以尝试减少num_workers或设置pin_memory=False
  • non_blocking=True搭配使用:在将数据转移到GPU时,使用data = data.cuda(non_blocking=True)。这允许异步传输,CPU可以继续执行后续指令,而不是等待传输完成。但需要注意,后续的计算操作(如model(data))需要与传输同步,通常通过一个torch.cuda.stream.synchronize()或等待下一个CUDA操作隐式同步。

4. 分布式数据并行(DDP)训练中的“深水区”

DDP是进行多卡训练的主流和推荐方式,它效率高,但配置复杂,错误信息往往不直观。

4.1 端口冲突与进程初始化失败

错误信息常类似于:RuntimeError: Address already in use,或subprocess.CalledProcessError

常见坑点

  • master_port冲突:在多机训练或同一台机器上同时启动多个DDP任务时,如果都使用默认的29500端口,会导致端口冲突。
  • 进程未正确清理:某个DDP训练任务异常退出(如被kill -9),可能导致进程组未正确销毁,端口仍被占用。
  • 环境变量设置错误MASTER_ADDR(主节点地址)和MASTER_PORT设置错误,导致进程间无法通信。

避坑实操

  1. 动态选择端口:在启动脚本中,使用随机端口或从某个范围选择。
    import socket def find_free_port(): with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: s.bind(('', 0)) # 绑定到任意地址和随机端口 return s.getsockname()[1] master_port = find_free_port()
    然后在启动命令或初始化dist.init_process_group时传入这个端口。
  2. 使用torch.distributed.run/torchrun(推荐):这是PyTorch 1.9+推荐的启动方式。它自动处理端口发现、环境变量设置等。启动命令如:
    torchrun --nnodes=1 --nproc_per_node=2 --master_port=12345 train_script.py
    在脚本中,可以直接通过环境变量获取rankworld_size
    local_rank = int(os.environ['LOCAL_RANK']) rank = int(os.environ['RANK']) world_size = int(os.environ['WORLD_SIZE'])
  3. 确保进程组最终被销毁:在训练脚本的末尾,加入dist.destroy_process_group()。使用try...except...finally结构确保异常时也能执行清理。

4.2 模型参数同步与buffer注册

DDP的核心是前向传播后同步各卡间的梯度。但有些模型状态不需要计算梯度,却需要在各卡间保持一致,比如BatchNorm层中的running_mean和running_var。

常见坑点

  • 非参数张量(Buffer)未同步:如果你在模型中自定义了一些持久化状态(如一个可学习的温度参数tau,但将其定义为self.tau而不是nn.Parameter),DDP默认不会在进程间同步它们。这会导致不同卡上的模型状态不一致。
  • .forward()内部修改模型参数:这会导致梯度计算图混乱,并可能引发同步问题。

避坑实操

  • 正确注册Buffer:所有需要在进程间同步、且不参与梯度计算的状态,都应该用self.register_buffer('name', tensor)注册。
    class MyModel(nn.Module): def __init__(self): super().__init__() self.linear = nn.Linear(10, 5) # 这是一个需要同步的buffer self.register_buffer('running_feature_mean', torch.zeros(5)) # 这是一个需要学习和同步的参数 self.temperature = nn.Parameter(torch.ones(1) * 0.07)
  • 避免在前向传播中修改成员变量:将需要计算的状态作为局部变量或返回值。如果必须修改buffer,需注意其在不同进程中的一致性,必要时在修改后手动进行跨进程通信(dist.all_reduce),但这通常不是好设计。

4.3 损失函数与指标计算的“聚合坑”

在DDP中,损失是在每个GPU上独立计算的。如果你简单地打印每个GPU上的损失,会发现它们不一样,这是正常的,因为每个GPU处理的数据不同。

常见坑点

  • 错误地只在rank 0上计算指标:如果你只在排名为0的进程(通常是主进程)上计算验证集准确率或召回率,那么你计算的只是主进程看到的那个数据子集的指标,不能代表整个验证集。
  • 对损失值进行不必要的all_reduceloss.backward()内部已经完成了梯度的同步。如果你只是为了打印一个平均损失,可以在所有进程上计算loss.item(),然后在rank 0进程上收集并平均。但更简单的做法是使用torch.distributed的聚合函数。

避坑实操

  • 使用torch.distributed聚合指标:对于需要全局平均的标量(如损失、准确率),可以使用all_reduce
    def reduce_tensor(tensor): rt = tensor.clone() dist.all_reduce(rt, op=dist.ReduceOp.SUM) rt /= dist.get_world_size() return rt # 在训练循环中 loss = criterion(output, target) reduced_loss = reduce_tensor(loss.data) # 注意,这里聚合的是loss的数值,不是计算图 if rank == 0: print(f"Epoch {epoch}, Global Avg Loss: {reduced_loss.item():.4f}")
  • 对于需要整体统计的指标(如分类任务的Accuracy):更安全的做法是让每个进程计算自己数据子集的预测结果和标签,然后使用dist.all_gather将所有进程的结果收集到一起,再在rank 0上进行统一计算。这样可以避免因数据分布不均导致的指标偏差。

5. 内存管理与调试的“微观战场”

GPU内存不足(OOM)是训练大模型或处理大图像时最常见的错误。错误信息可能是CUDA out of memory

5.1 张量累积与中间变量驻留

常见坑点

  • 在循环中不断创建张量而未释放:例如,在验证阶段,将每个batch的预测结果追加到一个列表中(predictions.append(output.cpu())),如果验证集很大,这个列表会消耗大量CPU内存,甚至间接影响GPU(因为可能涉及GPU到CPU的传输缓存)。
  • 计算图未及时释放:为了进行梯度计算,PyTorch会保留前向传播的中间变量(计算图)。如果在训练循环外或不需要梯度的地方(如验证、测试)没有使用with torch.no_grad():,这些中间变量会一直累积,导致内存泄漏。
  • 大张量作为模型属性:如果将一个大张量(如一个巨大的嵌入表)作为模型的普通属性(而非nn.Parameter),并且在不同设备(CPU/GPU)间移动模型时处理不当,可能导致意外的内存复制。

避坑实操

  1. 使用torch.no_grad()model.eval():在验证和测试时,务必使用:
    @torch.no_grad() def validate(): model.eval() for data, target in val_loader: output = model(data.cuda()) # ... 计算指标,避免将output保留到大的列表中 model.train()
  2. 及时释放不需要的引用:在循环内,对于不再需要的大张量,可以显式将其设为None,并调用torch.cuda.empty_cache()(注意,这个函数会释放所有未使用的缓存,可能会带来性能波动,谨慎使用)。
  3. 使用del关键字:在作用域结束时,使用del variable可以提示Python垃圾回收器尽快回收对象。但这不是立即生效的,对于PyTorch张量,更有效的是将其移出GPU:variable = variable.cpu(),然后再del
  4. 梯度累积技巧:当GPU内存不足以容纳大的batch_size时,可以使用梯度累积。即使用小的batch_size进行多次前向传播,但只进行反向传播不更新参数(loss.backward()),累积多次的梯度后再进行一次优化器更新(optimizer.step()),然后清空梯度(optimizer.zero_grad())。这相当于模拟了一个大的batch_size,但峰值内存消耗仅与小batch_size相关。

5.2 混合精度训练(AMP)的“甜蜜负担”

自动混合精度(AMP)训练通过使用torch.cuda.amp,让模型的部分计算使用半精度(FP16),从而减少内存占用、加速计算。但它引入了新的复杂性。

常见坑点

  • 梯度下溢(Underflow):FP16的数值表示范围远小于FP32。非常小的梯度值在转换为FP16时可能变为0,导致参数无法更新。
  • 损失缩放(Loss Scaling):为了缓解梯度下溢,AMP引入了损失缩放。即在计算损失后,将其放大一个倍数,再进行反向传播,这样梯度也会被放大,转换到FP16时就不容易下溢。在优化器更新参数前,再将梯度缩放回去。如果缩放因子不合适,要么无法解决下溢,要么导致梯度爆炸(Overflow)。
  • 某些操作不支持FP16:如一些复杂的数学函数、自定义的CUDA核函数可能不支持FP16,强制使用会导致错误或精度下降。

避坑实操

from torch.cuda.amp import autocast, GradScaler scaler = GradScaler() # 梯度缩放器 model = Model().cuda() optimizer = torch.optim.Adam(model.parameters()) for data, target in train_loader: optimizer.zero_grad() with autocast(): # 在这个上下文管理器中,PyTorch会自动选择FP16或FP32进行计算 output = model(data.cuda()) loss = criterion(output, target.cuda()) scaler.scale(loss).backward() # 缩放损失,然后反向传播 scaler.step(optimizer) # 先unscale梯度,如果梯度没有inf/NaN,则更新参数 scaler.update() # 根据梯度情况,动态调整缩放因子
  • 监控梯度:定期检查梯度是否包含infNaNscaler.step(optimizer)内部会处理,如果发现梯度有问题,这一步会跳过参数更新,并且scaler.update()会减小缩放因子。
  • 定制化:对于已知的、在FP16下不稳定的层(如某些BatchNorm的变体),可以将其强制保持在FP32精度下:
    with autocast(): # 大部分计算用FP16 x = self.conv1(x) # 强制某层用FP32 with autocast(enabled=False): x = self.special_layer(x.float()) # 显式转换为float32 x = self.conv2(x)

6. 模型保存、加载与部署的“最后一公里”

费尽千辛万苦训练好的模型,最终要拿来用。这里面的坑也不少。

6.1 状态字典(state_dict)的“钥匙”对不上

保存模型通常用torch.save(model.state_dict(), 'model.pth'),加载用model.load_state_dict(torch.load('model.pth'))。错误常出现在加载时:Missing key(s) in state_dictUnexpected key(s) in state_dict

常见坑点

  • 模型结构发生变化:训练后修改了模型类(增、删、改了层名),导致保存的state_dict的键与当前模型结构对不上。
  • 多GPU训练保存的模型:使用DataParallelDDP包装后,模型参数名会带有module.前缀。例如,原来的conv1.weight会变成module.conv1.weight。如果保存时是DataParallel模型,加载到一个普通的单GPU模型上,就会因为前缀不匹配而报错。
  • 优化器状态字典的加载:同样存在上述问题,且优化器的state_dict还包含了参数组、超参数等信息,结构更复杂。

避坑实操

  1. 保存完整模型(谨慎使用)torch.save(model, 'model_full.pth')。这种方法保存了整个模型对象(包括结构定义),加载时直接model = torch.load('model_full.pth')。缺点是文件大,且严重依赖于保存时的类定义和环境,可移植性差,一般不建议用于生产部署。
  2. 处理module.前缀
    # 情况1:保存的是DataParallel/DDP模型,加载到单卡模型 state_dict = torch.load('ddp_model.pth') # 移除前缀 from collections import OrderedDict new_state_dict = OrderedDict() for k, v in state_dict.items(): name = k[7:] if k.startswith('module.') else k # 去掉'module.' new_state_dict[name] = v model.load_state_dict(new_state_dict) # 情况2:保存的是单卡模型,要加载到DataParallel模型 model = nn.DataParallel(model).cuda() # 直接加载会报错,因为单卡state_dict没有`module.`前缀 # 需要先加载到单卡模型,再用DataParallel包装(推荐) # 或者,在加载前给state_dict的key加上前缀(较麻烦)
  3. 部分加载:如果只想加载部分参数(例如,用预训练的主干网络初始化一个新模型),可以使用strict=False参数:
    model.load_state_dict(pretrained_dict, strict=False)
    这允许键不匹配,能加载的就加载,不能加载的忽略。但务必打印出缺失和多余的键,确认是否符合预期。

6.2 训练模式与推理模式的切换

模型有model.train()model.eval()两种模式。这会影响DropoutBatchNorm等层的行为。

常见坑点

  • 推理时忘记model.eval():导致Dropout层仍然随机丢弃神经元,BatchNorm层使用当前batch的统计量而非训练好的running统计量,使得推理结果随机且不稳定。
  • torch.no_grad()model.eval()混淆model.eval()改变的是模型内部某些层的行为torch.no_grad()with torch.no_grad():上下文管理器,是告诉PyTorch不跟踪计算图、不计算梯度,以节省内存和计算。两者通常同时使用,但目的不同。

避坑实操

  • 养成习惯:在训练循环开始前调用model.train();在验证、测试或实际部署推理前,调用model.eval()并置于torch.no_grad()上下文中。
  • 注意BatchNormtrack_running_stats:在微调(finetune)时,如果目标任务的数据分布与源任务差异极大,有时会冻结BatchNorm层的参数(weightbias)但让其继续更新running mean/var(即设置affine=False或冻结weight/bias但保持track_running_stats=True),或者干脆使用model.eval()固定所有统计量。这需要根据实际情况实验。

6.3 模型部署时的“静态图”挑战

为了提升推理速度、降低依赖,常需要将动态图模型(PyTorch)转换为静态图模型(如ONNX、TorchScript、TensorRT)。

常见坑点

  • 动态控制流无法直接导出:模型中的if-elsefor循环(其循环次数依赖于输入数据)在静态图中难以表示。ONNX等格式对此支持有限。
  • 张量形状变化:如果模型中间层的张量形状不是固定的(例如,由于输入图像尺寸可变),在导出时也需要特殊处理。
  • 自定义算子的支持:模型中如果使用了PyTorch没有原生支持、或ONNX标准中没有的算子,需要自己实现并注册转换函数。

避坑实操

  1. 简化模型用于导出:将动态逻辑(如根据输入决定路径)在导出前固定下来。或者,准备一个专门用于导出的“脚本化”版本模型。
  2. 使用torch.jit.tracetorch.jit.script
    • torch.jit.trace:用一个示例输入“跟踪”模型的执行路径,记录所有操作。它适用于没有数据依赖控制流的模型。
    • torch.jit.script:直接编译模型源代码。它能处理一些控制流,但对Python语言的子集支持有限。
    # trace 示例 example_input = torch.randn(1, 3, 224, 224).cuda() traced_model = torch.jit.trace(model.eval(), example_input) traced_model.save("traced_model.pt") # script 示例 (如果模型支持) scripted_model = torch.jit.script(model.eval()) scripted_model.save("scripted_model.pt")
  3. 导出ONNX:使用torch.onnx.export。关键是提供input_namesoutput_names和一个示例输入example_inputs。对于动态轴(如batch size或序列长度),需要使用dynamic_axes参数指定。
    torch.onnx.export( model.eval(), example_input, "model.onnx", input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}}, # 第0维是动态的 opset_version=13, # 指定ONNX算子集版本 )
    导出后,务必使用ONNX Runtime或其他工具验证导出的模型,确保其输出与原始PyTorch模型在相同输入下一致。这个过程充满了试错,需要耐心和细致的调试。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询