Python 工程化重构:从 print 脚本到 pytest 项目
- 为什么做这件事
- 先拆再补
- 类型标注:给自己和机器一个共同的契约
- print 换 logging
- 写 pytest,不要靠人眼验证
- pre-commit:提交前把门关上
- pyproject.toml:别再弄五六个配置文件了
- 验证结果
- 和 Android 熟悉的那套对上
- 总结
- 系列文章
这是 AI 实操第 8 天。第 1 天的 MNIST 训练脚本,172 行塞一个文件里,print 满天飞,没有测试,没有类型标注。今天把它拆了重写:类型标注、logging、pytest、pre-commit 一条龙。 本文首发地址 https://h89.cn/archives/674.html 项目地址 https://gitee.com/chenjim/cockpit-ai-from-zero
为什么做这件事
第 1 天写了个能跑的 MNIST 识别——功能没问题,但交不出手。
不是我矫情。想象一下:一个月后自己回来改这个 172 行的文件,要花多久才能找到训练循环在哪?换个 batch size 要滚几屏?模型结构改了,怎么知道推不碎全连接层的维度?
更别说团队场景了。代码扔到 Git 上,同事怎么知道这个函数返回什么?改了数据加载怎么保证不影响训练?全凭自觉。
这些问题不解决,AI 工程师这岗位就别想着"工程化"三个字。这篇文章记录了我把一段能跑的脚本改成一个能交付的项目,踩了什么坑。
先拆再补
172 行不是问题,问题是这 172 行里混着 4 件事:
- 模型结构
- 数据加载
- 训练循环
- 训练入口
改一个就得看全部,不敢动。拆开:
src/mnist_engineering/
├── model.py # 模型
├── data.py # 数据
├── train.py # 训练
├── utils.py # 日志配置
└── main.py # 入口
每个文件 30~50 行。打开 model.py 就想看模型,打开 data.py 就是数据。这个拆法叫 src layout(src 目录布局,把代码统一放进 src/ 下的包目录),Python 社区用了很多年,确实好用。
拆完发现 main.py 里还有个 __init__.py——
from mnist_engineering.model import MNISTCNN
from mnist_engineering.data import get_dataloaders
from mnist_engineering.train import train_one_epoch, evaluate
from mnist_engineering.utils import configure_logging
__all__ = ["MNISTCNN", "get_dataloaders", "train_one_epoch", "evaluate", "configure_logging"]
外面的人 import mnist_engineering,IDE 就知道该补什么。顺手的事,但不做就是没有。
类型标注:给自己和机器一个共同的契约
翻一下第 1 天的代码:
def get_dataloaders(batch_size: int = 64):
返回值是什么?不清楚。看实现才发现返回两个 DataLoader。就这一行就能带歪读代码的人。
加全:
def train_one_epoch(
model: nn.Module,
loader: DataLoader[Any],
optimizer: torch.optim.Optimizer,
criterion: nn.Module,
device: torch.device,
) -> float:
-> float 明确告诉调用者:我返回的是平均 loss。看了签名不用读实现。
但类型标注只是给人看的,更关键的是给 mypy 看:
[tool.mypy]
strict = true
strict = true 开了,所有函数都得有返回类型,不能用 Any 糊弄(除非你真的不知道)。跑一遍:
$ mypy src tests
Success: no issues found in 11 source files
之后谁改了代码里类型不对,CI 直接拦住,不用等人 review 出来。
踩坑:PyTorch 的 DataLoader 是泛型类。在 mypy strict 模式下不写类型参数就报错:
def train(train_loader: DataLoader) -> None: ... # 错,缺类型参数
正确写法:
from __future__ import annotations
def train(train_loader: DataLoader[Any]) -> None: ...
这个 from __future__ import annotations 是 PEP 563(Python 增强提案 563,让类型注解成为惰性求值的字符串)。运行时不会有 import 开销,也不会因为循环引用炸掉。Python 4.0 会默认开启,现在加 __future__ import 是提前适配。
print 换 logging
12 处 print,关不掉。想少看日志?没法过滤。看到一行"模型已保存"谁打的?不知道。
换成 logging 之后:
旧:MNIST 数据已存在 (跳过下载)
新:16:40:17 [INFO ] mnist_engineering.data: MNIST 数据已存在,跳过下载
多了时间戳和模块名。别小看这个——生产环境里排查问题,差的就是这行信息。
有两个决策:
- 输出到 stderr 不是 stdout。stdout 留给模型结果,日志是运维信息,不该混一起。
- 模块级 logger。每个文件一行
logger = logging.getLogger(__name__),自动带模块路径。
配置写一次 configure_logging() 在入口调用,全局生效。
写 pytest,不要靠人眼验证
11 个测试,分成三组。
模型结构测试(4 个)——输入 28x28 灰度图,输出 [4, 10] 的 logits:
def test_output_shape(self, model, device):
x = torch.randn(4, 1, 28, 28).to(device)
with torch.no_grad():
output = model(x)
assert output.shape == (4, 10)
参数数量在合理范围(简单 CNN 应该 1~20 万)、所有参数都可训练、单张图前向不抛异常——每个都是 10 行以内的断言,但覆盖了模型结构最脆弱的几个点。
数据加载测试(5 个)——训练集 60000 张测试集 10000 张、batch 形状 [32, 1, 28, 28]、标签 0~9 范围内。数据层出问题基本不影响训练跑通,但影响结果正确性,这类测试防的是隐蔽错误。
训练流程测试(2 个)
最关键的是"过拟合一个 batch"。完整测试代码在 tests/test_train.py,这里只保留核心逻辑:
def test_overfit_single_batch(device, train_loader):
images, labels = next(iter(train_loader))
model = MNISTCNN().to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
for _ in range(200):
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
acc = (outputs.argmax(dim=1) == labels).float().mean().item()
assert acc >= 0.95
这个测试的原理是:如果你的模型、损失函数、优化器三者的组合能在一个 batch 上过拟合到 95% 以上,说明训练逻辑没写错。否则别跑 60000 张数据了,改 bug 先。
run 一下全部通过了,3.55 秒。用了 scope="session" 的 fixture——模型和数据只初始化一次,11 个测试共享,没必要每个测试重载一遍。
pre-commit:提交前把门关上
先说清楚,pre-commit 不是 git 自带的。
git 自带的是一个底层机制叫 git hooks——藏在 .git/hooks/ 目录下的 shell 脚本(名字就叫 pre-commit、pre-push、post-commit 等)。你可以在里面写 shell 代码,commit 时自动执行。问题是手写 shell 麻烦、跨平台容易炸、团队每个人都得自己配一遍。
pre-commit 是第三方工具,pip install pre-commit 安装。它的作用是:你写好 .pre-commit-config.yaml,它自动帮你生成跨平台的 git hook 脚本。之后就变成了:
git commit → black 格式化 → ruff 检查 → mypy 类型检查 → 全过才让提交
我的配置写了三个:
repos:
- repo: https://github.com/psf/black
rev: 24.4.2
hooks:
- id: black
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.5.1
hooks:
- id: ruff
args: [--fix]
- id: ruff-format
- repo: https://github.com/pre-commit/mirrors-mypy
rev: v1.10.0
hooks:
- id: mypy
additional_dependencies: [torch, torchvision]
激活就一行:pre-commit install
这行命令只做一件事:在 .git/hooks/pre-commit 写一个脚本。之后每次 git commit,这个脚本自动读 yaml 配置,逐个跑定义的 hook。
三个 hook 的执行顺序是:black → ruff → mypy。black 先格式化,ruff 再 lint(修复 import 排序等),最后 mypy 做类型检查。前面的修了文件,后面基于修好的版本检查,不会互相干扰。如果 black 或 ruff 自动改动了文件,commit 会被拦住,告诉你"文件被修改了,请重新 stage 再提交"——这也是好事,确保提交的都是格式化后的代码。
平时正常 commit 就行,hook 自动跑。想跳过某次:git commit --no-verify
想手动跑一遍看效果:
pre-commit run --all-files # 跑所有文件
pre-commit run ruff --all-files # 只跑 ruff
实际跑的时候有个坑:首次运行要下载工具环境。pre-commit 每个 hook 运行在隔离的 venv 里,第一次跑要下载 black、ruff、mypy 以及它们的依赖。mypy 的 pre-commit 环境还要装 torch+torchvision,这两个包加起来快 1G,下载超时是常事。有两种应对方式:
- 环境装好了就别删——pre-commit 下载的 venv 缓存到
~/.cache/pre-commit/,第二次跑不需要重新下 - mypy 单独跑命令行——
mypy src tests直接用项目现有的 .venv,不需要在隔离环境里再装一套 torch
pyproject.toml:别再弄五六个配置文件了
以前要管好一个 Python 项目可能需要:setup.py、setup.cfg、.flake8、mypy.ini、pytest.ini。PEP 621(Python 增强提案 621,把包元数据统一写进 pyproject.toml)之后统一进 pyproject.toml:
[build-system]
requires = ["setuptools>=64"]
build-backend = "setuptools.backends._legacy:_Backend"
[project]
name = "mnist-engineering"
requires-python = ">=3.10"
dependencies = ["torch", "torchvision"]
[tool.mypy]
strict = true
ignore_missing_imports = true
[tool.ruff]
line-length = 100
[tool.pytest.ini_options]
addopts = "-v --tb=short"
testpaths = ["tests"]
pythonpath = ["src"]
一个文件,四个 checkpoint:mypy 配置、ruff 配置、pytest 配置、包构建配置。用 pip install -e . 就能装成可编辑包。
验证结果
跑四条命令验证改造是否成功:
mypy src tests # 类型检查通过
pytest # 11 passed in 3.55s
ruff check src tests # lint 通过
python -m mnist_engineering.main # Epoch 5/5 | Accuracy: 99.10%
四个全绿,验收通过。
和 Android 熟悉的那套对上
| Python 工程化 | Android 工程化 |
|---|---|
| mypy | Kotlin 编译时类型检查 |
| pytest | JUnit |
| pre-commit | git hook + lint |
| pyproject.toml | build.gradle.kts |
| logging | Timber |
| ruff + black | ktlint + ktfmt |
本质是一样的,差别只在于 Python 的动态类型让这些约束都是可选的——你可以不写类型标注、不写测试、不配 pre-commit,代码照样能跑。但要不要对自己负责,那是另一回事。
总结
172 行的入门脚本 → 6 个模块的 src 包 + 11 个 pytest + mypy strict + pre-commit 门禁 + 统一配置。每项都是常规操作,合在一起就是工程化的底线。这套底线会直接用到第 9 天的端到端语音指令识别项目里。
完整的代码在 day08-H-python-engineering/ 目录下,有需要的可以直接当模板用。
系列文章
本系列「车载端侧 AI 工程化从零上手」共 10 篇,建议按序阅读:
- 零基础用 PyTorch 识别手写数字:MNIST 实战入门
- 把 PyTorch 模型变成 ONNX:导出、验证和可视化一次学会
- MLflow 实验追踪:从入门到上手
- 用 MNN 在 CPU 上跑 AI 推理:阿里端侧推理框架上手记
- 小白用 TensorRT 给模型加速:ONNX 转 Engine 踩坑实录
- 零基础上手:高通 SNPE 模型转换实战(从 ONNX 到 DLC)
- 零基础搞懂 SNPE 模型量化:INT8 精度损失 = 0% 的秘密
- Python 工程化重构:从 print 脚本到 pytest 项目
- 车载语音指令识别:从 44% 到 96% 的调优之路
- 车载语音指令识别:SNPE 转换与真机部署实测记录
本文链接:Python 工程化重构:从 print 脚本到 pytest 项目 - https://h89.cn/archives/674.html
版权声明:原创文章 遵循 CC 4.0 BY-SA 版权协议,转载请附上原文链接和本声明。