
1. 为什么机器学习项目需要继承从复制粘贴到系统抽象先说说我自己的经历。几年前我还在做CV相关的算法工作项目里模型越堆越多——baseline的朴素贝叶斯、特征工程后的XGBoost、后来上线的CNN、再后来又加了Transformer。当时项目代码是我和另一个同事各自用自己习惯的方式写的结果就是我这边有train_cnn.py他那边有train_xgb.py两个文件里80%的内容长得差不多只是中间那个model构造不一样数据加载和评估逻辑却有微妙差异。等到要统一加一个early stopping或者换一套评估指标我要改两个文件改了还可能不一致。更麻烦的是模型存档的格式也不统一有的存pickle有的存pt文件有的把整个object都存了上线的时候接预报接口的人每天来问我要说明书。后来我花了一个周末把这些全推倒重来核心就是用Python继承把整个训练、评估、预测的骨架统一起来。这篇内容就是想把那套整理清楚给正在被代码重复和实验管理折磨的人一个可落地的方案。不论你是在用scikit-learn做表格数据还是用PyTorch跑深度学习思路是通用的。继承这件事在教科书里一般跟在“封装、多态”后面例子永远是Animal、Dog、Cat听着好像懂了真到项目里不知道怎么下手。在机器学习项目里继承的核心价值不是少写几行代码那么表层而是把“流程”和“差异”分离流程是最好用的骨架差异是每个模型自己那点私货。基类负责公共的骨架——数据加载流程、训练循环、验证评估、模型保存加载、日志输出每个子类只需要定义自己真正不一样的东西——网络结构、特征处理方式、损失函数。这就是典型的“模板方法模式”。当你有5个模型、3套评测脚本、2种数据源时没有这层抽象每次实验都像在缝补丁有了这层抽象新增一个模型就是新增一个文件跑实验就是一行命令。我后面会逐步展开这套设计的关键点和实际代码中间穿插不少我踩过坑之后总结的经验尤其是那些不跑一遍根本想不到的坑。2. 基类骨架怎么搭接口设计决定你能省多少事2.1 先定义抽象基类把“必须做”和“允许改”分开我习惯用标准库的abc模块来做基类这样能在类实例化的时候直接拦住那些没实现关键方法的子类而不是等到运行时某个奇怪的地方才崩。我的BaseModel大概长这样from abc import ABC, abstractmethod from typing import Any, Dict, Optional import numpy as np class BaseModel(ABC): def __init__(self, model_name: str, device: str cpu, **kwargs): self.model_name model_name self.device device self.history: Dict[str, list] {train_loss: [], val_loss: [], val_metric: []} self._model None self._setup(**kwargs) def _setup(self, **kwargs): pass abstractmethod def build_model(self): ... abstractmethod def _train_step(self, batch): ... def train(self, train_loader, val_loader, epochs10, lr1e-3): raise NotImplementedError def evaluate(self, loader): raise NotImplementedError def predict(self, x): raise NotImplementedError def save(self, path: str): raise NotImplementedError def load(self, path: str): raise NotImplementedError几个设计要点__init__只负责接收通用参数如model_name、device然后把带不确定性的参数通过**kwargs转给_setup处理。这样做的好处是未来再加一个模型需要num_classes也好、需要hidden_dim也好都不用动基类的构造函数。build_model和_train_step标记为抽象方法意味着子类必须实现。这两个是整个系统里差异最大的地方强制实现它们等于把“你必须告诉我网络长什么样、一个batch怎么算loss”这条规则定死。train方法在基类里直接raise NotImplementedError是允许子类完全重写的。如果你所有模型共用一套训练循环那完全可以放到基类里实现但如果有的模型比如sklearn里的SVM不是基于batch训练的重写一下反而自然。为什么不用普通的duck typing直接约定因为在一个多模型、多人协作的项目里显式声明abstractmethod能在实例化时立刻给出错误信息而不是让你的同事运行五分钟后才在forward里收到一个AttributeError。这个时间差在调试中是非常宝贵的。2.2 关键方法的使用约定签名越稳调用方越舒服接口设计最怕的就是每个模型对外表现不一致。预测接口尤其重要因为在生产中调用预测的是另一拨人他们不关心你内部是神经网络还是树模型他们就要一个predict(x)。我的约定是train接收train_loader和val_loader返回self.history历史指标统一记录到history字典里。evaluate接收loader返回{loss: ..., acc: ..., ...}字典。predict接收单个样本或者一个batch返回预测结果不做任何概率输出和标签转换。save和load接收路径内部负责把模型权重、配置、类别标签映射一起打包。有一个容易翻车的地方是predict输入格式。有人习惯接收原始文本有人习惯接收向量最稳妥的做法是在子类内部做转换让外部接口保持统一。我在基类的docstring里明确写过一句话“所有进入predict的数据必须已经是模型可以吃的格式”这样至少同一个项目内部不会出现“你的predict接收字符串、我的predict接收numpy数组”的混乱局面。还有一个被很多人忽略的细节基类的history字典里字段名要固定。后续画loss曲线、对比实验、写实验报告全部依赖这个字段名。你可以在子类里补充其他指标但不要改掉公共字段的名称。我见过有同事把train_loss改成了loss_train结果画图脚本全部跑错排查了半天才找到是字段名的问题。2.3 模板方法把训练循环放进基类用钩子方法留出扩展点如果训练循环比较统一我建议把它实现到基类里通过钩子方法hook来扩展。比如这样class BaseModel(ABC): def train(self, train_loader, val_loader, epochs10, lr1e-3): optimizer self._create_optimizer(lr) for epoch in range(epochs): epoch_loss 0.0 self._on_epoch_start(epoch) for batch in train_loader: batch_loss self._train_step(batch, optimizer) epoch_loss batch_loss avg_loss epoch_loss / len(train_loader) val_metrics self.evaluate(val_loader) self._log_epoch(epoch, avg_loss, val_metrics) self._on_epoch_end(epoch, avg_loss, val_metrics) return self.history def _create_optimizer(self, lr): return torch.optim.Adam(self._model.parameters(), lrlr) def _on_epoch_start(self, epoch): pass def _on_epoch_end(self, epoch, avg_loss, val_metrics): pass def _log_epoch(self, epoch, avg_loss, val_metrics): self.history[train_loss].append(avg_loss) self.history[val_loss].append(val_metrics.get(loss, 0)) self.history[val_metric].append(val_metrics.get(acc, 0)) print(fEpoch {epoch 1}: loss{avg_loss:.4f}, val{val_metrics})这里的_on_epoch_start和_on_epoch_end就是钩子。子类如果想在每个epoch开始时调整学习率、或者在每个epoch结束后保存当前最优模型只需要重写对应方法即可。这就是印刷电路板和插槽的关系公共流程是电路板钩子就是插槽你想加什么模块插上去就行完全不用改板子本身。但这里也要清醒一点不是所有模型都适合模板方法。如果你做的是强化学习、或者需要分阶段训练的GAN训练逻辑差异太大硬把循环塞进基类反而别扭。遇到这种情况我的建议是让train在子类里完全重写基类只保留统一的save、load和evaluate这些“外围能力”。一个系统里允许不同风格的模型存在只要对外接口稳定即可。3. 从数据到模型继承体系如何贯通整个项目3.1 模型层的继承CNN和XGBoost用同一套对外接口接下来看一个实际例子。假设我的机器学习项目里既有PyTorch实现的CNN分类器也有scikit-learn实现的SVM基线我想让它们都能直接进同一个训练、存盘、预测的流程。先定义一个CNN子类import torch import torch.nn as nn class SimpleCNN(BaseModel): def _setup(self, num_classes10, hidden_dim128): self.num_classes num_classes self.hidden_dim hidden_dim def build_model(self): self._model nn.Sequential( nn.Conv2d(3, 32, kernel_size3, stride1, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, stride1, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Flatten(), nn.Linear(64 * 8 * 8, self.hidden_dim), nn.ReLU(), nn.Linear(self.hidden_dim, self.num_classes) ) return self._model def _train_step(self, batch, optimizer): x, y batch x, y x.to(self.device), y.to(self.device) logits self._model(x) loss nn.functional.cross_entropy(logits, y) optimizer.zero_grad() loss.backward() optimizer.step() return loss.item() def evaluate(self, loader): self._model.eval() total_loss, total_correct, total_num 0.0, 0, 0 with torch.no_grad(): for x, y in loader: x, y x.to(self.device), y.to(self.device) logits self._model(x) total_loss nn.functional.cross_entropy(logits, y, reductionsum).item() total_correct (logits.argmax(dim1) y).sum().item() total_num y.size(0) return {loss: total_loss / total_num, acc: total_correct / total_num}再来一个SVM子类from sklearn.svm import SVC class SVMBaseline(BaseModel): def _setup(self, kernelrbf, C1.0): self.kernel kernel self.C C def build_model(self): self._model SVC(kernelself.kernel, Cself.C, probabilityTrue) return self._model def train(self, X, y, **kwargs): self._model.fit(X, y) self.history[val_metric].append(self._model.score(X, y)) return self.history def predict(self, x): return self._model.predict(x)这两个子类风格差异很大但因为都继承自BaseModel对外暴露的方法名是一致的。构建模型时都调build_model()预测时都调predict(x)保存时都调save(path)。在模型管理脚本里你可以用完全相同的代码操作这两个模型这就是多态在机器学习项目里的价值。实际项目中我最常用到这个能力的地方是模型对比。实验脚本写一个for循环遍历所有模型实例每个模型走同一套train → evaluate → save → load → predict流程最后产出一张对比表。这个流程在没有继承的时候实现不了因为每个模型的函数名和调用方式都不一样。3.2 数据层的继承把数据集预处理也纳入统一框架模型只是机器学习项目的一半数据预处理往往更占时间。数据集的继承同样值得做。一个常见的模式是定义一个BaseDataset统一数据文件读取、缓存、索引这些通用逻辑子类只负责实现“给定index返回样本和标签”这个核心操作。from torch.utils.data import Dataset import os class BaseDataset(Dataset): def __init__(self, root_dir, transformNone): self.root_dir root_dir self.transform transform self.samples [] self.labels [] self._load_data() abstractmethod def _load_data(self): pass def __len__(self): return len(self.samples) def __getitem__(self, idx): sample, label self.samples[idx], self.labels[idx] if self.transform: sample self.transform(sample) return sample, label写一个猫狗分类的数据集子类class CatDogDataset(BaseDataset): def _load_data(self): for fname in os.listdir(self.root_dir): if fname.startswith(cat): self.samples.append(os.path.join(self.root_dir, fname)) self.labels.append(0) elif fname.startswith(dog): self.samples.append(os.path.join(self.root_dir, fname)) self.labels.append(1)这样设计的好处是数据增强、归一化、类别均衡这些通用操作可以在基类统一加不用在每个数据集子类里重复实现。数据层和模型层都用继承形成一个前后呼应的体系。3.3 配置与参数的继承不是只有模型需要继承说到扩展还有一个方向值得提——配置类。机器学习实验的参数量特别大全放着不现实全写进代码也不合适。我的做法是做一个BaseConfig基类把模型类型、数据路径、学习率、batch size、epochs等公共配置放在基类每个模型自己的特殊参数放在继承的子类配置里。class BaseConfig: model_name base data_dir data/raw output_dir outputs device cuda epochs 30 batch_size 64 lr 1e-3 class CNNConfig(BaseConfig): model_name simple_cnn num_classes 10 hidden_dim 256 kernel_size 3继承配置类有一个容易被忽略的好处当你用Conifg参数化跑实验时不同模型的配置文件天然就是继承关系公共参数统一调整模型专属参数各改各的。在实验记录和追踪上这个结构会让你很舒服。我在实际项目里因为有个阶段的配置全部写在字典里改一个公共参数要grep全部脚本后来改成继承式配置这一块才清静下来。4. 加一个新模型要多快继承的低成本扩展实践4.1 新模型接入的典型三步用继承组织之后给项目加一个新模型通常只需要三步写一个继承BaseModel的类、实现build_model和_train_step、注册进模型工厂。以加一个ResNet为例import torchvision.models as models class ResNetClassifier(BaseModel): def _setup(self, num_classes10, pretrainedTrue): self.num_classes num_classes self.pretrained pretrained def build_model(self): self._model models.resnet18(pretrainedself.pretrained) self._model.fc nn.Linear(self._model.fc.in_features, self.num_classes) return self._model def _train_step(self, batch, optimizer): x, y batch x, y x.to(self.device), y.to(self.device) loss nn.functional.cross_entropy(self._model(x), y) optimizer.zero_grad() loss.backward() optimizer.step() return loss.item()然后注册到模型工厂class ModelRegistry: _models {} classmethod def register(cls, name): def wrapper(klass): cls._models[name] klass return klass return wrapper classmethod def create(cls, name, **kwargs): return cls._models[name](**kwargs)以后跑实验就是model ModelRegistry.create(resnet, devicecuda) model.build_model() model.train(train_loader, val_loader, epochs30)整个过程中实验脚本、日志逻辑、模型存档方式全部复用你真正要写的核心代码就是网络结构和损失计算那几十行。我实测下来一个新模型从写好到能跑实验跑对比半小时足够瓶颈通常在数据格式对齐上而不是代码结构。4.2 预训练模型复用继承公共能力的典型场景再用热词里的“机器学习 认识猫 标签”来举个例子。假设你接了一个猫狗识别任务想快速验证几个预训练模型的效果。用上面的继承体系VGG、MobileNet、EfficientNet这些模型之间的差异也只是一行models.xxx()的区别class MobileNetClassifier(BaseModel): def _setup(self, num_classes2, pretrainedTrue): self.num_classes num_classes self.pretrained pretrained def build_model(self): self._model models.mobilenet_v2(pretrainedself.pretrained) self._model.classifier[1] nn.Linear(self._model.classifier[1].in_features, self.num_classes) return self._model这样你把ImageNet上预训练好的特征提取部分全部继承下来只需要修改最后分类头就能在猫狗数据集上快速做迁移学习。核心的冻层、微调策略放在基类的_on_epoch_start或者训练循环里统一处理。这类“公共能力”的复用比单纯复制粘贴代码价值要大得多。4.3 实验对比和模型存档的顺带优化模型多了以后还有一个问题怎么统一存档才能省心我在基类的save方法里统一做了这样几件事保存model_name、保存模型权重、保存类别映射。有了这些信息加载模型时就能自动恢复元信息。很多机器学习项目模型到上线阶段会出问题有一半是因为存档只有权重没有类别映射。def save(self, path: str): if self._model is None: raise ValueError(model is not built yet) checkpoint { model_name: self.model_name, state_dict: self._model.state_dict() if hasattr(self._model, state_dict) else self._model, classes: self._class_names if hasattr(self, _class_names) else None, } torch.save(checkpoint, path) def load(self, path: str): checkpoint torch.load(path, map_locationself.device) if self._model is None: self.build_model() if state_dict in checkpoint: self._model.load_state_dict(checkpoint[state_dict]) else: self._model checkpoint这个统一的存档接口影响非常大因为后续做模型版本对比、模型部署、做灰度切换全部依赖这一个方法就够了。子类不需要知道存盘格式细节只管自己的模型逻辑。我在实际项目中就用这个save接口把几十个实验模型统一存档然后写了一个回溯脚本可以一键对比任何两个模型的指标这在之前遍地散落的pkl和pt文件时代是不可想象的。5. 继承的坑与排查多继承、状态保存与动态加载5.1 常见问题速查表继承不是银弹踩坑的时候也有。我整理了一份速查表基本都是这几年真正遇到的问题问题现象根本原因解决方案报错Cant instantiate abstract class子类漏实现了抽象方法检查子类方法名是否与基类抽象方法完全一致包括方法名拼写调super().__init__()报错忘记在子类的__init__里调父类构造子类自定义构造时必须显式调用super().__init__(...)模型参数全部随机加载权重后效果不对build_model和load的顺序问题先build_model再load_state_dict顺序反了等于白load子类属性访问出错X object has no attribute y子类_setup里没给基类需要的属性赋值_setup中完成所有依赖属性的初始化修改基类后其他模型报错基类行为变更影响所有子类大改动前先写单元测试或者尽量通过新增钩子方法扩展双下划线私有变量在子类中无法访问Python名称改写机制跨类访问不要用双下划线用单下划线即可5.2 双下划线这个坑值得单独拿出来说很多人写父类时习惯用__private_var认为这样封装更彻底。但在继承体系里双下划线的行为可能会出乎意料。Python会把__var改写为_ClassName__var所以你在子类中写self.__var访问的实际上是另一个东西。这会导致极其隐蔽的bug——父类存一个__var子类再存一个self.__var两者互不相干代码里看起来一模一样实际各管各的。我的建议是在需要被继承的类里统一用单下划线_var表示受保护属性不在类外部直接访问用双下划线只在你确定这个属性绝对不允许子类覆盖时才用而且最好在注释里写清楚原因。这个经验是我在一个多级继承的项目里踩雷踩出来的排查了两个小时最后发现是__model这个名字被父类和子类各自改写了。5.3 动态加载模型时继承关系如何保持再来聊一个部署场景的问题。当你把模型序列化保存下来在另一个环境里加载时需要确保类定义可用。如果直接对模型实例做pickle.dump序列化的是整个对象包括它的类引用这要求加载环境中import路径完全一致。一旦改了目录结构或者模块名加载就崩。比较稳妥的方案是只保存权重和配置而不是保存对象。我上面的save方法只保存state_dict和model_name加载时通过ModelRegistry按名称实例化正确的子类再灌入权重。这样迁移环境时只需要导入模型类定义不需要保证对象序列化的兼容性。我见过不止一个团队上线前模型文件用的是pickle保存整个sklearn pipeline导致换一台机器就报模块找不到最后只能在新环境里重新训练风险非常大。关于动态加载还有一个细节值得提如果你用__file__或者相对路径在基类里加载配置文件在继承体系下要注意路径解析的基准。基类文件在models/base.py子类在models/cnn.py那么os.path.dirname(__file__)解析出来的目录可能不是你想象的目录。最稳妥的办法是把所有路径都定义在配置对象里避免在模型代码内部拼路径。5.4 多继承与Mixin谨慎但有用机器学习项目里偶尔会遇到多继承的需求比如一个模型既要BaseModel的训练能力又要LoggingMixin的日志能力还要DistributedMixin的分布式支持。Python的MRO方法解析顺序能处理这种情况但容易把人绕晕。我的经验是多继承用可以但只在“补充能力”的场景下用不要用多继承来组织“核心逻辑”。核心逻辑应该在单一基类链里跑能力型功能可以通过Mixin以有限的方式混入。Mixin类不要定义自己的__init__不要调用super().__init__只提供额外方法这样能避免初始化顺序带来的麻烦。class LoggingMixin: def log_metrics(self, metrics: dict): print(f[INFO] {self.model_name}: {metrics}) class DistributedMixin: def to_distributed(self): self._model torch.nn.DataParallel(self._model) return self用的时候class DistributedResNet(BaseModel, DistributedMixin, LoggingMixin): ...这种方式在管理大量模型时很有用。我把日志、可视化、断点续训这些能力都做成了Mixin模型类按需组合不会出现“所有模型都被迫拥有但大部分用不到”的冗余功能。5.5 关于继承深度的教训最后说一个可能被忽视的问题继承层级不要太深。我见过一张三层以上的继承图每一个类叠加几个方法到最后想要搞清楚一个调用实际指向哪个方法要一层层翻源码。在机器学习项目里代码的可读性和可调试性比精简更重要。我的原则是继承层级一般不超过两层基类跟具体模型之间最多隔一个中间抽象类例如BaseModel→ImageClassificationModel→ResNetClassifier。如果超过三层我会优先考虑用组合或者混入Mixin来重构而不是继续往下加子类。我自己在重构早期项目的时候就把一个五层继承的模型体系压回了三层代码总量反而变少了因为很多中间层只是在贴标签并没有提供真正有价值的内容。面向对象设计里的“组合优于继承”这句老话在机器学习项目里同样成立继承解决“是一个”的关系组合解决“有一个”的关系。模型和数据的关系、模型和日志的关系本质上是“有一个”用组合更自然模型的训练接口统一这才是“是一个”的关系用继承才合理。6. 最后再分享一个我常用的调试小技巧我在用继承改模型代码时最常做的调试动作是给BaseModel加一个summary()方法def summary(self): print(fModel: {self.model_name}) print(fDevice: {self.device}) print(fClass: {self.__class__.__name__}) print(fTrainable params: {sum(p.numel() for p in self._model.parameters() if p.requires_grad)})这个方法虽然没有太复杂但在调试基类和子类交互时特别管用。每次实例化一个模型第一件事就是调summary()检查类名、设备、参数量是否符合预期能很快发现很多初始化阶段的错误。比如基类_setup里某步忘赋值了_model还是None调用summary()就会立刻报错而不是等到训练循环走一步才崩。这套继承体系的整理本质上就是把机器学习项目中“会变的东西”和“不变的东西”分开。模型结构、损失函数、数据增强这些是会变的训练循环、评估接口、存档格式、日志记录这些基本是不变的。用继承把不变的部分沉淀到基类让变的部分以子类和钩子的形式自由生长项目才能越做越不累。你不需要一上来就设计得多完美从一个混乱的项目里先抽象出一个最小可用基类然后随着新模型不断接入慢慢调整骨架这个过程本身就是值得的。