You need to agree to share your contact information to access this model

This repository is publicly accessible, but you have to accept the conditions to access its files and content.

Log in or Sign Up to review the conditions and access this model content.

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),这是一个明确的负结果:该模型不能开箱即用地做零样本变异效应预测。

需要说明两点,以免被误读:

  1. 这是 12 亿参数 MLM 模型与 7B 级专用模型之间的规模与训练目标差距。 表中表现好的方法
    (CADD、GPN-MSA、phyloP)都是基于多序列比对或系统发育的专用方法,
    其信号来自跨物种保守性,而非单序列语言建模。
  2. NT / HyenaDNA 同样接近随机(0.5303 / 0.5070),说明这是同类 DNA 语言模型在零样本
    VEP 上的共性局限,而非本模型独有。

若需要 VEP 能力,应在本模型基础上做有监督微调,或改用多序列比对类方法。


八、适用场景

直接可用

  • 提取上下文化的核苷酸嵌入,用于下游分类器。
  • 掩码碱基补全 / 计算机模拟突变(in-silico mutagenesis)研究。
  • 作为有监督基因组学任务的微调初始化权重。

适合适配的下游任务

启动子与增强子预测、剪接位点与剪接连接点预测、转录因子结合位点预测、染色质可及性预测、
基因 / ORF 注释、变异效应预测,以及无需比对(alignment-free)的跨物种序列比较。

不适用场景

  • 临床或诊断决策。 本模型是未经充分验证的研究性产物。
  • 任何未经独立实验验证的医学、农业或生物安全决策。
  • 生成用于合成或表达的新型功能序列,且未经严格的生物安全审查。
  • 未经任务特定微调直接用于蛋白层面或 RNA 结构层面的任务。

风险与偏差

模型会继承 NCBI RefSeq 中固有的偏差:参考基因组过度代表研究充分的模式物种和特定人群,
对非模式物种及采样不足的类群覆盖较弱,同时也包含组装错误和残留污染。
因此,对于 RefSeq 中代表性充分的物种和基因组区域,嵌入与预测会更可靠,
不能假定其在整个生命之树上均匀泛化。在任何新类群上部署前,都应做实证验证。

Downloads last month
-
Safetensors
Model size
1B params
Tensor type
BF16
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support