scGPT关于细胞类型注释的预训练模型微调的代码过程如下:
https://scgpt.readthedocs.io/en/latest/tutorial_annotation.html
整个过程系统全面,但这复杂的代码过程看了让人有点不禁头皮发麻。
PyTorch 深度学习的“绝对内核”,也是 scGPT 微调过程中所有高级特性(混合精度、梯度裁剪、多任务损失)最终服务的核心过程,其实就是一次迭代中“前向传播 → 计算损失 → 反向传播 → 更新参数”这四步。
optimizer.zero_grad()
outputs = model(gene_ids, values, ...)
loss = criterion(outputs["cls_output"], labels)
loss.backward()
optimizer.step()
简单来说,微调过程就是用数据一轮一轮的更新模型参数,每一轮最最核心的步骤就是上面这5行代码。为了让这几行代码能够在每一轮的执行过程中更智能、更顺畅,又添加了一些高级特性如混合精度、梯度裁剪、多任务损失来辅助。
自动混合精度 (AMP) 训练的前端调度器。核心作用是自动为不同的GPU操作选择最优的数值精度,从而在不牺牲模型稳定性的前提下,显著加速训练并降低显存占用。
with torch.cuda.amp.autocast(enabled=hyperparameter_defaults['amp']):
outputs = model(gene_ids, values, ...)
在混合精度训练中,为了加速计算和节省显存,模型的部分前向和反向计算会使用 float16 格式。然而,float16 表示的数值范围远小于 float32,导致那些数值很小的梯度在反向传播时,可能会被直接“四舍五入”为 0。这种现象被称为“梯度下溢”,它会使模型的浅层参数无法得到有效更新,最终导致训练无法收敛。
scale(loss).backward()
scaler.step(optimizer)
scaler.update()
在深度学习中,尤其是在 Transformer 架构 (如 scGPT) 中,反向传播时梯度可能因为连乘效应变得极大 (超过 1e3 或 1e6),导致参数更新步长过大,模型权重瞬间被“撞飞”到数值不稳定区域 (出现 NaN 或 Loss 震荡)。梯度裁剪就是一道“安全阀门”。此步骤一定要在scale(loss).backward()和scaler.step(optimizer)之间执行。
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
在模型训练的不同阶段,需要不同大小的学习率来配合模型的“学习状态”。所以,使用调度器可以解决了固定学习率无法跨越的三大矛盾:早期防止梯度爆炸,防止模型跑偏;中期快速收敛,保证模型能快速向最优解移动;后期精细调参,精细调控找到最优点。
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, 1, gamma=0.9)
scheduler.step()
# 或
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=2, min_lr=1e-6)
scheduler.step(val_loss) # 需要验证损失
有了这些高级特性的守护,一个自主学习的过程便可以顺利地进行。以下是一个极简但完整的 scGPT 微调核心代码框架,只保留细胞类型注释的基础流程:
# ========== 1. 基础配置 ==========
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
epochs = 10
batch_size = 32
lr = 1e-4
weight_decay = 0.01
max_grad_norm = 1.0
use_amp = True # 混合精度
# ========== 2. 模型、优化器、调度器 ==========
model = load_pretrained_model() # 假设已加载预训练 scGPT,并添加了分类头
model.to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=2, min_lr=1e-6)
scaler = GradScaler(enabled=use_amp)
criterion = nn.CrossEntropyLoss()
# ========== 3. 数据加载器(假设已准备好) ==========
train_loader = DataLoader(...) # 返回 gene_ids, values, padding_mask, labels
valid_loader = DataLoader(...)
# ========== 4. 训练和验证函数 ==========
def train_one_epoch(loader):
model.train()
total_loss = 0
for batch in loader:
gene_ids = batch["gene_ids"].to(device)
values = batch["values"].to(device)
padding_mask = batch["padding_mask"].to(device)
labels = batch["labels"].to(device)
optimizer.zero_grad(set_to_none=True)
with autocast(enabled=use_amp):
outputs = model(gene_ids, values, src_key_padding_mask=padding_mask)
loss = criterion(outputs["cls_output"], labels)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)
scaler.step(optimizer)
scaler.update()
total_loss += loss.item()
return total_loss / len(loader)
def evaluate(loader):
model.eval()
total_loss, correct, total = 0, 0, 0
with torch.no_grad():
with autocast(enabled=use_amp):
for batch in loader:
gene_ids = batch["gene_ids"].to(device)
values = batch["values"].to(device)
padding_mask = batch["padding_mask"].to(device)
labels = batch["labels"].to(device)
outputs = model(gene_ids, values, src_key_padding_mask=padding_mask)
logits = outputs["cls_output"]
loss = criterion(logits, labels)
total_loss += loss.item()
preds = logits.argmax(dim=1)
correct += (preds == labels).sum().item()
total += labels.size(0)
return total_loss / len(loader), correct / total
# ========== 5. 主训练循环 ==========
best_val_acc = 0.0
for epoch in range(1, epochs + 1):
train_loss = train_one_epoch(train_loader)
val_loss, val_acc = evaluate(valid_loader)
print(f"Epoch {epoch}: Train Loss={train_loss:.4f}, Val Loss={val_loss:.4f}, Val Acc={val_acc:.4f}")
scheduler.step(val_loss) # ReduceLROnPlateau 需要验证损失
if val_acc > best_val_acc:
best_val_acc = val_acc
torch.save(model.state_dict(), "best_model.pt")
print(" -> Best model saved.")
微调过程中还有一个重要的因素需要考虑——批次效应。关于批次效应scGPT有三个相关的参数:DSBN,DAB,ADV。
简单来说,DSBN 是“结构适应”,DAB 是“特征对齐”,而 ADV 更像是一个底层的、实验性的实现方式。三个参数的核心区别,对比如下:
| 特性 |
DSBN (Domain-Specific BatchNorm) |
DAB (Domain-Adversarial Training) |
ADV (Adversarial Training) |
|---|---|---|---|
| 核心思想 | 结构适应:让模型的不同部分(域)拥有独立的批归一化参数,以适配各自的数据分布。 | 特征对齐:通过对抗训练,让模型学习一种“领域不变”的表征,使判别器无法区分数据来源。 |
对抗训练:与DAB目标一致,通过一个独立的判别器与模型进行对抗,来消除批次效应。 |
| 实现方式 | 修改网络结构,为每个批次学习专属的 BN 参数。 | 在损失函数中加入梯度反转层(GRL),实现端到端的对抗学习。 | 需要手动管理判别器和主模型的交替训练过程,实现更复杂。 |
| 主要优势 | 实现简单,计算开销小,能有效处理批次间分布差异大的情况。 | 端到端训练,与模型融为一体,通常效果更稳定,是官方整合任务的首选。 | 控制灵活,可以分别调整判别器和主模型的学习率等超参数。 |
| 主要劣势 | 无法泛化到训练中未见过的全新批次。 | 训练更复杂,需要调节dab_weight等超参数。 |
训练最不稳定,超参数敏感,且可能与DAB冲突,不推荐常规使用。 |
| 适用场景 | 批次已知且固定,主要关注拟合当前数据。 | 批次整合(Integration) 任务的标准选择,追求跨批次的泛化能力。 | 需要进行精细控制对抗训练过程的实验性场景。 |
指南
主要任务是“批次整合”(Batch Integration)时 → 选择
DAB
这是 scGPT 官方推荐的配置。它的目标是将不同批次的数据“对齐”到一个统一的隐空间中,让相同类型的细胞聚在一起。例如,当你需要合并多个实验的 PBMC 数据时,DAB是首选。只想让模型更好地“适应”已知的多个批次时 → 选择
DSBN
如果你的目标不是将批次完全对齐,而是希望模型在训练时能更好地处理来自不同批次的数据,DSBN是一个简单有效的选择。它通过为每个批次学习独立的 BN 参数,让模型能更好地拟合各个批次内部的数据分布。想进行精细的对抗训练控制时 → 可考虑
ADV(但不推荐)
ADV提供了更底层的对抗训练实现,你可以分别控制判别器和主模型的学习率。但请注意,它的训练更复杂、更不稳定。通常情况下,使用DAB就足够了,无需考虑ADV。
要点
-
DAB和ADV是互斥的:它们是同一种去批次思路的不同实现,绝不能同时开启,否则会导致训练冲突。 -
DSBN可与DAB或ADV联用:DSBN是从网络结构层面解决问题,而DAB/ADV是从损失函数层面,两者可以结合使用,可能带来更好的效果。 - 先去批次,再微调:对于批次效应极强的数据,一个更稳健的策略是先用传统的 Harmony、ComBat 等方法进行预处理,消除明显的批次差异,然后再用 scGPT 进行微调。
总结
-
首选
DAB:对于大多数批次整合任务,这是最标准、最可靠的选择。 -
辅助
DSBN:如果希望模型能更好地适应已知的多个批次,可以考虑在DAB的基础上开启DSBN。 -
避免
ADV:除非你有非常特殊的实验需求,否则不建议使用ADV。
总的来说,DAB 是官方推荐的、用于批次整合任务的标准配置,可以把它作为首选方案。