春江暮客

春江暮客的个人学习分享网站

蛋白质机器学习实战:用 Python 构建氨基酸组成基线

2026-09-23 技术
蛋白质机器学习实战:用 Python 构建氨基酸组成基线

nanoBERT 与 VHHBERT 对比中提到,微调编码器之前,可以先跑一个简单基线。这里把它做成可运行的实验:统计氨基酸组成,加上序列长度,再训练一个小型回归模型。

本文走通从序列到留出集预测的完整流程。程序只需 CPU,支持读取 CSV,并保存数据划分,方便后续用同一批记录比较蛋白质向量。内置数据和目标值均为人工生成,仅用于演示流程,不代表真实蛋白质性质的预测效果。

1. 明确基线能衡量什么

对每条序列,计算 20 种标准氨基酸各自所占的比例,再加入序列长度,共得到 21 个特征。例如,AAGC 中丙氨酸的比例为 2 / 4 = 0.5,长度为 4。

这种表示丢弃了残基顺序。AAGCCGAA 会得到完全相同的特征。它的表达能力有限,适合用来判断:更昂贵的序列表征能否为当前任务增加预测价值。

在相同测试记录上比较两个预测器:

预测器 输入 用途
训练集目标中位数 不使用序列信息 建立不依赖序列特征的误差参照
氨基酸组成 + Ridge 20 个比例和长度 检查简单特征能否降低误差

采用中位数策略的 DummyRegressor始终预测训练集目标的中位数。Ridge 回归通过 L2 惩罚拟合线性预测器。这里固定使用 alpha=1.0 演示,没有进行调参。

2. 安装依赖并准备输入

在 macOS 或 Linux 终端运行:

mkdir composition-demo
cd composition-demo
python3 -m venv .venv
source .venv/bin/activate
python -m pip install 'scikit-learn==1.9.1'

示例已在 Python 3.14.7、scikit-learn 1.9.1 环境测试,无需 GPU 或预训练模型文件。

使用自己的数据时,CSV 需要包含以下列:

列名 含义
id 唯一记录编号
sequence 仅使用 20 种标准氨基酸的非空序列
group 预先确定的家族、谱系或其他评价分组
target 含义和单位一致的有限数值测量结果

程序会统一大小写,去除首尾空白;遇到 X 等模糊残基、缺口或序列内部空格时会拒绝输入。运行前应明确这些记录的处理办法,直接删除残基会改变特征。

3. 运行完整基线

将以下代码保存为 composition_baseline.py,也可以下载脚本

import argparse
import csv
import random
from pathlib import Path

import numpy as np
import sklearn
from sklearn.dummy import DummyRegressor
from sklearn.linear_model import Ridge
from sklearn.metrics import mean_absolute_error
from sklearn.model_selection import GroupShuffleSplit
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler

AA = "ACDEFGHIKLMNPQRSTVWY"


def demo_rows():
    rng = random.Random(17)
    rows = []
    for family in range(60):
        length = rng.randrange(60, 101)
        weights = [rng.uniform(1, 15) if a == "A" else 1 for a in AA]
        parent = "".join(rng.choices(AA, weights=weights, k=length))
        for variant in range(3):
            seq = parent[:-3] + "".join(rng.choices(AA, k=3))
            # Artificial target designed to depend on composition and length.
            target = 4 * seq.count("A") / len(seq) + 0.01 * len(seq)
            target += rng.gauss(0, 0.03)
            rows.append(dict(id=f"f{family}_v{variant}", sequence=seq,
                             group=f"f{family}", target=target))
    return rows


def prepare(rows):
    if not rows:
        raise ValueError("No data rows")
    ids, sequences, groups, targets = [], [], [], []
    for row in rows:
        ident, group = str(row["id"]).strip(), str(row["group"]).strip()
        seq = str(row["sequence"]).strip().upper()
        target = float(row["target"])
        if not ident or not group or not seq or set(seq) - set(AA):
            raise ValueError("Each row needs id, group, and a canonical sequence")
        if not np.isfinite(target):
            raise ValueError("Targets must be finite numbers")
        ids.append(ident)
        sequences.append(seq)
        groups.append(group)
        targets.append(target)
    if len(set(ids)) != len(ids):
        raise ValueError("IDs must be unique")
    if len(set(groups)) < 4:
        raise ValueError("Provide at least four groups for this demo workflow")
    X = np.array([[s.count(a) / len(s) for a in AA] + [len(s)]
                  for s in sequences], dtype=float)
    return X, np.array(targets), np.array(groups), ids, sequences


def main():
    parser = argparse.ArgumentParser(description="Grouped composition baseline")
    parser.add_argument("--csv", type=Path, help="id,sequence,group,target columns")
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--split-out", type=Path, default=Path("split.csv"))
    args = parser.parse_args()
    if args.csv:
        with args.csv.open(newline="", encoding="utf-8-sig") as handle:
            reader = csv.DictReader(handle)
            required = {"id", "sequence", "group", "target"}
            if not required <= set(reader.fieldnames or []):
                raise ValueError("Missing CSV columns: id,sequence,group,target")
            rows = list(reader)
    else:
        rows = demo_rows()
    X, y, groups, ids, sequences = prepare(rows)
    splitter = GroupShuffleSplit(n_splits=1, test_size=0.25,
                                 random_state=args.seed)
    train, test = next(splitter.split(X, y, groups))
    overlap = set(groups[train]) & set(groups[test])
    if overlap:
        raise ValueError("Group overlap")
    if {sequences[i] for i in train} & {sequences[i] for i in test}:
        raise ValueError("Identical sequences cross the split; revise groups")
    # Exclusive creation protects an existing split record from replacement.
    with args.split_out.open("x", newline="", encoding="utf-8") as handle:
        writer = csv.writer(handle)
        writer.writerow(["id", "group", "split"])
        for label, indices in [("train", train), ("test", test)]:
            writer.writerows((ids[i], groups[i], label) for i in indices)
    print(f"scikit-learn={sklearn.__version__}; synthetic={args.csv is None}")
    print(f"rows: train={len(train)} test={len(test)}; features={X.shape[1]}")
    print(f"groups: train={len(set(groups[train]))} "
          f"test={len(set(groups[test]))}; overlap={len(overlap)}")
    models = {
        "median": DummyRegressor(strategy="median"),
        "composition_ridge": make_pipeline(StandardScaler(), Ridge(alpha=1.0)),
    }
    for name, model in models.items():
        model.fit(X[train], y[train])
        error = mean_absolute_error(y[test], model.predict(X[test]))
        print(f"{name}: MAE={error:.4f}")
    print(f"saved split: {args.split_out}")


if __name__ == "__main__":
    main()

默认生成器构造 60 个虚拟家族,每个家族有三个变体。目标值由丙氨酸比例、序列长度和噪声直接构造,因此适合检查代码是否跑通;这种设计本身就有利于组成模型。

查看帮助并运行演示:

python composition_baseline.py --help
python composition_baseline.py

测试环境中的实际输出:

scikit-learn=1.9.1; synthetic=True
rows: train=135 test=45; features=21
groups: train=45 test=15; overlap=0
median: MAE=0.3388
composition_ridge: MAE=0.0275
saved split: split.csv

程序在当前目录生成 split.csv。它不会覆盖已有文件,再次运行时请换一个输出名称:

python composition_baseline.py --split-out split-repeat.csv
python composition_baseline.py --csv measurements.csv --split-out split-real.csv

第二条命令需要你自己的完整数据集。程序要求至少四组,只是输入检查的下限,不代表四组足以支持可靠研究。

4. 正确理解划分与误差

GroupShuffleSplit将一个组的全部记录放在划分的同一侧。test_size=0.25 留出四分之一的组,组数向上取整;各组大小不同时,测试行数比例可能并非 25%。

代码检查组间交叉,以及完全相同的序列是否跨越训练集与测试集。它不会自动发现相似序列,也无法判断分组是否具有生物学意义。训练前必须先建立合理分组。若目标是预测新抗原,只有谱系划分还不够,见抗体模型评价教程

标准化器放在 pipeline 中,只用训练行拟合参数,测试行沿用这些参数。这遵循 scikit-learn 关于预处理数据泄漏的说明。

mean_absolute_error计算预测值与目标值之差的绝对值,再求平均。数值越小越好,单位与目标一致;本例单位没有实际生物学含义。输出的分数按行等权计算,因此大组的贡献高于小组。

5. 用同一基线比较蛋白质向量

保存 split.csv、输入数据和环境记录:

python -m pip freeze > requirements-tested.txt

下一次实验中,按 id 将向量与记录连接,并读取已经保存的 traintest 分配。检查每个预期编号是否恰好出现一次。不要依赖数组行顺序,也不要在 CSV 重新排序后重新生成划分。

比较组成特征与蛋白质向量时,保持目标、留出记录、评价指标和调参预算一致。如需选择正则化强度,应在训练部分内部做分组验证,不能用最终测试目标选参数。单次划分适合初步检查;正式比较还需要足够多的独立分组,以及符合采样设计的不确定性分析。

常见问题

错误或现象 处理办法
Missing CSV columns 使用准确列名:id,sequence,group,target
标准氨基酸序列错误 检查空序列、缺口、空格和非标准残基,按明确的数据规则处理。
Identical sequences cross the split 训练前重新检查重复记录的分组。
FileExistsError --split-out 指定新路径,保留原划分。
Ridge 不优于中位数 检查目标单位、样本量、分组和组成特征是否包含有效信息,不要为了提高分数更换划分。

下一步实验

先在测量口径一致的性质数据上运行基线,再用保存的划分比较冻结的蛋白质向量。重点看新表征在符合实际用途的评价中,相比中位数和组成基线降低了多少误差。

友情链接

其它