422

MLOps for 药物发现:实验追踪、版本管理与部署

MLOps 让药物发现的模型可复现、可维护、可协作。这篇给出实验追踪、版本管理与部署的实践方案。

药物发现的模型有几个特点:数据量小、更新频繁、生命周期长、决策后果重。这让 MLOps 实践比在互联网场景中更重要——一个半年前的预测可能需要在今天被追溯与解释。

需要管理的四类资产

资产 工具 关键要求
代码 Git 提交哈希可追溯
数据 DVC / LakeFS / 内容哈希 数据集版本化
实验 MLflow / W&B 参数、指标、产物全记录
模型 模型注册表 版本、血缘、部署状态
环境 容器 + lock 文件 可精确重建

数据版本化是药物发现场景中最容易被忽略也最重要的一环:训练集会随着新实验数据不断增长,如果不做版本管理,「用哪批数据训的这个模型」就无从追溯。

实验追踪

import mlflow
import mlflow.sklearn
import hashlib, json

def train_with_tracking(X, y, params, dataset_df, split_info):
    mlflow.set_experiment("herg_prediction")

    with mlflow.start_run():
        # 1) 记录数据版本 —— 关键
        data_hash = hashlib.sha256(
            dataset_df.to_csv(index=False).encode()).hexdigest()
        mlflow.log_param("dataset_hash", data_hash)
        mlflow.log_param("n_samples", len(dataset_df))
        mlflow.log_param("data_snapshot_date", dataset_df.attrs.get("date"))

        # 2) 记录划分方式 —— 这决定了指标的含义
        mlflow.log_param("split_type", split_info["type"])   # scaffold/random/time
        mlflow.log_param("split_seed", split_info["seed"])

        # 3) 记录超参
        mlflow.log_params(params)

        # 4) 训练
        model = train_model(X, y, params)

        # 5) 记录指标(含多种子的均值与标准差)
        metrics = evaluate(model, X_test, y_test)
        mlflow.log_metrics(metrics)

        # 6) 记录适用域信息
        mlflow.log_dict(
            {"training_chemical_space": describe_chemical_space(dataset_df)},
            "applicability_domain.json")

        # 7) 保存模型与环境
        mlflow.sklearn.log_model(model, "model",
                                 registered_model_name="herg_classifier")
        mlflow.log_artifact("requirements.lock")

        # 8) 记录代码版本
        import subprocess
        commit = subprocess.check_output(
            ["git", "rev-parse", "HEAD"]).decode().strip()
        mlflow.log_param("git_commit", commit)

        return model

数据版本化

# DVC 的基本用法
pip install dvc

dvc init
dvc remote add -d storage s3://bucket/dvc-store

# 跟踪数据集
dvc add data/chembl_egfr_v3.csv
git add data/chembl_egfr_v3.csv.dvc .gitignore
git commit -m "数据集 v3:新增 120 条内部数据"
dvc push

# 恢复某个历史版本
git checkout <commit>
dvc checkout

# 定义数据处理流水线(可复现)
dvc stage add -n clean \
  -d scripts/clean.py -d data/raw.csv \
  -o data/clean.csv \
  python scripts/clean.py

dvc repro          # 只重跑变更影响的部分

模型注册与生命周期

from mlflow.tracking import MlflowClient

client = MlflowClient()

# 注册模型版本
client.create_model_version(
    name="herg_classifier",
    source="runs:/<run_id>/model",
    run_id=run_id,
)

# 阶段管理
client.transition_model_version_stage(
    name="herg_classifier", version=3, stage="Staging")
# Staging → 验证通过 → Production

# 关键:记录每个版本的元信息
client.set_model_version_tag(
    "herg_classifier", 3, "validation_report", "s3://.../report_v3.pdf")
client.set_model_version_tag(
    "herg_classifier", 3, "applicability_domain", "MW 150-600, 常见有机官能团")
client.set_model_version_tag(
    "herg_classifier", 3, "intended_use", "内部化合物优先级排序")

# 生产环境永远引用具体版本,不引用 "latest"
#   → 保证可复现

预测服务与记录

from datetime import datetime
import json, hashlib

class TrackedPredictor:
    """带完整审计记录的预测服务"""

    def __init__(self, model, model_version, ad_checker, log_sink):
        self.model = model
        self.model_version = model_version
        self.ad_checker = ad_checker      # 适用域检查器
        self.log_sink = log_sink

    def predict(self, smiles, requested_by=None):
        ad = self.ad_checker(smiles)
        pred, unc = self.model.predict_with_uncertainty(smiles)

        record = {
            "timestamp": datetime.utcnow().isoformat(),
            "requested_by": requested_by,
            "model_version": self.model_version,
            "input_smiles": smiles,
            "input_hash": hashlib.sha256(smiles.encode()).hexdigest()[:16],
            "prediction": float(pred),
            "uncertainty": float(unc),
            "applicability_domain": ad,
            "needs_review": (ad["confidence"] != "高") or (unc > 0.5),
        }
        self.log_sink.write(record)
        return record

# 每次预测都留痕,且明确标出「需要人工复核」的情况

性能监测与漂移

# 药物发现场景的漂移特点:
#   项目推进 → 化学空间移动 → 新化合物逐渐偏离训练分布
#   这是必然发生的,需要主动监测

def monitor_drift(recent_predictions, training_data, window_days=30):
    """监测输入分布漂移"""
    ad_scores = [p["applicability_domain"]["avg_top_k"]
                 for p in recent_predictions]
    out_of_domain_rate = sum(1 for s in ad_scores if s < 0.4) / len(ad_scores)

    alerts = []
    if out_of_domain_rate > 0.3:
        alerts.append("超过 30% 的预测落在适用域外,建议用新数据重训")

    return {"out_of_domain_rate": out_of_domain_rate, "alerts": alerts}

def monitor_performance(predictions, later_measurements):
    """前瞻性性能监测 —— 最重要的监测"""
    merged = match_predictions_to_results(predictions, later_measurements)
    if len(merged) < 10:
        return {"status": "数据不足"}

    import numpy as np
    err = merged["predicted"] - merged["measured"]
    spearman = merged["predicted"].corr(merged["measured"], method="spearman")

    status = ("正常" if spearman > 0.5 else
              "性能下降" if spearman > 0.3 else "失效,应停用")
    return {
        "n": len(merged), "rmse": float(np.sqrt((err**2).mean())),
        "bias": float(err.mean()), "spearman": float(spearman),
        "status": status,
    }

# 建议:每轮实验数据回流后自动运行,结果进项目报告(见 355)

重训策略

  • 触发条件:新数据积累到一定量、适用域外比例升高、前瞻性能下降、化学空间明显移动;
  • 不要盲目自动重训:药物发现的数据量小,新增少量数据可能让模型变差。每次重训都应该与当前生产版本做对比验证
  • 保留旧版本:新版本表现不如旧版本时能回滚;
  • 重训后重新做适用域定义:训练数据变了,适用域也变了。

药物发现场景的特殊考量

  • 数据量小:几百到几千条,模型对数据变化敏感,重训需谨慎验证;
  • 数据缓慢积累:每轮实验只新增几十条,而非互联网场景的持续大流量;
  • 反馈周期长:预测到实测可能间隔数周到数月,性能监测的滞后必须接受;
  • 决策后果重:一个错误预测可能导致几个月的错误方向,因此人工复核的价值远高于追求全自动;
  • 长期可追溯:项目周期以年计,需要能追溯多年前的预测依据。

关键要点

  • 数据版本化是药物发现 MLOps 中最易忽略也最重要的一环
  • 实验记录必须包含划分方式——它决定了指标的含义;
  • 生产环境引用具体版本号而非 “latest”,保证可复现;
  • 不要盲目自动重训——小数据下新版本可能更差,必须对比验证。

延伸资源