Python 编程

把脚本写成能复用的工具

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?固定随机种子解决什么问题?

登录 后可看答案