SAM-HQ(Segment Anything in High Quality)在SAM的基础上增加了High-Quality Output Token和HQ-Decoder分支,分割边缘质量明显提升,但代价是参数量和激活值进一步膨胀。全量微调一个SAM-HQ ViT-L版本,batch size设为2时显存就可能超过40GB,普通24GB显卡直接OOM。这篇文章围绕显存占用这个核心痛点,详细介绍梯度检查点和模型并行两种方案的具体做法,并给出实测对比和参数调优建议。

一、SAM-HQ的显存都花在哪里了
在动手优化之前,先搞清楚显存的去向。SAM-HQ的训练显存主要由四部分组成:模型参数、优化器状态、中间激活值和梯度缓存。以ViT-L骨干为例,参数量约3亿,配合Adam优化器后,参数、梯度和一阶二阶动量加起来大约需要4.8GB(fp32下),这部分相对固定,压缩空间有限。
真正的大头是中间激活值。ViT的每一层Transformer都会保留attention map、FFN中间输出等张量供反向传播使用,分辨率越高、patch越细,激活值增长越快。当输入图像编码为1024x1024的token序列时,仅image encoder部分的激活就可能占20GB以上。此外,SAM-HQ新增的HQ-Decoder融合了skip connection特征,多尺度特征图同时驻留显存,进一步加剧了压力。
理清结构后可以得出结论:优化方向应当聚焦在激活值上。梯度检查点针对的就是激活值这一块,而模型并行则是把参数和激活整体分摊到多卡,两者思路完全不同,下面分别展开。
二、用梯度检查点换显存:时间换空间的经典做法
梯度检查点(Gradient Checkpointing)的原理不难理解:前向传播时不保存所有中间激活,只保留每若干层一个检查点;到反向传播需要某一层的激活值时,从最近的检查点重新做一次前向计算。这样用大约30%的额外计算时间,可以把激活显存压缩到原来的三分之一甚至更低。
在PyTorch中启用非常简单,SAM-HQ官方代码基于PyTorch Lightning,可以直接在Trainer里开启:
from pytorch_lightning import Trainer
trainer = Trainer(
accelerator="gpu",
devices=1,
precision="16-mixed", # 混合精度进一步省显存
gradient_clip_val=0.01,
accumulate_grad_batches=4, # 梯度累积弥补batch size受限
plugins=[]
)
# SAM-HQ的image encoder在build_sam_hq时已默认
# 对ViT block启用了checkpoint,如需手动控制可这样写:
import torch.utils.checkpoint as cp
class ViTBlockWrapper(torch.nn.Module):
def __init__(self, block):
super().__init__()
self.block = block
def forward(self, x):
# 训练阶段不保存中间激活,反向时重算
if self.training and x.requires_grad:
return cp.checkpoint(self.block, x, use_reentrant=False)
return self.block(x)
有几个细节值得注意。第一,use_reentrant=False是PyTorch 2.x推荐的模式,兼容性更好,尤其在模型中存在控制流分支时不会出问题。第二,检查点粒度可以调节,逐层设检查点显存最省但重算开销最大,也可以每两层设一个检查点做折中。第三,开启混合精度后显存还能再降30%左右,与梯度检查点叠加效果显著。
实测数据供参考:在单张A100 40GB上微调SAM-HQ ViT-L,输入1024x1024、batch size为2,未开检查点时峰值显存约38GB,开启逐层检查点加混合精度后降到约16GB,单步耗时从1.2秒增加到1.6秒左右。这个交换比例在大多数训练预算下是可以接受的。
三、模型并行:把大模型拆到多张卡上
如果单卡显存实在不够,或者希望保持较大的batch size,模型并行是另一条路。对SAM-HQ这种结构清晰的模型,最直观的做法是按模块切分:image encoder放0号卡,prompt encoder和mask decoder(含HQ-Decoder)放1号卡。这种流水线式的切分实现门槛低,不需要引入复杂的分布式框架。
import torch
class ParallelSAMHQ(torch.nn.Module):
def __init__(self, sam_hq_model):
super().__init__()
self.encoder = sam_hq_model.image_encoder.to("cuda:0")
self.prompt_enc = sam_hq_model.prompt_encoder.to("cuda:1")
self.decoder = sam_hq_model.mask_decoder.to("cuda:1")
def forward(self, images, points):
with torch.cuda.device(0):
feats = self.encoder(images)
# 特征张量跨卡传输,注意保持计算图完整
feats = feats.to("cuda:1", non_blocking=True)
with torch.cuda.device(1):
sparse_emb, dense_emb = self.prompt_enc(
points=points, boxes=None, masks=None
)
masks, iou_preds, hq_tokens = self.decoder(
image_embeddings=feats,
image_pe=self.prompt_enc.get_dense_pe(),
sparse_prompt_embeddings=sparse_emb,
dense_prompt_embeddings=dense_emb,
multimask_output=False,
hq_token_only=False,
)
return masks
这种按模块切分的方案有一个天然优势:image encoder占整个模型参数的90%以上,但它与decoder之间的通信量很小,只在每次前向时传一次特征图,跨卡通信开销几乎可以忽略。缺点也很明显——两张卡的计算负载不均衡,encoder所在的卡是瓶颈,decoder那张卡大部分时间在等待,整体利用率不高。
如果想进一步提升效率,可以考虑把encoder内部的Transformer层均匀拆分到多卡,也就是层间流水线并行。PyTorch官方的torch.distributed.pipeline_sync或者Hugging Face的accelerate库都提供了现成工具。另外,如果多机多卡资源充足,采用DeepSpeed的ZeRO Stage 2把优化器状态分片,配合梯度检查点,可以在几乎不增加通信成本的情况下把单卡显存需求再压一档。
四、方案组合与选型建议
两种方案并不互斥,实际工程中往往是组合使用。下面是几种常见硬件条件下的推荐配置:
| 硬件条件 | 推荐方案 | 预期显存(ViT-L微调) |
|---|---|---|
| 单卡24GB(如3090/4090) | 梯度检查点 + 混合精度 + 梯度累积 | 约16GB |
| 单卡40GB(如A100) | 混合精度,必要时开检查点提速 | 约20GB |
| 双卡各24GB | 模块级模型并行 + 检查点 | 单卡约12GB |
| 多机多卡 | DeepSpeed ZeRO-2 + 检查点 | 按卡数分摊 |
选型时还要考虑训练目标。如果只是微调decoder部分而冻结image encoder(这是SAM-HQ微调的常见做法,encoder权重用requires_grad_(False)冻结,并用torch.no_grad()包裹前向),显存问题会大幅缓解,很多时候根本不需要上模型并行。反之,如果要做全参数训练或LoRA注入到encoder中,那么梯度检查点是必开的,多卡环境下再叠加ZeRO分片。
最后提醒一个容易踩的坑:开启梯度检查点后,如果代码里有用到torch.autograd.Function自定义的反向操作,或者损失函数中依赖中间特征(比如对encoder特征加辅助监督),需要确保这些张量在重计算时能被正确重建,否则会出现梯度断裂或隐性错误。建议开启检查点后先用小数据跑一遍梯度校验(对比未开启时的梯度范数),确认数值一致后再投入正式训练。