DNA-HybridMoE-1.2B
一个 12 亿参数的 DNA 基础模型,在 NCBI RefSeq 多物种基因组上以掩码语言建模(MLM)
方式预训练。模型架构融合了 DeepSeek-V3 式的多头潜在注意力(MLA)与稀疏混合专家(MoE)
前馈网络,训练序列长度为 8,192 bp,在 16 张 NVIDIA A100 上训练了 116,000 步。
本仓库提供的是原始预训练权重(base model),用于在下游基因组学任务上继续微调,例如
启动子/增强子预测、剪接位点预测、变异效应预测、基因注释以及序列嵌入提取。
一、模型概要
| 项目 | 值 |
|---|---|
| 模型类型 | 双向(非因果)编码器式 Transformer,MLA + MoE 架构 |
| 训练目标 | 核苷酸序列的掩码语言建模(MLM) |
| 总参数量 | 1,200,695,452(约 12.0 亿) |
| 每 token 激活参数量 | 约 3.82 亿(每个 MoE 层激活 top-2 路由专家 + 1 个共享专家) |
| 层数 | 18 层(其中 5 层稠密 FFN,13 层 MoE) |
| 隐藏维度 | 1024 |
| 注意力机制 | 多头潜在注意力(MLA),8 个头,非因果(双向) |
| 上下文长度 | 8,192 bp |
| 词表大小 | 64(实际仅使用 9 个 token ID) |
| 分词方式 | 字符级(A / C / G / T / N) |
| 数值精度 | BF16 |
| 训练数据 | NCBI RefSeq 多物种基因组 |
| 训练硬件 | 16 × NVIDIA A100 |
| 训练步数 | 116,000 步 |
| 累计训练 token | 约 5,474 亿(547.4 B) |
| Transformers 版本 | 4.57.1 |
二、架构细节
2.1 顶层结构
HybridMoEModel
├── embed_tokens nn.Embedding(64, 1024) # 字符级,padding_idx = 1
├── layers ModuleList,18 个 HybridMoELayer
├── norm HybridMoERMSNorm(1024)
└── rope_emb DeepseekV3RotaryEmbedding
HybridMoeForMaskLM
├── model HybridMoEModel
└── lm_head nn.Linear(1024, 64, bias=False) # 与词嵌入不共享权重
每个 HybridMoELayer 都是前置归一化(pre-norm)残差块:x = x + 自注意力(RMSNorm(x)),x = x + 稠密FFN或MoE(RMSNorm(x))。
2.2 层类型布局
通过 first_k_dense_replace / last_k_dense_replace 调度,模型在输入侧与输出侧保留稠密通路,
把容量集中在中段:
| 层索引 | 层数 | FFN 类型 | 中间维度 |
|---|---|---|---|
| 0、1、2 | 3 | 稠密 HybridMoEFFN |
4096 |
| 3 – 15 | 13 | 稀疏 HybridMoE |
每专家 2048 |
| 16、17 | 2 | 稠密 HybridMoEFFN |
4096 |
2.3 注意力机制(MLA)
| 参数 | 值 |
|---|---|
num_attention_heads |
8 |
num_key_value_heads |
8 |
q_lora_rank |
256 |
kv_lora_rank |
256 |
qk_nope_head_dim |
128 |
qk_rope_head_dim |
128 |
qk_head_dim |
256 |
v_head_dim |
256 |
rope_theta |
50000.0 |
is_causal |
false |
attention_bias |
false |
attention_dropout |
0.0 |
Query 与 Key/Value 都经过低秩瓶颈投影。RoPE 只施加在 Q/K 的 qk_rope_head_dim = 128
这一部分子维度上,其余 128 维不含位置信息。
注意力计算直接调用 flash_attention_forward,因此本模型只有 FlashAttention-2
一种后端,flash-attn 是运行必需依赖。
2.4 混合专家(MoE)
| 参数 | 值 |
|---|---|
n_routed_experts |
12 |
n_shared_experts |
1 |
num_experts_per_tok |
2 |
moe_intermediate_size |
2048 |
n_group |
3 |
topk_group |
2 |
norm_topk_prob |
true |
routed_scaling_factor |
1.0 |
| 路由打分方式 | sigmoid + 可学习的 e_score_correction_bias |
路由采用 DeepSeek-V3 的分组受限、偏置校正的 top-k 方案(noaux_tc):用 sigmoid
(而非 softmax)得到路由分数,为选择加上一个无梯度的可学习偏置e_score_correction_bias,再把 12 个专家分为 3 组、每组取分数最高的 2 个求和,
只在组分数最高的 2 组内取全局 top-2,最后对选中权重重新归一化。
偏置机制使模型能在不引入辅助负载均衡损失的前提下平衡专家负载,
从而避免辅助损失项污染语言建模目标。
最终输出是加权路由专家输出之和,再加上一个对所有 token 无条件生效的共享专家:
y = Σ_{i ∈ top2} w_i · Expert_i(x) + SharedExpert(x)
2.5 参数分解
| 组件 | 参数量 |
|---|---|
词嵌入(embed_tokens) |
65,536 |
输出头(lm_head) |
65,536 |
| 最终归一化 | 1,024 |
| 其他参数(未逐项展开) | 1,048,576 |
| 注意力(× 18 层) | 4,063,744 × 18 = 73,147,392 |
| 稠密 FFN(× 5 层) | 12,582,912 × 5 = 62,914,560 |
| MoE 路由专家(× 13 层 × 12 个) | 6,291,456 × 156 = 981,467,136 |
| MoE 共享专家(× 13 层) | 6,291,456 × 13 = 81,788,928 |
| MoE 路由门控(× 13 层) | 12,300 × 13 = 159,900 |
| 层归一化(× 18 层 × 2) | 2,048 × 18 = 36,864 |
| 合计 | 1,200,695,452 |
每 token 激活参数量约 381,757,596(3.82 亿),占总参数的 **31.8%**。
三、分词器
分词器(HybridMoETokenizer)是字符级的:一个核苷酸对应一个 token,
没有 k-mer 合并、没有 BPE。
| Token | ID | 说明 |
|---|---|---|
<s> |
0 | bos_token |
</s> |
0 | eos_token,与 <s> 共用 ID 0 |
<pad> |
1 | pad_token,同时是词嵌入的 padding_idx |
<mask> |
2 | mask_token,用于 MLM |
A |
10 | |
C |
11 | |
G |
12 | |
T |
13 | |
N |
14 | 同时作为 unk_token |
vocab_size 声明为 64 是为了对齐 kernel / 便于硬件处理,实际只会输出
0、1、2、10、11、12、13、14 这 8 个 ID,其余为词嵌入矩阵中从未使用的填充行。
归一化行为:输入会先转成大写,随后所有不属于 [ACGT] 的字符都被替换为 N(ID 14)。
这意味着 IUPAC 简并碱基(R、Y、S、W、K、M、B、D、H、V)、小写碱基、U(RNA)、gap 字符(-、.)与空白**全部会被折叠成 N**。若需保留简并碱基信息,
必须自行扩展词表并调整词嵌入尺寸。
>>> tokenizer("acgtNryk-")
{'input_ids': [10, 11, 12, 13, 14, 14, 14, 14, 14]}
四、训练细节
模型权重为随机初始化后从零预训练:训练脚本只读取模型配置并据此实例化模型,
不加载任何已有权重。
4.1 数据
- 来源:NCBI RefSeq 参考基因组组装,覆盖多个物种。
- 格式:预处理后的序列以 Parquet 存放,每条记录是一条定长 8,192 bp 的字符编码序列。
- 序列长度:固定 8,192 bp,不做 padding。
- 字符集:只接受
A/C/G/T/N;含其他字符的记录在读取时被跳过。 - 反向互补增强:每条序列有 50% 的概率被替换为其反向互补序列(A↔T、C↔G 并整体翻转),
使模型对 DNA 双链方向保持不变的归纳偏置。
掩码策略
对每条序列按 20% 的比例随机选点(mask_prob = 0.20),被选中的位置再按 BERT 式
80/10/10 规则处理:
| 处理方式 | 占被选中位置的比例 | 说明 |
|---|---|---|
替换为 <mask>(ID 2) |
80% | 主要监督信号 |
| 替换为随机碱基(ID 10–14) | 10% | 迫使模型不依赖 [MASK] 标记本身 |
| 保持原碱基不变 | 10% | 同上 |
未被选中的位置其标签被置为 pad_token_id(ID 1),在损失计算中被忽略,
因此损失只在这 20% 的位置上计算。
4.2 训练目标
掩码语言建模(MLM),注意力为双向,每个被掩码位置都能同时看到上游与下游上下文。
损失为交叉熵,且只在 5 个核苷酸 token(A/C/G/T/N,ID 10–14)上做 softmax:
计算前先把其余词表项(ID 0–9、15–63)的 logits 置为 finfo.min,
避免模型把概率质量分配到从未作为预测目标出现过的词表项上。
4.3 MoE 负载均衡
训练中不引入任何 aux-loss 项,而是每 balance_interval = 10 个优化步调整一次路由门控中的e_score_correction_bias:统计该窗口内每个专家被路由到的 token 数,按target_bias = -log((count + 1) / mean_load) 计算目标偏置(负载越高偏置越负,从而被压低),
再以 gamma = 0.10 缩放、裁剪到 ±bias_clamp = 0.20,最后以 ema_decay = 0.95 做指数滑动平均。
该 buffer 不参与梯度更新,只影响专家的选择,不修改路由权重本身。
| 超参数 | 值 |
|---|---|
gamma |
0.10 |
ema_decay |
0.95 |
bias_clamp |
0.20 |
warmup_steps |
500 |
balance_interval |
10 |
五、使用方法
5.1 环境安装
pip install "transformers>=4.57.1" torch accelerate safetensors
# 必需:HybridMoEAttention 中硬编码调用 FlashAttention-2,缺少 flash-attn 无法前向
pip install flash-attn --no-build-isolation
本模型支持 FlashAttention-2,且它是唯一可用的注意力后端:注意力层直接调用flash_attention_forward,加载时无需显式传 attn_implementation(默认即走 FlashAttention-2)。
5.2 掩码碱基预测
import torch
from transformers import AutoTokenizer, AutoModelForMaskedLM
repo_id = "zhangchao162/HybridDNA"
tokenizer = AutoTokenizer.from_pretrained(repo_id, trust_remote_code=True)
model = AutoModelForMaskedLM.from_pretrained(
repo_id,
trust_remote_code=True,
dtype=torch.bfloat16,
).cuda().eval()
seq = "ACGTACGTACGTACGTACGT" + tokenizer.mask_token + "GGGTTTAAACCCGGGTTTAA"
inputs = tokenizer(seq, return_tensors="pt").to(model.device)
with torch.no_grad():
logits = model(**inputs).logits # 形状 (1, L, 64)
mask_pos = (inputs["input_ids"][0] == tokenizer.mask_token_id).nonzero(as_tuple=True)[0]
probs = logits[0, mask_pos[0]].softmax(dim=-1)
for score, idx in zip(probs.topk(5).values.tolist(), probs.topk(5).indices.tolist()):
print(f"{tokenizer.convert_ids_to_tokens(idx):>6} {score:.4f}")
5.3 提取序列嵌入
AutoModel 映射到编码器本体 HybridMoEModel,它会同时返回最终归一化后的隐状态,
以及一个包含各层隐状态的字典(适合做线性探针):
import torch
from transformers import AutoTokenizer, AutoModel
repo_id = "zhangchao162/HybridDNA"
tokenizer = AutoTokenizer.from_pretrained(repo_id, trust_remote_code=True)
encoder = AutoModel.from_pretrained(
repo_id, trust_remote_code=True, dtype=torch.bfloat16
).cuda().eval()
inputs = tokenizer("ACGT" * 512, return_tensors="pt").to(encoder.device)
with torch.no_grad():
out = encoder(**inputs)
print(out.last_hidden_state.shape) # (1, 2048, 1024),已过最终归一化
print(list(out.all_hidden_states)[:3]) # ['layer.0', 'layer.1', 'layer.2'],未过最终归一化
print(out.all_hidden_states["layer.9"].shape)
all_hidden_states是一个以"layer.{i}"为键的 字典(不是元组),
每个条目是经过最终model.norm之前的块输出。
六、性能说明
- 数值精度:checkpoint 以 BF16 存储。在 Ampere 及更新架构的 GPU 上请以 BF16 加载;
转为 FP32 虽然可行但没有必要,且显存占用大致翻倍。 - 注意力后端:只有 FlashAttention-2 一种后端,
flash-attn是必需依赖。 - MoE 分发:参考实现中的
HybridMoE.moe在 Python 层循环遍历 12 个专家并使用index_add_。这种写法简单且正确,但不是融合的 grouped-GEMM 内核,
吞吐会低于生产级 MoE 内核(如 Megatron / vLLM 的 grouped experts)。
若需大规模服务,建议导出到支持 MoE 的推理运行时。 - KV 缓存:MLA 将每层每 token 的 K/V 压缩为一个 256 维潜向量,远小于缓存 8 个完整
注意力头,对长序列推理以及沿染色体滑窗扫描是明显优势。 - 显存占用:BF16 下权重约 2.4 GB,另需加上激活值。
七、评测结果
本节报告 3 组评测,全部在 116k 最终权重上运行。7.1 与 7.3 的置信区间为
1000 次序列级 cluster bootstrap——同一序列内的碱基高度相关,按位置重采样会严重低估区间宽度。
7.1 掩码碱基预测
留出分片,2000 条序列 × 8192 bp;字符级分词,按 BERT 式 80/10/10 掩码,
只统计被掩码位置上 5 类核苷酸(A/C/G/T/N)的 argmax 准确率。
mask_prob |
被掩码 token | accuracy | 95% CI | macro-F1 | cross-entropy |
|---|---|---|---|---|---|
| 0.05 | 818,498 | 0.5070 | [0.5033, 0.5113] | 0.4926 | 1.1004 |
| 0.10 | 1,640,467 | 0.5026 | [0.4992, 0.5069] | 0.4880 | 1.1075 |
| 0.20(训练设定) | 3,276,091 | 0.4919 | [0.4885, 0.4960] | 0.4764 | 1.1253 |
| 0.30 | 4,916,363 | 0.4799 | [0.4766, 0.4837] | 0.4633 | 1.1466 |
| 0.50 | 8,191,432 | 0.3978 | [0.3949, 0.4009] | 0.3855 | 1.5243 |
准确率随掩码率单调下降,符合预期(掩码越密,可用上下文越少)。模型在训练设定(0.20)
下的准确率为 0.4919,比「永远预测最高频碱基」的常数基线(0.2965)高约 19.5 个百分点,
交叉熵低于四类均匀分布(1.3863)0.26 nat;掩码率降到 0.05 时升至 0.5070,
说明它确实在利用上下文,而不是输出语料的边缘分布。
多物种评测集(另一个独立评测集,3 个物种各 500 条)
mask_prob |
留出分片 accuracy | 多物种集 accuracy |
|---|---|---|
| 0.05 | 0.5070 | 0.5491 |
| 0.10 | 0.5026 | 0.5435 |
| 0.20(训练设定) | 0.4919 | 0.5328 [0.5286, 0.5373] |
| 0.30 | 0.4799 | 0.5194 |
| 0.50 | 0.3978 | 0.3974 |
7.2 Genomic Benchmarks 线性探针
冻结模型,取平均池化后的序列嵌入,训练逻辑回归探针(max_train = 20000,max_len = 1024)。
| 数据集 | 训练 / 测试 | accuracy | MCC | AUROC |
|---|---|---|---|---|
demo_coding_vs_intergenomic_seqs |
20000 / 25000 | 0.8912 | 0.7825 | 0.9557 |
demo_human_or_worm |
20000 / 25000 | 0.9411 | 0.8822 | 0.9859 |
human_nontata_promoters |
20000 / 9034 | 0.8498 | 0.7003 | 0.9208 |
human_enhancers_ensembl |
20000 / 29744 | 0.7391 | 0.4783 | 0.8124 |
human_enhancers_cohn |
20000 / 6948 | 0.7337 | 0.4675 | 0.8104 |
human_ocr_ensembl |
20000 / 34952 | 0.6683 | 0.3374 | 0.7279 |
drosophila_enhancers_stark |
5184 / 1730 | 0.6549 | 0.3098 | 0.7217 |
| 宏平均 | 0.7826 | 0.5654 | 0.8478 |
除 human_nontata_promoters(多数类基线 0.5441)外,其余数据集的多数类基线均为 0.5000
(测试集完全平衡),因此上述准确率可直接与 0.5 比较。
任务间差异很大:跨物种判别(人 vs 线虫)几乎饱和(0.9411),而 OCR / 增强子这类调控区任务
明显更难(0.65–0.74)。这与任务本身的可分性一致,说明模型学到的是通用的序列表示,
而不是针对某一类调控元件的特化特征。
7.3 零样本变异效应预测
在 ClinVar 上做零样本致病性判别:对每个变异取 512 bp 窗口,比较参考碱基与替代碱基的
掩码边缘概率(masked_marginal,并对反向互补取平均),用得分排序计算 AUROC。
| 方法 | n | AUROC |
|---|---|---|
| CADD | 900 | 0.9526 |
| GPN-MSA | 900 | 0.9514 |
| phyloP-241m | 900 | 0.9143 |
| Evo2-7B | 900 | 0.8943 |
| NT-v2 | 900 | 0.5863 |
| NT | 900 | 0.5303 |
| 本模型(零样本) | 1018 | 0.5157 [0.4802, 0.5515] |
| HyenaDNA | 900 | 0.5070 |
本模型的零样本 VEP 表现接近随机(AUROC 0.5157 [0.4802, 0.5515],AUPRC 0.4807,n = 1018,
95% CI 覆盖 0.5),这是一个明确的负结果:该模型不能开箱即用地做零样本变异效应预测。
需要说明两点,以免被误读:
- 这是 12 亿参数 MLM 模型与 7B 级专用模型之间的规模与训练目标差距。 表中表现好的方法
(CADD、GPN-MSA、phyloP)都是基于多序列比对或系统发育的专用方法,
其信号来自跨物种保守性,而非单序列语言建模。 - NT / HyenaDNA 同样接近随机(0.5303 / 0.5070),说明这是同类 DNA 语言模型在零样本
VEP 上的共性局限,而非本模型独有。
若需要 VEP 能力,应在本模型基础上做有监督微调,或改用多序列比对类方法。
八、适用场景
直接可用
- 提取上下文化的核苷酸嵌入,用于下游分类器。
- 掩码碱基补全 / 计算机模拟突变(in-silico mutagenesis)研究。
- 作为有监督基因组学任务的微调初始化权重。
适合适配的下游任务
启动子与增强子预测、剪接位点与剪接连接点预测、转录因子结合位点预测、染色质可及性预测、
基因 / ORF 注释、变异效应预测,以及无需比对(alignment-free)的跨物种序列比较。
不适用场景
- 临床或诊断决策。 本模型是未经充分验证的研究性产物。
- 任何未经独立实验验证的医学、农业或生物安全决策。
- 生成用于合成或表达的新型功能序列,且未经严格的生物安全审查。
- 未经任务特定微调直接用于蛋白层面或 RNA 结构层面的任务。
风险与偏差
模型会继承 NCBI RefSeq 中固有的偏差:参考基因组过度代表研究充分的模式物种和特定人群,
对非模式物种及采样不足的类群覆盖较弱,同时也包含组装错误和残留污染。
因此,对于 RefSeq 中代表性充分的物种和基因组区域,嵌入与预测会更可靠,
不能假定其在整个生命之树上均匀泛化。在任何新类群上部署前,都应做实证验证。
- Downloads last month
- -