所有博客
最佳实践训练Jupyter

从 Jupyter 到生产训练:5 个常见错误和解决方案

很多团队在 Jupyter Notebook 里实验通过后,直接拿来跑生产训练。这篇文章总结了 5 个最容易踩的坑,以及如何避免。

RRunWhere Team·2025年5月10日·6 分钟阅读
从 Jupyter 到生产训练:5 个常见错误和解决方案

Jupyter Notebook 是实验神器,但直接拿来跑生产训练,你会遇到比想象中多得多的问题。我们在平台上看到大量训练任务失败,根因都可以归结为「把实验代码当成了生产代码」。

错误一:硬编码文件路径

Jupyter 里写 /home/user/data/train.csv 很正常,但换一台机器就会报 FileNotFoundError

解决方案:使用环境变量或配置文件管理路径。

import os

# ❌ 硬编码路径
data_path = "/home/user/data/train.csv"

# ✅ 环境变量
data_path = os.environ.get("DATA_PATH", "./data/train.csv")

错误二:忽略随机种子

Notebook 里反复运行 cell,每次结果都不一样,但你可能从来没注意过。到了生产训练,实验不可复现会让你浪费大量时间排查「为什么效果变差了」。

解决方案:在脚本入口统一设置随机种子。

import random
import numpy as np
import torch

def set_seed(seed=42):
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)

set_seed()

错误三:没有 Checkpoint 机制

Notebook 里训练通常不超过几十分钟,挂了重跑就好。生产训练动辄几小时甚至几天,没有 checkpoint 意味着一旦中断就要从头开始。

解决方案:每 N 个 epoch 保存一次 checkpoint,训练脚本启动时自动检测并恢复。

checkpoint_dir = os.environ.get("CHECKPOINT_DIR", "./checkpoints")

# 保存
if epoch % save_every == 0:
    torch.save({
        "epoch": epoch,
        "model_state": model.state_dict(),
        "optimizer_state": optimizer.state_dict(),
        "loss": loss,
    }, f"{checkpoint_dir}/epoch_{epoch}.pt")

# 恢复
latest = find_latest_checkpoint(checkpoint_dir)
if latest:
    checkpoint = torch.load(latest)
    model.load_state_dict(checkpoint["model_state"])
    start_epoch = checkpoint["epoch"] + 1

错误四:忽略内存泄漏

Notebook 的 cell 之间共享命名空间,变量不会被回收。跑几十个 cell 后内存占用越来越高,但你可能以为这是正常的。到了生产训练,同样的代码可能在第 50 个 epoch 时 OOM。

解决方案

  • del 及时释放不再需要的大 tensor
  • 训练循环中用 torch.cuda.empty_cache() 清理缓存
  • 使用 gradient_accumulation_steps 减小单步显存占用

错误五:没有日志和监控

Notebook 里用 print() 输出结果,训练结束后结果就在页面上。生产训练跑在远程 GPU 上,没有实时日志你根本不知道训练到哪了、loss 是否正常。

解决方案

  • 使用 Python logging 模块替代 print
  • 集成 TensorBoard 或 Weights & Biases 做训练监控
  • 关键指标(loss、lr、GPU 利用率)每 N 步记录一次
import logging
logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s [%(levelname)s] %(message)s",
    handlers=[
        logging.FileHandler("train.log"),
        logging.StreamHandler(),
    ],
)
logger = logging.getLogger(__name__)

for epoch in range(num_epochs):
    for step, batch in enumerate(dataloader):
        loss = train_step(model, batch)
        if step % log_every == 0:
            logger.info(f"Epoch {epoch} Step {step} Loss {loss:.4f}")

总结

从 Jupyter 到生产训练,本质上是从「交互式实验」到「自动化流水线」的转变。把以上 5 个问题解决好,你的训练任务会稳定得多,排查问题的时间也会大幅减少。

RunWhere 的任务提交流程天然解决了其中一部分问题:环境隔离避免了路径硬编码,自动停机避免了资源浪费,日志自动收集避免了信息丢失。但代码层面的最佳实践,仍然需要你自己来落实。

想体验 RunWhere.ai?

Free 永久免费,自带云账号,无需信用卡,30 秒启动第一个 GPU 任务。

免费开始 →