把脚本写成能复用的工具
约 14 分钟
前面的代码大多是「写一次跑一次」。这一节讲怎么让脚本能被反复使用、参数可调、结果可复现。
一、用命令行参数代替改代码
❌ 每次改需求就改源码:
INPUT = "data/train.csv" # 换个文件就要改这里
EPOCHS = 10
✅ 用 argparse 把它们变成参数:
import argparse
def main():
p = argparse.ArgumentParser(description="训练一个分类模型")
p.add_argument("--input", required=True, help="训练数据 CSV 路径")
p.add_argument("--epochs", type=int, default=10)
p.add_argument("--lr", type=float, default=0.001)
p.add_argument("--verbose", action="store_true", help="打印详细日志")
args = p.parse_args()
print(f"训练 {args.epochs} 轮,学习率 {args.lr}")
if __name__ == "__main__":
main()
python train.py --input data/train.csv --epochs 20 --lr 0.01 --verbose
python train.py --help # 自动生成帮助文档
--help 是白送的——argparse 根据你写的 help 文字自动生成,别人不用读源码就知道怎么用。
二、用 logging 代替 print
import logging
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(message)s",
)
log = logging.getLogger(__name__)
log.debug("细节信息,调试时才看")
log.info("正常进度:已处理 1000 条")
log.warning("发现 3 条缺失值,已跳过")
log.error("文件读取失败")
比 print 强在哪:
| 好处 | 说明 |
|---|---|
| 分级 | 改一个参数就能切换详细程度,不用删代码 |
| 带时间戳 | 知道每一步花了多久 |
| 可写文件 | 加一个 handler 就能存日志 |
| 可按模块控制 | 只看某个模块的日志 |
一条经验:临时调试用 print,正式脚本用 logging。
三、保证可复现:固定随机种子
import random, numpy as np, torch
def set_seed(seed=42):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
为什么重要:不固定种子,同一份代码每次跑出的结果都不一样,你无法判断「这次好一点」是改进起了作用,还是纯属运气。
科研和调参的前提,是结果可复现。
注意:GPU 上的某些运算本身不确定,完全复现还需要额外设置,但固定种子已经能解决绝大部分问题。
四、把配置和代码分开
参数多了之后,命令行会很长。改用配置文件:
import json
with open("config.json", encoding="utf-8") as f:
cfg = json.load(f)
print(cfg["epochs"], cfg["lr"])
好处:每次实验存一份配置文件,就知道当时用的什么参数。 这是实验管理最基础的做法。
五、一个完整的脚本骨架
"""训练脚本:读数据、训练、保存模型。"""
import argparse, logging, json
from pathlib import Path
log = logging.getLogger(__name__)
def load_data(path: Path):
if not path.exists():
raise FileNotFoundError(f"数据文件不存在:{path}")
...
def main():
p = argparse.ArgumentParser()
p.add_argument("--input", type=Path, required=True)
p.add_argument("--out", type=Path, default=Path("model.pt"))
args = p.parse_args()
logging.basicConfig(level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(message)s")
set_seed(42)
log.info("开始加载数据:%s", args.input)
data = load_data(args.input)
log.info("完成,共 %d 条", len(data))
if __name__ == "__main__":
main()
注意用了 pathlib.Path 而不是字符串拼路径——它跨平台、能直接判断存在性、拼接用 / 运算符(base / "data" / "train.csv"),比 os.path.join 清爽得多。
练习:把你之前写过的一个脚本改造一下——加上 argparse 参数、把 print 换成 logging、固定随机种子。
小纸条
为什么正式脚本要用 logging 而不是 print?固定随机种子解决什么问题?
登录 后可看答案