1. 从零手搓AI工程:为什么我不建议你直接调包
很多人一听到“AI工程”这四个字,第一反应就是打开某个云平台,拖几个组件,调几个API,然后跑通一个Demo,就觉得自己已经入门了。我刚开始接触这个方向的时候也是这么想的,直到有一次线上环境出了个诡异的问题——模型推理结果忽好忽坏,日志里没有任何报错,监控指标也一切正常。排查了整整两天,最后发现是特征预处理阶段的一个归一化参数在并发场景下被意外覆盖了。那一刻我才意识到,如果我对底层的数据流转、模型加载、服务编排没有足够的掌控力,我连问题出在哪一层都定位不到。
这就是我决定从零开始搭建一套AI工程体系的原因。不是因为我排斥现成的框架和工具,而是因为我需要真正理解每一个环节在做什么、为什么这么做、出了问题该从哪里下手。这套“从零手搓”的实践,覆盖了从数据接入、特征处理、模型训练、模型服务化到监控告警的完整链路。它适合那些已经会用现成工具跑Demo、但想进一步搞清楚底层原理的开发者,也适合那些在工作中被各种“黑盒”问题折磨、想建立自己技术判断力的工程师。
我所说的“从零”,不是让你去手写矩阵乘法或者重新实现一个深度学习框架,那没有意义。我的定义是:不依赖高度封装的端到端平台,用相对底层的组件和清晰的代码逻辑,把AI系统的每个关键环节显式地搭建出来。你可以用NumPy做数值计算,用Flask或FastAPI做服务暴露,用SQLite或Parquet做数据存储,用简单的轮询或消息队列做任务调度。重点不在于工具多高级,而在于你对整条链路的掌控程度。
在这篇文章里,我会按照我实际搭建的顺序,把每个环节的设计思路、踩过的坑、以及那些“看起来能跑但生产环境一定会出问题”的细节,逐一拆开来讲。我不会给你一个完美的架构图,因为真实项目里从来没有完美的架构,只有不断演进的妥协方案。但我会告诉你,在每一个决策点上,我是怎么权衡的,以及你可能会遇到什么。
2. 数据管道的搭建:从原始文件到可训练样本
2.1 为什么数据加载比模型训练更值得花时间
我见过太多人把80%的精力花在调模型结构上,结果数据管道写得一塌糊涂。训练集和验证集的划分逻辑有漏洞、特征归一化的统计量是在全量数据上算的、类别不平衡的处理方式引入了未来信息——这些问题在Demo阶段不会暴露,因为数据量小、场景简单,但一旦上到真实业务,模型效果会莫名其妙地差,而且你根本找不到原因。
我的做法是,把数据管道当成一个独立的、可测试的模块来对待。它的输入是原始数据文件(CSV、JSON、数据库导出等),输出是经过清洗、转换、划分后的训练样本集。这个模块的代码量往往比模型定义部分还多,但我觉得非常值得。因为数据管道一旦稳定了,后面换模型、调参数都是在这个稳定基础上做增量实验,效率会高很多。
具体来说,我会把数据管道拆成四个阶段:原始数据读取、数据质量检查、特征工程、数据集划分与持久化。每个阶段都有明确的输入输出契约,阶段之间通过中间文件或内存中的DataFrame传递。这样做的好处是,任何一个阶段出问题,我都可以单独调试,而不需要跑完整条链路。
2.2 原始数据读取中的编码与类型陷阱
读取原始数据听起来很简单,但实际操作中坑非常多。最常见的问题是字符编码。我遇到过CSV文件里混了GBK和UTF-8两种编码的文本字段,用pandas默认的读取方式会直接报错或者产生乱码。我的处理方式是,先用二进制模式读取文件的前几KB,用chardet之类的库检测编码,然后显式指定编码格式读取。如果检测结果不确定,就尝试用几种常见编码分别读取,看哪种能成功解析出预期的列数。
另一个坑是数值类型的自动推断。pandas在读取CSV时,如果某一列全是数字,它会自动推断为int或float,但如果这一列里混了一个空值或者一个非数字字符,整列就会变成object类型。这在后续做数值计算时会直接报错。我的习惯是,在读取阶段就显式指定每一列的数据类型,对于不确定的列,先按字符串读取,然后在数据质量检查阶段再做类型转换和异常处理。
import pandas as pd import chardet def detect_encoding(file_path, sample_size=10000): with open(file_path, 'rb') as f: raw = f.read(sample_size) result = chardet.detect(raw) return result['encoding'] def load_raw_data(file_path, dtype_map=None): encoding = detect_encoding(file_path) df = pd.read_csv(file_path, encoding=encoding, dtype=dtype_map) return df注意:不要迷信自动编码检测,对于关键数据文件,最好人工确认一下编码格式。我一般会在读取后打印前几行和每列的数据类型,肉眼扫一遍,确认没有明显的解析错误。
2.3 数据质量检查:那些不检查就一定会后悔的指标
数据质量检查是我在踩过几次坑之后强制加入的环节。具体检查哪些指标,取决于你的业务场景,但有几项是通用的:缺失值比例、唯一值数量、数值列的分布范围、类别列的取值集合。我会把这些检查结果输出成一个简单的报告,每次数据更新后都跑一遍,对比历史报告,看有没有异常变化。
举个例子,有一次我处理一个用户行为数据集,某个类别特征原本只有十几个取值,结果某天数据更新后突然变成了上千个取值。排查后发现是上游系统的一个字段格式变了,把原本的枚举值改成了自由文本。如果没有这个检查,这个变化会直接进入训练流程,导致模型学出一堆无意义的类别,效果大幅下降。
缺失值的处理也需要根据业务含义来定。数值列的缺失,有时候填0是合理的,有时候填均值更合理,有时候应该直接丢弃这条样本。我的做法是,在数据质量检查阶段先统计缺失比例,对于缺失比例超过一定阈值(比如30%)的列,直接标记为不可用;对于缺失比例较低的列,根据业务含义选择填充策略,并在代码里写清楚注释。
2.4 特征工程:在训练之前就把变换逻辑固定下来
特征工程是数据管道里最需要小心的地方,因为这里最容易引入数据泄露。所谓数据泄露,就是你在训练阶段用到了预测阶段拿不到的信息。最典型的错误是,在划分训练集和验证集之前,就对全量数据做了归一化或者标准化。这样验证集的统计量已经影响了训练过程,导致验证结果过于乐观。
我的原则是:任何依赖数据统计量的变换,都必须只在训练集上拟合,然后应用到验证集和测试集。具体操作上,我会把特征变换分成两类:一类是无状态变换,比如取对数、做多项式组合,这类变换对每条样本独立进行,不依赖全局统计量;另一类是有状态变换,比如标准化、归一化、目标编码,这类变换需要先在训练集上计算统计量,然后保存下来,后续对任何新数据都用这个保存的统计量做变换。
from sklearn.preprocessing import StandardScaler import joblib # 只在训练集上拟合 scaler = StandardScaler() scaler.fit(X_train) # 保存变换器,供后续推理使用 joblib.dump(scaler, 'scaler.pkl') # 应用到验证集和测试集 X_train_scaled = scaler.transform(X_train) X_val_scaled = scaler.transform(X_val) X_test_scaled = scaler.transform(X_test)这个流程看起来简单,但在实际项目中,很多人会因为图省事而直接在全量数据上做变换。我建议你把“拟合”和“变换”这两个步骤在代码里显式分开,并且把拟合好的变换器持久化保存。这样在模型上线时,你可以确保推理阶段用的变换逻辑和训练阶段完全一致。
2.5 数据集划分与持久化:别小看随机种子的作用
数据集划分看似只是调用一个train_test_split,但有几个细节需要注意。首先是随机种子,一定要固定,否则每次运行划分结果都不一样,实验无法复现。其次是分层抽样,对于分类问题,要确保训练集和验证集的类别分布一致,尤其是类别不平衡的场景。最后是划分后的数据持久化格式,我一般用Parquet,因为它读取速度快、支持列式存储、能保留数据类型信息。
from sklearn.model_selection import train_test_split X_train, X_temp, y_train, y_temp = train_test_split( X, y, test_size=0.3, random_state=42, stratify=y ) X_val, X_test, y_val, y_test = train_test_split( X_temp, y_temp, test_size=0.5, random_state=42, stratify=y_temp ) # 持久化为Parquet train_df = pd.concat([X_train, y_train], axis=1) train_df.to_parquet('train.parquet', index=False)提示:Parquet文件在跨版本读取时偶尔会有兼容性问题,建议在项目里固定pandas和pyarrow的版本,并在README里写清楚依赖版本。
3. 模型训练环节:把实验管理当成一等公民
3.1 为什么你的实验结果总是无法复现
模型训练这部分,很多人觉得只要把数据丢进去、调几个超参数、看准确率就行了。但我在实际工作中发现,实验无法复现是最让人头疼的问题之一。你上周跑了一个实验,准确率85%,这周想在此基础上再调一调,结果同样的代码跑出来只有82%。你开始怀疑是数据变了、环境变了、还是自己记错了参数。
我的解决方案是,把每次训练都当成一次完整的实验记录。具体来说,我会在训练脚本里自动记录以下信息:代码版本(git commit hash)、数据版本(数据文件的哈希值或版本号)、超参数配置、环境依赖版本、训练开始和结束时间、最终的评估指标。这些信息统一写入一个实验记录文件,可以是JSON、CSV或者简单的SQLite数据库。
import json import hashlib import subprocess from datetime import datetime def get_git_commit(): return subprocess.check_output(['git', 'rev-parse', 'HEAD']).decode('utf-8').strip() def get_file_hash(file_path): hasher = hashlib.md5() with open(file_path, 'rb') as f: hasher.update(f.read()) return hasher.hexdigest() experiment_record = { 'timestamp': datetime.now().isoformat(), 'git_commit': get_git_commit(), 'data_hash': get_file_hash('train.parquet'), 'hyperparameters': {'learning_rate': 0.01, 'max_depth': 6}, 'metrics': {'accuracy': 0.85, 'f1': 0.83} } with open('experiments/exp_001.json', 'w') as f: json.dump(experiment_record, f, indent=2)这样做的好处是,当你发现某个实验结果异常时,可以快速定位到当时的代码和数据状态,判断是哪个环节发生了变化。我甚至会在实验记录里保存模型文件的路径,方便后续做对比分析。
3.2 训练循环中的早停与检查点策略
训练循环本身不复杂,但早停和检查点的策略需要根据实际情况来定。早停的目的是防止过拟合,但早停的耐心值(patience)设多少,需要看你的训练曲线。如果验证集指标波动很大,耐心值设小了会导致过早停止;设大了又浪费计算资源。我的经验是,先跑一次完整的训练,观察验证集指标的波动幅度,然后根据波动幅度来设定耐心值。一般来说,耐心值设为波动周期的2到3倍比较合适。
检查点策略也很重要。我一般会保存两个检查点:最佳验证集指标对应的模型和最后一个epoch的模型。最佳检查点用于后续部署,最后一个检查点用于分析训练过程是否还有提升空间。保存时除了模型参数,还要保存优化器的状态,这样如果训练中断了,可以从检查点恢复继续训练。
best_val_loss = float('inf') patience_counter = 0 patience = 5 for epoch in range(num_epochs): train_loss = train_one_epoch(model, train_loader, optimizer) val_loss = evaluate(model, val_loader) if val_loss < best_val_loss: best_val_loss = val_loss patience_counter = 0 torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'val_loss': val_loss }, 'best_checkpoint.pt') else: patience_counter += 1 if patience_counter >= patience: print(f'Early stopping at epoch {epoch}') break注意:保存检查点时,如果模型很大,频繁保存会占用大量磁盘空间。我一般只保留最近3个检查点,旧的自动删除。另外,检查点文件最好加上时间戳或实验编号,避免不同实验之间互相覆盖。
3.3 超参数搜索:网格搜索之外的实用策略
超参数搜索是模型训练中最耗时的环节之一。网格搜索虽然简单,但计算量随参数数量指数增长,实际项目中往往不可行。我常用的策略是随机搜索加手动精调。先用随机搜索在较大的参数空间里采样几十组配置,快速筛选出表现较好的区域,然后在这个区域附近做小范围的网格搜索或手动调整。
另一个实用技巧是逐步缩小搜索范围。比如先固定其他参数,只调学习率,找到最优学习率后再调正则化系数,依次进行。这种方法虽然不能保证找到全局最优,但在实际项目中往往能在可接受的时间内找到足够好的配置。
import numpy as np from sklearn.model_selection import ParameterSampler param_dist = { 'learning_rate': np.logspace(-4, -1, 100), 'max_depth': [3, 5, 7, 9], 'min_child_weight': [1, 3, 5, 7] } sampler = ParameterSampler(param_dist, n_iter=30, random_state=42) for params in sampler: # 训练并评估 score = train_and_evaluate(params) # 记录结果我还会把每次超参数搜索的结果可视化出来,比如用散点图看学习率和验证集准确率的关系,这样能直观地判断参数的影响趋势,比单纯看数字表格更有感觉。
3.4 模型评估:准确率之外你还需要看什么
准确率是最直观的指标,但在很多场景下它会产生误导。比如类别极度不平衡时,一个把所有样本都预测为多数类的模型也能拿到很高的准确率,但它没有任何实用价值。所以我在评估模型时,会根据业务场景选择多个指标:精确率、召回率、F1分数、AUC-ROC、AUC-PR,以及混淆矩阵。
对于回归问题,除了均方误差和平均绝对误差,我还会看预测值与真实值的散点图,观察模型在哪些区间预测偏差较大。有时候整体误差不大,但在某个关键区间误差很大,这在业务上可能是不可接受的。
from sklearn.metrics import classification_report, confusion_matrix, roc_auc_score print(classification_report(y_val, y_pred)) print(confusion_matrix(y_val, y_pred)) print(f'AUC-ROC: {roc_auc_score(y_val, y_pred_proba):.4f}')我还会做一个错误分析:把预测错误的样本单独拿出来,看看它们有什么共同特征。是某些类别的样本特别容易混淆?还是某些特征区间内的样本预测偏差大?这些分析结果往往能指导下一步的特征工程或数据采集方向。
4. 模型服务化:从训练脚本到可调用的API
4.1 为什么我不推荐直接用Flask裸奔
训练好的模型要产生价值,必须能被其他系统调用。最直接的方式是用Flask写一个简单的HTTP接口,加载模型,接收请求,返回预测结果。但我在生产环境里踩过几次坑之后,发现裸奔的Flask服务有几个致命问题:没有并发控制、没有请求队列、没有超时处理、没有优雅关闭。当请求量稍微大一点,服务就会变得不稳定。
我的做法是,在Flask或FastAPI前面加一层WSGI服务器(比如Gunicorn或Uvicorn),用多个worker进程来处理并发请求。同时,在应用层实现请求队列和超时机制,避免某个慢请求拖垮整个服务。对于模型推理这种计算密集型任务,我还会考虑用单独的进程或线程池来执行,避免阻塞Web服务器的IO处理。
from fastapi import FastAPI from pydantic import BaseModel import joblib import numpy as np app = FastAPI() model = joblib.load('model.pkl') scaler = joblib.load('scaler.pkl') class PredictRequest(BaseModel): features: list @app.post('/predict') def predict(request: PredictRequest): features = np.array(request.features).reshape(1, -1) features_scaled = scaler.transform(features) prediction = model.predict(features_scaled) return {'prediction': prediction.tolist()}启动命令:
uvicorn main:app --host 0.0.0.0 --port 8000 --workers 4提示:worker数量不是越多越好,一般设置为CPU核心数的1到2倍。如果模型推理本身很耗CPU,worker设太多反而会导致频繁的上下文切换,降低整体吞吐量。
4.2 模型加载与版本管理:别让旧模型污染新服务
模型服务化过程中,一个容易被忽视的问题是模型版本管理。当你更新了模型文件,但服务还在用旧模型,或者多个服务实例加载了不同版本的模型,就会出现预测结果不一致的情况。我的做法是,在模型文件命名中加入版本号或时间戳,服务启动时显式指定要加载的模型版本,并在健康检查接口中返回当前加载的模型版本。
import os MODEL_VERSION = os.environ.get('MODEL_VERSION', 'v1') MODEL_PATH = f'models/model_{MODEL_VERSION}.pkl' model = joblib.load(MODEL_PATH) @app.get('/health') def health(): return {'status': 'ok', 'model_version': MODEL_VERSION}这样,当需要更新模型时,只需要重新部署一个新的服务实例,指定新的版本号,然后通过负载均衡逐步切换流量。如果新模型有问题,可以快速回滚到旧版本。
4.3 输入校验:那些你以为不会发生的异常请求
线上服务收到的请求,永远比你想象的更离谱。我遇到过特征数量不对的、特征值超出正常范围的、甚至传了空数组的。如果不做输入校验,这些异常请求会导致模型报错,进而返回500错误,影响用户体验。我的做法是,在API层做严格的输入校验:检查特征数量是否匹配、检查数值范围是否合理、检查是否有缺失值。对于不合法的请求,返回明确的错误信息,而不是让模型去处理。
from fastapi import HTTPException EXPECTED_FEATURE_COUNT = 10 @app.post('/predict') def predict(request: PredictRequest): if len(request.features) != EXPECTED_FEATURE_COUNT: raise HTTPException( status_code=400, detail=f'Expected {EXPECTED_FEATURE_COUNT} features, got {len(request.features)}' ) features = np.array(request.features).reshape(1, -1) if np.isnan(features).any(): raise HTTPException(status_code=400, detail='Features contain NaN values') # 继续处理注意:输入校验的严格程度需要根据业务场景来定。有些场景下,缺失值可以用默认值填充,而不是直接拒绝请求。关键是要在文档里写清楚接口的输入要求,让调用方知道该怎么传参。
4.4 性能优化:批处理与缓存的取舍
当请求量增大时,逐个处理请求的效率很低。一个常见的优化是批处理:把多个请求攒在一起,一次性送给模型推理,然后拆分结果返回。这在GPU推理场景下效果尤其明显,因为GPU的并行计算能力很强,批处理能大幅提升吞吐量。
但批处理也引入了延迟:你需要等待足够多的请求才能组成一个批次。如果请求量本身不大,等待时间可能会超过单个请求的处理时间,反而降低了响应速度。我的做法是,设置一个最大等待时间和最大批次大小,哪个条件先满足就触发推理。这样在请求量大时能充分利用批处理优势,在请求量小时也能保证响应速度。
import asyncio from collections import deque batch_queue = deque() MAX_BATCH_SIZE = 32 MAX_WAIT_TIME = 0.05 # 50ms async def process_batch(): while True: await asyncio.sleep(MAX_WAIT_TIME) if batch_queue: batch = list(batch_queue) batch_queue.clear() # 执行批量推理 results = model.predict(np.array([item['features'] for item in batch])) for item, result in zip(batch, results): item['future'].set_result(result)缓存是另一个优化手段。如果某些请求的特征组合经常重复出现,可以把预测结果缓存起来,下次遇到相同的请求直接返回缓存结果。但缓存需要设置合理的过期策略,避免模型更新后还在返回旧结果。
5. 监控与迭代:上线只是开始
5.1 服务指标监控:延迟、吞吐量与错误率
模型服务上线后,必须持续监控它的运行状态。我关注的三个核心指标是:请求延迟(P50、P95、P99)、吞吐量(QPS)、错误率。延迟反映了用户体验,吞吐量反映了系统容量,错误率反映了服务稳定性。这三个指标中任何一个出现异常,都需要立即排查。
我一般用Prometheus加Grafana来做监控。在服务代码里埋点,记录每个请求的处理时间和状态码,然后通过Prometheus的客户端库暴露指标接口,Grafana负责可视化。这样我可以随时看到服务的实时状态,并设置告警规则,比如P99延迟超过500ms就发通知。
from prometheus_client import Histogram, Counter import time REQUEST_LATENCY = Histogram('request_latency_seconds', 'Request latency') REQUEST_COUNT = Counter('request_count', 'Total request count', ['status']) @app.middleware('http') async def monitor_requests(request, call_next): start_time = time.time() response = await call_next(request) latency = time.time() - start_time REQUEST_LATENCY.observe(latency) REQUEST_COUNT.labels(status=response.status_code).inc() return response提示:监控指标不要只盯着平均值,平均值会掩盖很多问题。P95和P99延迟更能反映真实用户体验,因为少数慢请求往往才是用户抱怨的来源。
5.2 数据漂移检测:模型效果下降的早期信号
模型上线后,效果不会一直保持不变。随着时间推移,输入数据的分布可能会发生变化,导致模型在新数据上的表现下降。这种现象叫做数据漂移。如果不及时发现,模型可能会在不知不觉中变得不可用。
我的做法是,定期统计线上请求的特征分布,和训练时的特征分布做对比。如果某个特征的分布发生了显著变化(比如均值偏移超过一定阈值,或者类别分布差异过大),就触发告警,提醒我可能需要重新训练模型。常用的检测方法包括KL散度、PSI(群体稳定性指标)、KS检验等。
from scipy.stats import ks_2samp def detect_drift(train_feature, online_feature, threshold=0.05): statistic, p_value = ks_2samp(train_feature, online_feature) if p_value < threshold: return True, statistic return False, statistic除了特征分布,我还会监控预测结果的分布。如果模型输出的预测值分布发生了明显变化,比如原本预测为正类的比例是10%,突然变成了30%,这往往意味着输入数据或者业务场景发生了变化,需要进一步排查。
5.3 模型重训练:什么时候该更新模型
模型重训练的时机,不能只靠固定周期来决定。我一般会结合三个信号来判断:数据漂移检测触发、线上评估指标下降、业务规则变化。如果数据漂移检测发现特征分布显著变化,或者线上监控发现模型效果指标持续下降,就需要考虑重新训练。业务规则变化则是指,比如业务方调整了正负样本的定义,或者新增了重要的特征维度。
重训练不是简单地用新数据跑一遍训练脚本。我会先做一次离线评估,用新数据训练一个候选模型,在历史测试集上对比新旧模型的表现。如果新模型在离线指标上明显优于旧模型,再考虑上线。上线时采用灰度发布策略,先让新模型处理一小部分流量,观察一段时间,确认没有问题后再逐步扩大流量比例。
# 离线评估对比 old_model_score = evaluate(old_model, test_data) new_model_score = evaluate(new_model, test_data) if new_model_score > old_model_score + 0.01: # 至少提升1个百分点 print('New model is better, ready for canary deployment') else: print('New model does not show significant improvement')5.4 日志与追踪:出问题时怎么快速定位
线上服务出问题时,最怕的是没有足够的日志来定位原因。我的做法是,在服务的每个关键环节都打日志:请求接收、输入校验、特征变换、模型推理、结果返回。日志里包含请求ID、时间戳、关键参数和耗时信息。这样当某个请求出现异常时,我可以根据请求ID把整条链路的日志串起来,快速定位是哪个环节出了问题。
对于更复杂的系统,我会引入分布式追踪,用OpenTelemetry之类的工具记录请求在各个服务之间的流转路径。这样不仅能定位单个服务内部的问题,还能看到服务之间的调用关系和耗时分布。
import logging import uuid logger = logging.getLogger(__name__) @app.post('/predict') def predict(request: PredictRequest): request_id = str(uuid.uuid4()) logger.info(f'[{request_id}] Received request with {len(request.features)} features') try: features = np.array(request.features).reshape(1, -1) logger.info(f'[{request_id}] Features validated') features_scaled = scaler.transform(features) logger.info(f'[{request_id}] Features scaled') prediction = model.predict(features_scaled) logger.info(f'[{request_id}] Prediction completed: {prediction}') return {'prediction': prediction.tolist(), 'request_id': request_id} except Exception as e: logger.error(f'[{request_id}] Error: {str(e)}') raise注意:日志里不要记录敏感信息,比如用户的原始特征值。如果确实需要记录用于调试,可以对敏感字段做脱敏处理,或者只记录特征的统计量而不是具体值。
6. 一些让我少走了很多弯路的实践习惯
6.1 配置文件与代码分离
我早期写代码时喜欢把参数硬编码在脚本里,改一个参数就要改代码、重新运行。后来我养成了把配置抽离到单独文件的习惯,用YAML或JSON来管理。这样切换实验配置时只需要改配置文件,代码不用动。而且配置文件可以纳入版本管理,方便追溯每次实验用了什么参数。
# config/train_config.yaml data: train_path: 'data/train.parquet' val_path: 'data/val.parquet' feature_columns: ['age', 'income', 'score'] target_column: 'label' model: type: 'xgboost' params: learning_rate: 0.01 max_depth: 6 n_estimators: 200 training: batch_size: 256 epochs: 100 early_stopping_patience: 5import yaml with open('config/train_config.yaml', 'r') as f: config = yaml.safe_load(f) model = create_model(config['model']['type'], config['model']['params'])6.2 单元测试:数据管道和特征变换的守护者
数据管道和特征变换的代码,我强烈建议写单元测试。因为这些代码的逻辑往往比较复杂,而且一旦出错影响面很大。我会针对每个特征变换函数写测试用例,验证输入输出是否符合预期。对于数据质量检查函数,我会构造一些包含缺失值、异常值、类型错误的测试数据,确保检查逻辑能正确识别这些问题。
import pytest import numpy as np def test_normalize_features(): scaler = StandardScaler() train_data = np.array([[1.0], [2.0], [3.0]]) scaler.fit(train_data) test_data = np.array([[4.0]]) result = scaler.transform(test_data) # 验证变换后的均值和标准差 assert abs(result.mean()) < 1e-6 assert abs(result.std() - 1.0) < 1e-6 def test_missing_value_check(): df = pd.DataFrame({'a': [1, 2, None], 'b': [4, 5, 6]}) report = check_missing_values(df) assert report['a']['missing_ratio'] == pytest.approx(1/3) assert report['b']['missing_ratio'] == 0.06.3 版本管理:代码、数据和模型一个都不能少
版本管理不只是代码的git commit。数据和模型也需要版本管理。我的做法是,数据文件用DVC或类似的工具管理,每次数据更新都打上版本标签。模型文件在保存时带上训练数据的版本号和代码的commit hash,这样任何一个模型都能追溯到它的训练来源。
# 数据版本管理示例 dvc add data/train.parquet git add data/train.parquet.dvc git commit -m "Update training data to v2"6.4 文档:写给三个月后的自己
我写文档的原则是:假设三个月后的我已经忘记了所有细节。所以文档里要写清楚每个模块的职责、输入输出格式、关键参数的含义、以及常见的坑。特别是那些“看起来很奇怪但必须这么写”的代码,一定要注释清楚原因,否则三个月后自己都会想把它改掉,然后重新踩一遍坑。
# 注意:这里必须用float64而不是float32 # 因为下游的模型推理库对float32的精度处理有bug # 会导致预测结果在小数点后第6位出现偏差 features = features.astype(np.float64)这套从零搭建的AI工程体系,我前后迭代了大概半年时间。最开始只是想搞清楚模型服务化到底在做什么,后来逐渐扩展到数据管道、实验管理、监控告警。每一步都是遇到问题、解决问题、然后把解决方案固化下来的过程。它不一定适合所有人,但如果你也想建立自己对AI系统的完整掌控力,我觉得这条路值得走一遍。