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