代码跑通之后怎么提升?模型改进、损失函数调优与实验管理完整指南
深度学习代码能跑通,说明你已经跨过了入门阶段最难的坎——环境没问题、数据能加载、前向反向能走通、loss 在下降。但一个很残酷的事实是:跑通和出成果之间,还隔着很长一段路。很多人在代码第一次跑起来之后,就不知道该干什么了。继续训练?不懂怎么调。改模型?怕改坏。换损失函数?连代码在哪都不知道。最后只能反复跑同一个脚本,看着差不多的指标,心态一点点崩掉。
这篇文章要解决的,就是“代码跑通后做什么、怎么做”这个阶段的问题。我会按照一条可执行的路径,把跑通之后的工作拆成几个阶段:先建立基线和诊断流程,再谈模型改进和修改损失函数的具体方法,然后是实验管理、项目演示和工程化封装。适合已经能跑通基础代码、但还没有系统改进经验的读者,也适合研究生开题、项目交付、竞赛复现阶段需要“从能跑到跑好”的人。
文章里会给出通用的改进方法论,以及基于 PyTorch 的损失函数修改代码示例。所有代码都是模板,需要按你自己的模型结构、任务类型和数据集做调整。
1. 核心工作流速览
“代码跑通后”不是一个动作,而是一套流程。你先要知道这套流程里包含哪些阶段,再对照自己的项目看缺了哪一段。
| 阶段 | 核心问题 | 输入 | 输出 |
|---|---|---|---|
| 基线锁定 | 当前代码的真实效果是多少? | 训练脚本、数据集、验证脚本 | 一组可复现的 baseline 指标 |
| 诊断分析 | 模型是欠拟合还是过拟合? | 训练日志、验证集、样本可视化 | 改进方向清单 |
| 模型改进 | 网络结构哪里可以改? | baseline 代码、诊断结论 | 结构改动后的新实验 |
| 修改损失 | 现有 loss 是否匹配任务目标? | 原损失函数、业务指标 | 新损失函数代码和实验 |
| 实验管理 | 多组实验如何对比?如何避免忘记参数? | 实验配置、日志、权重文件 | 可追溯的实验记录表 |
| 项目演示 | 如何让别人看懂你的改进? | 指标对比、可视化结果、推理 Demo | 可复现的演示脚本和说明文档 |
| 工程封装 | 如何把代码交给别人或部署上线? | 训练好的权重、推理代码 | 接口服务、批量推理脚本、部署文档 |
这七个阶段不需要严格串行。实际项目里,诊断分析和模型改进经常是循环的:改一次 loss,跑一次实验,看指标,再决定下一步。但基线和诊断必须最先做,否则后面的所有“改进”都缺少判断依据。
2. 第一步:锁定基线与复现记录
代码第一次跑通时,你可能并没有认真记录当时的配置和结果。这是后续改进最大的隐患——你改了三个地方,指标涨了 2%,但你根本不知道是哪个改动起了作用。
2.1 先记录一组完整的基线信息
无论代码多么原始,都要立刻把当时的环境和结果固定下来:
# 记录当前环境信息 pip freeze > requirements_lock.txt # 记录 GPU 驱动和 CUDA 版本 nvidia-smi # 记录 PyTorch 版本 python -c "import torch; print(torch.__version__, torch.version.cuda)"建议同时记录以下内容到项目根目录的BASELINE.md文件里:
- 数据集名称、训练集/验证集划分方式、是否存在数据预处理顺序
- 优化器名称、学习率、batch size、训练轮数
- 随机种子
- 模型的参数量
- 训练集 loss 最终值、验证集指标最终值
- 单轮训练耗时、总训练时长
- 显卡型号和显存占用峰值
- 是否使用了混合精度、梯度累积等加速手段
这个过程不能省。没有基线,后面做任何模型改进都等于在黑箱里猜。
2.2 复现验证
如果项目脚本允许,建议在锁定基线后,用相同的随机种子重新跑一次短训练,确认指标能稳定复现到相近水平。特别要注意数据加载时是否真的有 shuffle,以及验证集是否固定。有些项目跑通时指标“很好”,其实是验证集被 shuffle 了,每次评估都在用不同的数据,这种基线没有任何参考价值。
从复现验证开始,你就进入了“以实验记录驱动改进”的工作模式,而不是“凭感觉改代码”。
3. 第二步:诊断——先判断模型处于什么状态
直接在原代码上大改结构、换损失函数,常常是白费功夫。你需要先回答一个问题:
- 如果训练集 loss 很高,验证集 loss 也很高:模型欠拟合,能力不够或没训练充分。
- 如果训练集 loss 很低,验证集 loss 很高:模型过拟合,泛化能力不足。
- 如果训练集和验证集 loss 都很低,但业务指标不达标:训练目标和评估目标不匹配,通常需要修改损失函数或评估方式。
这是一个基础而关键的诊断。很多同学上来就换BCEWithLogitsLoss为Focal Loss,结果发现模型本来就是欠拟合,换什么损失都救不回来。
3.1 数据样本可视化检查
在动手改模型之前,先把数据检查一遍。图像任务可以保存一批训练样本和标注叠加图,文本任务随机打印几条样本和标签。常见问题是:
- 标签是否错位
- 归一化是否写错(例如图像增强后像素范围不是 0~1 而是 0~255,或反过来)
- 类别不均衡但代码没有处理
- 验证集和训练集存在数据泄露
数据层面的问题,往往比模型结构问题更影响最终效果。一上来就改网络结构,忽略了数据本身,是改进阶段最常踩的坑。
3.2 日志可视化
如果你还没有把 loss、学习率、验证指标记录到 TensorBoard 或 wandb,建议尽快补上。最少也要把每个 epoch 的 train loss、val loss、val metric 写道 CSV 文件里:
import csv with open("training_log.csv", "w", newline="") as f: writer = csv.writer(f) writer.writerow(["epoch", "train_loss", "val_loss", "val_acc"])有了曲线,你才能判断改进是否真的有效,也才能在论文或项目汇报中展示变化过程。
4. 模型改进:不要推翻重写,要做模块化替换
模型改进的核心不是“换一个更大的网络”,而是找到当前网络上可替换、可验证的模块。以图像分割常见的 UNet 为例,改进思路通常有四类:
4.1 编码器替换
把原来自己写的卷积编码器换成预训练的 ResNet、EfficientNet、ConvNeXt,借助 ImageNet 预训练权重提升特征提取能力。改动范围通常是网络初始化和forward的输入输出通道,不改整体 U 形结构。
4.2 特征融合改进
在编码器和解码器之间的跳跃连接处做改进,例如加入注意力机制、特征金字塔结构、多尺度特征融合模块。这类改动的风险低,因为主结构不变,只是增加或调整跳跃连接的实现。
4.3 解码器改进
改进上采样方式,例如把简单的双线性插值换成可学习的转置卷积,或者加入深度可分离卷积降低参数量。
4.4 模块替换之后必须做对照实验
任何改动都要保持“单一变量”原则:一次只改一个模块,其他保持不变。改完之后,用完全相同的训练配置跑实验,对比新 baseline 和旧 baseline 的验证指标。如果指标涨了,再叠加下一个改动。
建议用如下格式记录每个改动:
| 实验编号 | 改动内容 | 参数量 | 验证指标 | 与基线对比 |
|---|---|---|---|---|
| baseline | 原始代码 | 12.3M | 0.812 | - |
| exp1 | 编码器换成 ResNet34 | 21.8M | 0.835 | +2.3% |
| exp2 | 跳跃连接加注意力 | 13.0M | 0.824 | +1.2% |
| exp3 | exp1 + exp2 | 22.5M | 0.847 | +3.5% |
模型改进是实验科学,不是灵感创作。所有改动都要能回溯、能对比、能复现。
5. 修改损失函数:从公式到 PyTorch 实现
修改损失函数是“代码跑通后”最常见的改进方式,也是最容易写错的地方。这一节给你一套可操作的流程。
5.1 先明确为什么改损失函数
只有当现有 loss 和业务目标不一致,或者现有 loss 在训练中表现异常时,才需要修改。举例来说:
- 目标检测中正负样本极度不平衡,交叉熵 loss 训练困难,可以改用 Focal Loss。
- 医学图像分割中前景区域很小,Dice Loss 往往比 BCE Loss 更稳定。
- 人脸识别任务需要拉近同类、推远异类,直接用交叉熵不够,要做基于 margin 的 loss。
- 回归任务对离群点敏感,MSE Loss 会让模型被个别大误差样本主导,可以换 Huber Loss。
所以,修改损失函数的第一步,是写出当前任务的评估指标,然后反推损失函数应该重点优化什么。
5.2 损失函数修改的最小实现流程
以 PyTorch 为例,假设你现在用的是torch.nn.CrossEntropyLoss,想改成带 difficulty 权重的 Focal Loss,完整流程如下:
5.2.1 定义新的损失函数类
import torch import torch.nn as nn import torch.nn.functional as F class FocalLoss(nn.Module): """ Focal Loss for multi-class classification. 公式: FL = -alpha_t * (1 - p_t)^gamma * log(p_t) 其中 p_t 是正确类别的预测概率,alpha_t 是类别权重。 """ def __init__(self, alpha=None, gamma=2.0, reduction="mean"): super().__init__() self.alpha = alpha self.gamma = gamma self.reduction = reduction def forward(self, logits, targets): ce_loss = F.cross_entropy( logits, targets, weight=self.alpha, reduction="none" ) # 获取正确类别的预测概率 p_t prob = F.softmax(logits, dim=1) p_t = prob.gather(1, targets.unsqueeze(1)).squeeze(1) # Focal loss 权重因子 focal_weight = (1 - p_t) ** self.gamma loss = focal_weight * ce_loss if self.reduction == "mean": return loss.mean() elif self.reduction == "sum": return loss.sum() else: return loss5.2.2 在训练脚本中替换
# 原来的写法 criterion = nn.CrossEntropyLoss() # 新的写法 # alpha 可以传入一个 Tensor,形状是 [num_classes] criterion = FocalLoss( alpha=None, # 如果不做类别加权,保持 None gamma=2.0, # gamma 越大,对难样本的关注越大 reduction="mean" )5.2.3 小规模验证新损失
换 loss 之后,不要直接启动完整训练。先跑 5~10 个 epoch 的小实验,观察 train loss 是否在合理范围内下降,以及数值是否出现 NaN 或爆炸。
5.3 调试损失函数的通用技巧
- 打印每次 forward 的 loss 数值,确认不是
nan或固定值。 - 用一个小 batch 数据跑一次
backward(),检查梯度是否存在。 - 先关闭权重衰减和高级优化器,用最朴素的 SGD 验证 loss 能否下降。
- 如果新 loss 是多个 loss 的组合,例如
L = L_cls + 0.1 * L_dice,建议分别记录每个子 loss 的数值,确认各项都在合理范围。
# 多 loss 组合示例 loss = cls_loss + 0.1 * dice_loss print(f"cls_loss: {cls_loss.item():.4f}, dice_loss: {dice_loss.item():.4f}")5.4 修改损失函数时最容易犯的错误
| 错误类型 | 现象 | 排查方式 |
|---|---|---|
| 维度错误 | gather或unsqueeze报错 | 打印logits.shape和targets.shape |
| 类型错误 | label 是 float 而 loss 需要 long | targets.long()转换 |
| 数值不稳定 | 出现nan | 检查是否对prob=0取 log,加eps平滑 |
| 权重不匹配 | alpha维度不等于类别数 | 检查alpha的 shape |
| 梯度断链 | loss 来自torch.no_grad()区域 | 检查自定义 loss 中是否误用.detach() |
修改损失函数是高风险高收益的操作。写正确的收益是把训练目标拉向真实的评估指标;写错了,整个训练过程都会无效,而且不容易发现。
6. 实验管理与可重复性
当你开始做模型改进和损失函数修改之后,一定会产生大量实验。手动管理几个目录很容易,但实验数量超过 10 个后,就会开始混淆“哪个脚本对应哪个结果”。建议从第一天就建立轻量级实验管理机制。
6.1 目录结构示例
experiments/ ├── baseline/ │ ├── config.yaml │ ├── train_log.csv │ ├── best_model.pth │ └── eval_result.json ├── exp1_resnet34/ │ ├── config.yaml │ ├── train_log.csv │ ├── best_model.pth │ └── eval_result.json └── exp2_focal_loss/ ├── config.yaml ├── train_log.csv ├── best_model.pth └── eval_result.json每次启动训练前,把当前配置写到一个config.yaml文件中,训练结束后把验证集指标写入eval_result.json。不需要复杂的平台,文件系统就是最简单的实验管理工具。
# config.yaml 示例 model: name: unet_resnet34 in_channels: 3 num_classes: 5 dataset: name: custom_seg_dataset train_dir: ./data/train val_dir: ./data/val training: batch_size: 8 epochs: 50 optimizer: adamw learning_rate: 0.001 scheduler: cosine seed: 42 mixed_precision: true6.2 每次训练必须记录的指标
- 训练集 loss 曲线
- 验证集 loss 曲线
- 验证集业务指标(精确率、召回率、mAP、Dice、IoU 等,按任务选择)
- 单 epoch 耗时
- 显存占用峰值
- 最终权重文件路径
这些记录的价值会在 20 个实验之后体现出来。你会发现,真正能复现的实验,靠的都是这种“没什么技术含量”的记录习惯。
7. 项目演示与结果汇报
代码跑通后,改了一堆实验,但如果不能把结果清楚展示出来,你的工作在团队里或论文里都难以得到认可。项目演示不是简单的“发给别人代码”,而是要让人能快速理解你的改进有效在哪里。
7.1 指标对比表
把所有实验的指标汇总到一张表,按实验编号排列。对比维度包括:参数量、训练时长、验证指标、显存占用。这样评审人一眼就能看到每个改动带来的收益和成本。
7.2 可视化结果图
图像类任务,把 baseline 和最新改进模型的预测结果并排输出,每张图包含:
- 原始输入
- 真实标注
- baseline 预测
- 改进模型预测
这样可以直观看出改进模型在哪些样本上表现更好,在哪些样本上仍然失败。
import torch import matplotlib.pyplot as plt # 假设 model1 是 baseline,model2 是改进模型 model1.eval() model2.eval() with torch.no_grad(): pred1 = torch.argmax(model1(x).logits, dim=1).cpu().numpy() pred2 = torch.argmax(model2(x).logits, dim=1).cpu().numpy() # 保存对比图 fig, axes = plt.subplots(2, 2, figsize=(12, 12)) axes[0][0].imshow(x[0].permute(1, 2, 0).cpu().numpy()) axes[0][0].set_title("Input") axes[0][1].imshow(y[0].cpu().numpy()) axes[0][1].set_title("Ground Truth") axes[1][0].imshow(pred1[0]) axes[1][0].set_title("Baseline") axes[1][1].imshow(pred2[0]) axes[1][1].set_title("Improved") plt.savefig("comparison.png")7.3 写一个 README 说明文档
无论如何交付代码,都建议在项目根目录放一个README.md,内容包含:
- 环境安装命令
- 数据准备方式
- baseline 复现命令
- 改进实验复现命令
- 结果指标表
- 常见问题说明
这个文档本身就是项目演示的一部分,也是你“代码跑通后”工作成果的可交付形态。
8. 接口封装与批量推理
实验阶段验证通过之后,往往还需要把模型交给别人测试,或者对接业务系统。这时需要把训练代码和推理代码分离,封装成可独立调用的接口。
8.1 推理脚本分离
新建一个inference.py,只负责加载权重和推理,不再包含训练逻辑:
import torch from PIL import Image from torchvision import transforms def load_model(model, weight_path, device): model.load_state_dict(torch.load(weight_path, map_location=device)["model_state_dict"]) model.to(device) model.eval() return model def predict_image(model, image_path, device): transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.ToTensor(), transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]), ]) image = Image.open(image_path).convert("RGB") input_tensor = transform(image).unsqueeze(0).to(device) with torch.no_grad(): output = model(input_tensor) pred = torch.argmax(output, dim=1) return pred8.2 批量推理目录
如果要对一个目录下的所有图片做批量处理,可以写一个循环,并记录处理进度和失败情况:
import os from pathlib import Path input_dir = Path("./test_images") output_dir = Path("./test_outputs") output_dir.mkdir(exist_ok=True) for image_path in sorted(input_dir.glob("*.jpg")): result = predict_image(model, str(image_path), device) # 保存结果或写入记录 print(f"processed: {image_path.name}")批量处理时,建议每处理一张图片就把状态写入日志文件,便于中断后从断点继续,而不是重头再来。
8.3 简单 API 封装
如果对方需要以接口方式调用模型,可以用 FastAPI 封装一个极简服务。代码需要按实际模型和接口规范调整:
from fastapi import FastAPI, UploadFile, File import io from PIL import Image app = FastAPI() @app.post("/predict") async def predict(file: UploadFile = File(...)): image_bytes = await file.read() image = Image.open(io.BytesIO(image_bytes)) # 调用你的 predict_image 函数 # result = predict_image(model, image, device) # return {"prediction": result.tolist()} return {"message": "接口需要按实际模型实现"}接口封装的目的,是让“代码跑通”成为一个可被别人使用的能力,而不是停留在 Jupyter Notebook 里的试验品。
9. 常见问题与排查方法
从“跑通”到“跑好”,中间会遇到各种问题。下面是一份高频问题清单。
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 修改损失函数后 loss 为 NaN | 对 0 概率取 log,或学习率过大 | 检查损失函数内部数值,打印每一步输出 | 对概率加eps=1e-8,降低学习率 |
| 改了损失函数但指标不变 | 新 loss 数值上改了但梯度方向未变化,或代码中未真正使用新 loss | 检查训练脚本中criterion是否被替换 | 确认criterion变量指向新损失函数 |
| 增加模块后参数量暴涨 | 模块实现中使用了过大的中间通道 | 打印模型每层参数 | 调整通道数或使用深度可分离卷积 |
| 模型改进后指标下降 | 改动破坏了原有结构的初始化或梯度传播 | 回到 baseline 确认可复现,再做单变量实验 | 一次只改一个模块,使用预训练权重 |
| 训练速度变慢 | 混合精度未开启,或数据加载成瓶颈 | 查看 GPU 利用率和 CPU 数据加载时间 | 开启 AMP,使用num_workers和预加载 |
| 验证集指标与训练集差距大 | 过拟合 | 对比训练集和验证集 loss | 增加数据增强、正则化、早停 |
| 复现实验时指标对不上 | 随机种子未固定,或数据加载顺序变了 | 固定 seed,固定数据 shuffle | 设置torch.manual_seed(42)等 |
| 接口调用时显存溢满 | 服务常驻导致显存累积 | 查看nvidia-smi占用 | 每次请求后释放缓存,控制并发数 |
| 批量处理中途中断 | 缺少断点续跑机制 | 查看日志文件 | 增加跳过已处理文件逻辑 |
| 更换 loss 后训练 loss 不下降 | 损失函数公式写错或梯度被阻断 | 在 backward 之前打印 loss 和梯度 | 用随机小 batch 调试前向和反向 |
10. 最佳实践总结
这一节把前面所有内容浓缩成一份可直接照做的清单。代码跑通之后,按这套顺序执行,能最大程度减少无效实验和重复劳动。
- 不管代码多乱,先固定基线。记录环境、配置、指标、权重路径,把“可复现”作为第一优先级。
- 在改模型之前,先诊断当前是欠拟合、过拟合还是训练目标不匹配。没有诊断的改进都是盲改。
- 模型改进采用模块化替换,一次只改一个模块。使用预训练权重时注意输入输出通道对齐,保持与 baseline 相同的训练配置。
- 修改损失函数前,先写清楚业务指标和当前 loss 的关系。换成 Focal Loss、Dice Loss 等新 loss 时,先在小规模实验上验证数值稳定性,再跑完整训练。
- 多 loss 组合时,分别记录每个子 loss 的数值。不要用一个总的 loss 掩盖单个子 loss 的异常。
- 每次实验保存独立的配置文件和验证结果。实验数量增多后,你会发现这些文件比权重本身还重要。
- 使用混合精度训练可以明显降低显存占用和加速训练,但新加的损失函数需要在 AMP 下测试数值稳定性。
- 项目演示永远准备好三样东西:指标对比表、可视化结果、可复现命令。缺少任何一个,别人都无法判断你的工作是否可信。
- 批量推理必须支持中断恢复。数据量大时,不要把整个数据集一次性载入内存。
- 接口服务常驻时,要关注显存释放和并发控制。不要把训练代码直接复制成服务代码。
代码跑通那一刻的成就感,很容易让人误以为工作已经完成了大半。但真正的技术提升恰恰是从跑通之后开始的——你需要学会诊断模型、设计改进实验、修改损失函数并验证效果,这些能力没法靠“跑通一个开源项目”获得,只能靠一个一个受控实验积累。
如果你的代码已经能跑通,下一步不是急着换更大的模型,也不是反复调学习率,而是按这篇文章的流程,先把基线固定下来,再做一次诊断,然后选择一个最简单的模块改动或损失函数修改,验证“改动-效果”的因果链。只要这一条链路打通,后面所有模型改进就都有了方法可循。
推荐顺序是:先照着本文第四节做一次模块化模型改进,再按照第五节改一次损失函数。如果这两步走通,你的项目就从“能跑”进入“能改进”的阶段了。之后再考虑接口封装、批量推理和项目演示,把实验结果变成可交付的成果。
