11.代码模块化拆分
第 11 章 — 代码模块化拆分
第 10 章给单文件
train.py加了混合精度(AMP),文件已经长到"改一个超参数要滚动半天"的地步。本章把全部代码从一个文件拆成多文件工程:配置、数据、模型、工具、训练循环、入口各归各的目录。这是整套教程的转折点——前 10 章学的是"怎么把训练跑起来",从本章开始学的是"怎么让代码在变大之后还能被看懂、被维护"。拆完之后运行方式不变:pip install -r requirements.txt && python train.py。
一、本章要解决的问题
- 之前:从第 1 章到第 10 章,所有代码都写在同一个
train.py里,功能越加越多,文件越来越长,改一个参数要在一个文件里上下滚动。 - 现在:我们要把代码按职责拆成多个文件,让"改数据增强"只碰数据文件、"改网络结构"只碰模型文件、"改超参数"只碰配置文件。
- 判断标准:拆分后
pip install -r requirements.txt && python train.py跑出来的行为与第 10 章完全一致,但任何一个文件单独打开都短、清楚、只干一件事。
二、核心概念速览
这一章不引入新的深度学习概念,引入的是工程概念。下面 6 个词是本章全部"生词",先用生活里的例子看明白,再看代码。
1. 模块化与单一职责
把一个大文件拆成多个小文件,每个文件只负责一类事,就叫模块化;每个模块只干一件事、别的事一律不管,就叫单一职责。可以类比厨房分工:一个人做整桌菜,菜多了一定手忙脚乱;改成洗菜、切菜、掌勺、装盘各司其职,虽然人数变多了,但每个人只记自己那点事,出错好找、效率还高。代码也一样——一个 300 行的 train.py 就像一个人做满汉全席,拆成多个小文件后,哪个环节出问题,直接去那个文件找。
2. 工厂函数 build_model / build_dataloaders
"工厂函数"是一类特殊函数:你告诉它要什么(比如"给我一个模型"),它负责怎么造(读取配置、实例化、返回成品)。类比点餐出菜:你去餐厅只说"来一份红烧肉",至于猪肉切多大、火开多大,那是后厨的事,你不用管。build_model(cfg) 就是那个后厨——你只传配置,它把 SimpleCNN 拼好递给你。好处是:以后想换网络结构,只改工厂函数内部,调用方一行不用动。
3. 依赖注入
一个对象需要用到的东西(数据、日志、优化器……),由外面传进来,而不是自己内部去创建,就叫依赖注入。类比组装电脑:主板不自己生产 CPU、内存、显卡,而是留好插槽,你来把配件插上去。Trainer 就是那块主板——它不知道数据从哪来、日志写到哪去,构造函数里把所有"配件"收进来,它只负责用。好处:想换数据源、换日志方案,换个"插头"就行,Trainer 内部一字不改。
4. Trainer 类的职责封装
把"训练循环 + 验证 + 存档 + 早停 + 指标记录"整段逻辑收进一个类里,对外只暴露 train() 一个接口,就像遥控器:面板上只有电源键,按一下电视该干嘛自己干。调用方(train.py)不需要知道内部是先用 AMP 还是先存 checkpoint,它只需要 trainer.train(resume=args.resume)。这就是封装:接口简单,内部复杂但藏起来。
5. 循环导入风险与依赖方向
如果文件 A 导入了文件 B,同时文件 B 又导入了文件 A,Python 会在导入时报 ImportError: cannot import name ...,这叫循环导入。避免它的办法是定好依赖方向:谁在顶层、谁在底层,方向永远是"顶层 → 底层",底层绝不回头导入顶层。本章的方向是:train.py(入口,顶层)→ engine/trainer.py → utils/* 和 dataset/*、models/*(叶子,最底层),底层模块之间互不导入,循环导入自然不可能发生。
6. Python namespace package
Python 3.3 起,一个目录不带 __init__.py 也能被 import,这类目录叫 namespace package(命名空间包)。只要程序运行时该目录在 sys.path 里(本项目里 train.py 就在项目根目录运行,当前目录天然在搜索路径里),from dataset.datasets import build_dataloaders 就能找到 dataset/datasets.py。所以本项目不加空的 __init__.py——加了也不报错,但既然 3.3+ 不需要,我们就保持目录干净,只放真正的代码文件。
三、解决思路
- 按职责分目录:
config/(配置)、dataset/(数据)、models/(模型)、utils/(通用工具)、engine/(训练引擎),入口train.py留在根目录。 - 定死依赖方向:
train.py是唯一"上帝文件",它 import 所有模块、组装所有组件;engine/trainer.py只 importutils/*;dataset/、models/、utils/之间互不 import。方向永远是单向的,循环导入从根上杜绝。 - 用工厂函数把"构造"藏起来:
build_model(cfg)、build_dataloaders(cfg)接收配置、返回成品。调用方只提需求,不管实现。 - 用依赖注入组装
Trainer:模型、优化器、调度器、scaler、三个 loader、logger、writer 全部在train.py里造好,通过构造函数一次性传给Trainer。Trainer自己不创建任何外部资源,只负责"用"。 - namespace package:Python 3.3+ 的 namespace package 让目录无需
__init__.py即可 import,本项目因此不加空__init__.py——每个目录里只放真正有内容的文件。 - 不做什么:不引入类继承、抽象基类、插件注册、依赖注入框架等重型设计。拆分只是"物理分文件 + 明确接口",概念上还是前面 10 章的东西。
trade-off:模块化牺牲了"打开一个文件就能看到全部流程"的便利(新手视角确实更直观),换来的是"每个文件都小、可读、可测、可单独替换"。对 300 行的小项目拆分有点"杀鸡用牛刀",但真实项目都是几千上万行,现在练好习惯,后面章节扩展功能时会庆幸今天拆了。
四、代码变更
相对第 10 章的单文件 train.py,现在变成 10 个文件:
- train.py # 之前:唯一的文件,几百行塞了所有功能
+ config/config.yaml # 超参数(原来硬编码在 train.py 里)
+ dataset/datasets.py # 数据变换 + 数据加载
+ models/classifier.py # SimpleCNN 模型定义
+ utils/logger.py # 日志工具
+ utils/checkpoint.py # checkpoint 存档/读取
+ utils/metrics.py # 指标记录
+ utils/early_stopping.py # 早停
+ engine/trainer.py # 训练/验证/存档/早停的完整封装
+ train.py # 入口(瘦身:只做参数解析、配置加载、组件组装)
+ requirements.txt # 依赖清单对应关系(原 train.py 的功能块 → 现在的位置):
- 硬编码的超参数(epochs、lr、batch_size……)→
config/config.yaml SimpleCNN模型类 →models/classifier.py(同时新增统一工厂函数build_model)build_transforms与三个 DataLoader 的构建 →dataset/datasets.py(封装成build_dataloaders)setup_logger日志设置 →utils/logger.py- checkpoint 的保存/加载 →
utils/checkpoint.py MetricsTracker指标记录 →utils/metrics.pyEarlyStopping早停 →utils/early_stopping.pytrain_one_epoch、validate、主训练循环 →engine/trainer.py的Trainer类argparse、配置读取、组件组装、main()→ 瘦身后的train.py
功能上一行没丢、一行没加,纯粹是"搬家 + 封装"。
五、完整代码
按下面的目录结构创建文件。所有目录都不放 __init__.py(namespace package,见上面概念 6),只放真正的文件:
pytorch_learning/
├── train.py
├── requirements.txt
├── config/
│ └── config.yaml
├── dataset/
│ └── datasets.py
├── models/
│ └── classifier.py
├── utils/
│ ├── logger.py
│ ├── checkpoint.py
│ ├── metrics.py
│ └── early_stopping.py
└── engine/
└── trainer.pyconfig/config.yaml:
# CIFAR-10 图像分类实验配置
device: auto
# 上面:auto = 有 GPU 用 GPU,没有用 CPU(等价于第 1 章的 torch.cuda.is_available())
data:
root: ./data
batch_size: 64
num_workers: 2
val_ratio: 0.1
augmentation:
random_crop_padding: 4
horizontal_flip: true
color_jitter: [0.2, 0.2, 0.2, 0.1]
model:
num_classes: 10
train:
epochs: 30
lr: 0.001
lr_scheduler: step
lr_step_size: 15
lr_gamma: 0.1
patience: 7
use_amp: true
checkpoint:
dir: ./checkpoints
log:
dir: ./runs/cifar10dataset/datasets.py:
"""数据集与数据加载:build_transforms 构建变换,build_dataloaders 组装三个 loader。"""
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms
def build_transforms(cfg, train=False):
"""按配置构建训练/验证两套变换。训练集做增强,评估集只归一化。"""
aug = cfg.data.augmentation
transforms_list = []
if train:
if aug.random_crop_padding > 0:
transforms_list.append(
transforms.RandomCrop(32, padding=aug.random_crop_padding))
if aug.horizontal_flip:
transforms_list.append(transforms.RandomHorizontalFlip())
if aug.color_jitter is not None:
transforms_list.append(transforms.ColorJitter(*aug.color_jitter))
transforms_list += [
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
]
return transforms.Compose(transforms_list)
def build_dataloaders(cfg):
"""返回 (train_loader, val_loader, test_loader) 三元组。
工厂函数:调用方只要数据,怎么切分、怎么打乱都在这搞定。"""
train_dataset = datasets.CIFAR10(root=cfg.data.root, train=True, download=True,
transform=build_transforms(cfg, train=True))
test_dataset = datasets.CIFAR10(root=cfg.data.root, train=False, download=True,
transform=build_transforms(cfg, train=False))
# 从训练集里按 val_ratio 划出验证集(random_split 不复制数据,只是索引划分)
val_size = int(len(train_dataset) * cfg.data.val_ratio)
train_dataset, val_dataset = random_split(
train_dataset, [len(train_dataset) - val_size, val_size]
)
# 局部小函数:避免把 batch_size/num_workers 重复写三遍
def make_loader(ds, shuffle):
return DataLoader(ds, batch_size=cfg.data.batch_size, shuffle=shuffle,
num_workers=cfg.data.num_workers)
return (make_loader(train_dataset, shuffle=True),
make_loader(val_dataset, shuffle=False),
make_loader(test_dataset, shuffle=False))models/classifier.py:
"""模型定义:SimpleCNN 与统一的 build_model(cfg) 工厂接口。"""
import torch.nn as nn
class SimpleCNN(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
# 特征提取:3 个卷积块,跟第 1 章的 SimpleCNN 一模一样
self.features = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding=1),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(64, 128, kernel_size=3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
)
# 分类头:全局平均池化 -> 展平 -> 全连接
self.classifier = nn.Sequential(
nn.AdaptiveAvgPool2d(1),
nn.Flatten(),
nn.Linear(128, num_classes),
)
def forward(self, x):
return self.classifier(self.features(x))
def build_model(cfg):
"""按配置构建模型实例(后续可扩展为 cfg.model.name 选择不同结构)。
这是"工厂函数":调用方只要模型,具体构造细节都在这。"""
return SimpleCNN(num_classes=cfg.model.num_classes)utils/logger.py:
"""日志工具:控制台用 rich 渲染,文件用普通格式。"""
import logging
import os
from rich.logging import RichHandler
def setup_logger(log_dir):
"""构建并返回名为 train 的 logger;重复调用会重建,避免叠加重复 handler。
控制台输出带颜色(rich),文件输出纯文本方便事后翻看。"""
os.makedirs(log_dir, exist_ok=True)
logger = logging.getLogger("train")
logger.handlers.clear() # 清空旧 handler,防止热重载时重复打印
logger.setLevel(logging.INFO)
console = RichHandler(rich_tracebacks=True, markup=True)
console.setLevel(logging.INFO)
logger.addHandler(console)
file_handler = logging.FileHandler(os.path.join(log_dir, "train.log"),
encoding="utf-8")
file_handler.setFormatter(logging.Formatter(
"%(asctime)s | %(levelname)s | %(message)s", datefmt="%Y-%m-%d %H:%M:%S"))
logger.addHandler(file_handler)
return loggerutils/checkpoint.py:
"""Checkpoint 工具:保存与加载完整训练状态。"""
import os
import torch
def save_checkpoint(state, path):
"""把训练状态写入磁盘;目录不存在时自动创建。"""
os.makedirs(os.path.dirname(path), exist_ok=True)
torch.save(state, path)
def load_checkpoint(path, device):
"""加载 checkpoint 状态字典。weights_only=False 用于包含 optimizer 等
非 tensor 对象的状态;只应加载自己生成、可信的文件。"""
return torch.load(path, map_location=device, weights_only=False)utils/metrics.py:
"""指标管理:按名字记录每个 epoch 的历史值,支持序列化与恢复。"""
from collections import defaultdict
class MetricsTracker:
def __init__(self):
self.history = defaultdict(list) # 名字 -> 该指标每个 epoch 的值列表
def update(self, **kwargs):
"""例如 update(train_loss=0.5, val_acc=0.75)。"""
for name, value in kwargs.items():
self.history[name].append(value)
def latest(self, name):
"""最近一个 epoch 的指标值。"""
return self.history[name][-1]
def to_dict(self):
return dict(self.history)
def load(self, data):
"""从 checkpoint 恢复历史记录。"""
self.history = defaultdict(list, data)utils/early_stopping.py:
"""早停:验证集指标连续 patience 轮无提升时,置 early_stop=True。"""
class EarlyStopping:
def __init__(self, patience=7, delta=0.0):
self.patience = patience
self.delta = delta
self.best_score = None
self.counter = 0 # 连续"没进步"的轮数
self.early_stop = False
def __call__(self, val_acc):
"""每轮验证后调用一次,传入本轮的 val_acc。"""
score = val_acc
if self.best_score is None:
self.best_score = score # 第一轮直接记为最优
elif score < self.best_score + self.delta:
self.counter += 1 # 没超过历史最优(加 delta 容差),记一次
if self.counter >= self.patience:
self.early_stop = True # 连续 patience 轮没进步,触发早停
else:
self.best_score = score # 进步了,重置计数
self.counter = 0engine/trainer.py:
"""Trainer:把训练/验证/存档/早停/恢复/指标记录完整封装成一个类。
所有依赖(模型、数据、日志……)都从构造函数注入,本模块不自己创建。"""
import os
import torch
import torch.nn as nn
from tqdm import tqdm
from utils.checkpoint import save_checkpoint, load_checkpoint
from utils.metrics import MetricsTracker
from utils.early_stopping import EarlyStopping
class Trainer:
def __init__(self, cfg, model, optimizer, scheduler, scaler,
train_loader, val_loader, test_loader, logger, writer):
# 依赖注入:10 个"配件"全部由外部传进来,这里只负责收下并保存
self.cfg = cfg
self.model = model
self.optimizer = optimizer
self.scheduler = scheduler
self.scaler = scaler
self.train_loader = train_loader
self.val_loader = val_loader
self.test_loader = test_loader
self.logger = logger
self.writer = writer
self.device = next(model.parameters()).device # 模型在哪个设备,训练就在哪个设备
self.criterion = nn.CrossEntropyLoss()
self.metrics = MetricsTracker()
self.early_stopping = EarlyStopping(patience=cfg.train.patience)
self.start_epoch = 1
self.best_acc = 0.0
def _train_one_epoch(self):
self.model.train()
total_loss, correct, total = 0.0, 0, 0
for images, labels in tqdm(self.train_loader, desc="Train", leave=False):
images, labels = images.to(self.device), labels.to(self.device)
self.optimizer.zero_grad()
# AMP:自动混合精度(第 10 章学的),enabled 由 scaler 决定
with torch.amp.autocast(device_type=self.device.type,
enabled=self.scaler.is_enabled()):
outputs = self.model(images)
loss = self.criterion(outputs, labels)
self.scaler.scale(loss).backward()
self.scaler.step(self.optimizer)
self.scaler.update()
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def _validate(self, loader):
self.model.eval()
total_loss, correct, total = 0.0, 0, 0
with torch.no_grad(): # 验证/测试不更新参数,也不需要梯度
for images, labels in tqdm(loader, desc="Val", leave=False):
images, labels = images.to(self.device), labels.to(self.device)
outputs = self.model(images)
loss = self.criterion(outputs, labels)
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def _build_state(self, epoch):
"""把"能恢复训练的一切"打包成一个字典,供 save_checkpoint 使用。"""
return {
"model_state": self.model.state_dict(),
"optimizer_state": self.optimizer.state_dict(),
"scheduler_state": self.scheduler.state_dict() if self.scheduler else None,
"scaler_state": self.scaler.state_dict(),
"epoch": epoch,
"best_acc": self.best_acc,
"metrics": self.metrics.to_dict(),
}
def _save(self, epoch, is_best):
ckpt_dir = self.cfg.checkpoint.dir
state = self._build_state(epoch)
save_checkpoint(state, os.path.join(ckpt_dir, "last.pth"))
if is_best:
save_checkpoint(state, os.path.join(ckpt_dir, "best.pth"))
self.logger.info(f"[Checkpoint] best.pth 已更新(epoch {epoch})")
def _resume(self, resume_path):
"""从 checkpoint 恢复全部状态(模型/优化器/调度器/指标/进度)。"""
ckpt = load_checkpoint(resume_path, self.device)
self.model.load_state_dict(ckpt["model_state"])
self.optimizer.load_state_dict(ckpt["optimizer_state"])
if self.scheduler is not None and ckpt.get("scheduler_state") is not None:
self.scheduler.load_state_dict(ckpt["scheduler_state"])
self.scaler.load_state_dict(ckpt["scaler_state"])
self.start_epoch = ckpt["epoch"] + 1
self.best_acc = ckpt.get("best_acc", 0.0)
self.metrics.load(ckpt.get("metrics", {}))
self.logger.info(
f"[Resume] 从 epoch {ckpt['epoch']} 恢复,历史最优 val_acc={self.best_acc:.4f}")
def train(self, resume=None):
"""对外唯一的接口:start training。内部怎么做,调用方不用管。"""
if resume:
self._resume(resume)
self.logger.info("=" * 60)
self.logger.info("最终生效配置:")
for k, v in self.cfg.items():
self.logger.info(f" {k}: {v}")
self.logger.info("=" * 60)
for epoch in range(self.start_epoch, self.cfg.train.epochs + 1):
train_loss, train_acc = self._train_one_epoch()
val_loss, val_acc = self._validate(self.val_loader)
if self.scheduler is not None:
self.scheduler.step()
current_lr = self.optimizer.param_groups[0]["lr"]
self.metrics.update(train_loss=train_loss, train_acc=train_acc,
val_loss=val_loss, val_acc=val_acc)
self.logger.info(
f"Epoch {epoch:02d}/{self.cfg.train.epochs} | "
f"Train loss {train_loss:.4f} acc {train_acc:.4f} | "
f"Val loss {val_loss:.4f} acc {val_acc:.4f} | "
f"lr {current_lr:.2e}")
self.writer.add_scalar("train/loss", train_loss, epoch)
self.writer.add_scalar("train/acc", train_acc, epoch)
self.writer.add_scalar("val/loss", val_loss, epoch)
self.writer.add_scalar("val/acc", val_acc, epoch)
self.writer.add_scalar("lr", current_lr, epoch)
self._save(epoch, is_best=(val_acc > self.best_acc))
self.best_acc = max(self.best_acc, val_acc)
self.early_stopping(val_acc)
if self.early_stopping.early_stop:
self.logger.info(
f"[EarlyStopping] 连续 {self.cfg.train.patience} 轮无提升,"
f"在 epoch {epoch} 停止")
break
# 训练结束(正常结束或被早停)后用测试集评估一次
test_loss, test_acc = self._validate(self.test_loader)
self.logger.info(f"[Test] loss {test_loss:.4f} acc {test_acc:.4f}")train.py:
"""训练入口:解析参数 -> 加载配置 -> 构建组件 -> 启动训练。只做编排,不做实现。"""
import argparse
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.tensorboard import SummaryWriter
import yaml
# 唯一的"上帝文件":import 所有模块,下面这些 import 的方向决定整个工程的依赖方向
from dataset.datasets import build_dataloaders
from models.classifier import build_model
from engine.trainer import Trainer
from utils.logger import setup_logger
def parse_args():
"""命令行参数:全部有默认值,不传就用配置文件里的。"""
parser = argparse.ArgumentParser(description="CIFAR-10 图像分类训练")
parser.add_argument("--config", type=str, default="config/config.yaml",
help="YAML 配置文件路径")
parser.add_argument("--lr", type=float, default=None, help="覆盖配置文件中的 train.lr")
parser.add_argument("--epochs", type=int, default=None, help="覆盖配置文件中的 train.epochs")
parser.add_argument("--batch-size", type=int, default=None,
help="覆盖配置文件中的 data.batch_size")
parser.add_argument("--resume", type=str, default=None,
help="从指定 checkpoint 恢复训练")
return parser.parse_args()
class Config(dict):
"""让 dict 支持 cfg.train.lr 这样的点号访问,而不是 cfg["train"]["lr"]。"""
def __getattr__(self, key):
try:
return self[key]
except KeyError as e:
raise AttributeError(key) from e
def _to_config(obj):
"""把 yaml 读出来的普通 dict/list 递归转成 Config。"""
if isinstance(obj, dict):
return Config({k: _to_config(v) for k, v in obj.items()})
if isinstance(obj, list):
return [_to_config(v) for v in obj]
return obj
def load_config(path):
with open(path, "r", encoding="utf-8") as f:
return _to_config(yaml.safe_load(f))
def build_scheduler(optimizer, cfg):
"""按配置选择学习率调度器;不支持的名字返回 None(等价于不用调度器)。"""
if cfg.train.lr_scheduler == "step":
return optim.lr_scheduler.StepLR(
optimizer, step_size=cfg.train.lr_step_size, gamma=cfg.train.lr_gamma)
if cfg.train.lr_scheduler == "cosine":
return optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=cfg.train.epochs)
return None
def main():
args = parse_args()
cfg = load_config(args.config)
# 命令行参数优先:传了就覆盖配置文件里的值
if args.lr is not None:
cfg.train.lr = args.lr
if args.epochs is not None:
cfg.train.epochs = args.epochs
if args.batch_size is not None:
cfg.data.batch_size = args.batch_size
logger = setup_logger(cfg.log.dir)
writer = SummaryWriter(log_dir=cfg.log.dir)
# device: auto = 有 GPU 用 GPU,否则用 CPU
device = torch.device("cuda" if cfg.device == "auto" and torch.cuda.is_available()
else cfg.device)
logger.info(f"Using device: {device}")
# 组装阶段:把"配件"一个个造好,最后一次性注入 Trainer
train_loader, val_loader, test_loader = build_dataloaders(cfg)
model = build_model(cfg).to(device)
optimizer = optim.Adam(model.parameters(), lr=cfg.train.lr)
scheduler = build_scheduler(optimizer, cfg)
use_amp = cfg.train.use_amp and device.type == "cuda"
scaler = torch.amp.GradScaler("cuda", enabled=use_amp)
if use_amp:
logger.info("AMP 已启用(FP16 混合精度)")
trainer = Trainer(cfg, model, optimizer, scheduler, scaler,
train_loader, val_loader, test_loader, logger, writer)
trainer.train(resume=args.resume)
writer.close()
if __name__ == "__main__":
main()requirements.txt:
torch>=2.0
torchvision>=0.15
numpy>=1.24
pyyaml>=6.0
tensorboard>=2.13
matplotlib>=3.7
tqdm>=4.66
rich>=13.0
torchsummary>=1.5运行方式与前 10 章一致:
pip install -r requirements.txt
python train.py六、本章小结
- 学到了什么
- 模块化 = 按职责拆文件,每个文件只干一件事(单一职责)。从"看全部"变成"改哪里找哪里"。
- 工厂函数(
build_model、build_dataloaders)把"怎么造"藏起来,调用方只提需求;以后换模型/换数据源,改工厂内部即可。 - 依赖注入:
Trainer需要的组件全部从构造函数传入,自己不创建外部资源,替换组件不改类内部。 Trainer对外只暴露train()一个接口,内部细节(AMP、存档、早停、指标)全部封装起来。- 循环导入靠依赖方向根治:入口
train.py在顶层,底层模块(dataset/、models/、utils/)之间互不 import,方向永远单向。 - Python 3.3+ 的 namespace package 让目录不需要
__init__.py也能被 import,所以本项目刻意不加空的__init__.py。
- 常见坑
- 循环导入:在
engine里 importtrain.py的函数,一运行就ImportError。记住方向:底层永远不回头看顶层。 - 相对导入:
from .utils import ...这种带点的导入只能在"被打包"时用;脚本方式运行时from utils.xxx import ...更稳。 - 拆完文件后
ModuleNotFoundError:多半是train.py不在项目根目录运行,或者拼错目录名/文件名,先ls核对大小写。 - 想当然地补空
__init__.py:Python 3.3+ 不需要;加了也不报错,但既然用 namespace package,保持目录只有真文件更干净。
- 循环导入:在
- 下一章预告:本章拆完工程后,同一次训练跑两遍结果仍然不同——
random_split划分、模型初始化、DataLoader洗牌、数据增强全在随机。第 12 章将固定所有随机源(set_seed),让"同配置 = 同结果"。
七、动手练习
- 手绘依赖图:画出 9 个
.py文件之间的 import 关系,验证方向确实是train.py → engine → utils/dataset/models单向的;再想想如果在utils/logger.py里 importengine.trainer会发生什么,亲手试一下看报错信息。 - 模拟换模型:仿照
build_model的写法,在models/里新增一个build_model_v2(比如把第一个卷积改成 5x5),train.py里换一行调用,确认其他地方一行都不用改。 - 给
--resume传参:先跑 5 个 epoch 中断,再用python train.py --resume ./checkpoints/last.pth恢复,观察日志里的[Resume]信息,并核对start_epoch是从 6 开始的。 - 删掉一个
__init__.py实验:随便建一个空的dataset/__init__.py再删除它,两次都运行python train.py,验证 namespace package 下 import 行为完全一致,从而理解本章为什么特意不加空文件。 - 统计收益:分别给第 10 章的单文件
train.py和本章工程版做"找一个超参数(如patience)并修改它"的小比赛,体会模块化在"改一处"上的体验差异。
