多 GPU 跑 Nemotron-Labs-Diffusion 这类扩散模型时,负载均衡和容错是两件事:前者决定显存是否被均匀使用,后者决定某张卡故障后服务是全部中断还是能降级继续出图。很多部署卡在启动阶段,模型加载后只有主卡有显存占用,或者用多进程硬切导致显存溢出;这类问题通常不是模型不并行,而是并行策略和任务切分方式不匹配。另一个常见现象是单卡报错后整个服务直接退出,缺少“避开故障卡重试”的机制。下面的配置思路适用于 PyTorch 环境下、模型支持手工指定并行维度的情况;需要结合你本地的卡数、单卡显存和框架版本做调整。
核心判断:先确认模型原生支持哪种并行,再配置设备映射让每卡显存尽量接近,之后在请求层加入健康检查和重试,最后用观察命令确认各卡利用率。适用场景是多卡推理、单机多卡或同构多机;操作动作是修改模型加载参数、设置 device_map、增加运行时检测;验证方式是看每卡显存和利用率是否同时波动;风险边界在于模型本身不支持张量并行时,盲目切分会导致加载失败或生成结果错误。
先确认模型支持的并行方式
扩散模型的多卡并行主要有张量并行和数据并行。张量并行是把一个 Transformer 块或卷积层的权重切到多张卡上,适合单模型太大、单卡放不下的场景;数据并行则是每张卡放一个完整模型副本,各自处理不同 batch,适合单卡能放下模型但吞吐不够的场景。Nemotron-Labs-Diffusion 如果基于 Transformer 结构,通常两者都可能支持,但代码实现不一定都暴露了接口。要先读模型的加载库源码或说明,确认它是否实现了张量并行切分逻辑,而不是简单地把模型复制到多卡。
在模型加载时,通常会有一个并行度参数,比如 parallelism 或 tensor_parallel_size。以下是一种通用加载骨架,具体参数名以实际代码为准:
from your_model_library import load_model
model = load_model(
"Nemotron-Labs-Diffusion-checkpoint",
parallelism="tensor", # 或 "data"
tensor_parallel_size=4, # 使用 4 张卡做张量并行
dtype="bfloat16"
)
适用条件:当模型的加载函数明确支持 tensor_parallel_size,或底层调用了 model_parallel 相关模块时,才能用张量并行。如果只是用 nn.DataParallel 包了一层,那不是真正的张量并行,只是数据并行,模型无法跨卡切分显存。判断方法是加载后打印每卡显存占用:如果所有卡都有占用且数值接近,才是张量并行生效;如果只有 0 号卡有占用且其他卡是空的,说明只做了数据并行,需要检查并行参数是否被正确传入。
配置设备映射确保显存均匀分配
显存分配不均最常见的原因是框架默认把所有参数放到主卡,然后用通信把数据广播到其他卡。解决办法是显式指定 device_map。device_map 可以是一个字典,把每一层分配到指定 GPU,也可以写一个自动分配函数。下面是一个字典示例,适用于模型结构为若干连续层的简单情况:
device_map = {
"text_encoder": "cuda:0",
"unet.diffusion_model.encoder": "cuda:1",
"unet.diffusion_model.middle": "cuda:2",
"unet.diffusion_model.decoder": "cuda:3",
"vae": "cuda:0" # 小模块可以分担到已有卡
}
model = load_model(
"Nemotron-Labs-Diffusion-checkpoint",
device_map=device_map
)
更通用的做法是用一个函数让框架按显存占用自动分配:
from transformers import AutoModel
def balanced_device_map(module_sizes, max_memory_per_gpu="16GiB"):
# 按各模块权重大小排序,按顺序分配到当前显存最空的卡
import torch
available = {i: torch.cuda.get_device_properties(i).total_memory for i in range(torch.cuda.device_count())}
map = {}
used = {i: 0 for i in available}
for name, size in module_sizes.items():
target = min(available, key=lambda i: used[i])
map[name] = target
used[target] += size
return map
module_sizes = {
"text_encoder": 3500,
"unet": 12000,
"vae": 1000
}
print(balanced_device_map(module_sizes))
配置后要打印显存分布,直接看加载后的当前占用:
import torch
for i in range(torch.cuda.device_count()):
print(f"GPU {i}: {torch.cuda.memory_allocated(i) / 1024**3:.2f} GiB")
判断标准:各卡已分配显存差距不超过单卡容量的 10% 到 20%,就认为分配基本均匀。如果某一卡明显偏高,就把高占用的模块换成显存占用低的卡,再重新打印。
在请求层加入健康检查与重试
单卡故障时,如果推理程序直接崩溃,所有请求都会失败。更合理的做法是:在请求处理循环中持续检测每张卡的健康状态,遇到故障卡就将其从可用设备列表中移除,然后让当前请求在剩余卡上重试。这里做一个 cuda-runtime 层面的健康检查:
import torch
def healthy_gpus():
good = []
for i in range(torch.cuda.device_count()):
try:
torch.cuda.init()
# 做一次小张量 op,触发 runtime 错误
_ = torch.zeros(1, device=f"cuda:{i}")
torch.cuda.synchronize(i)
good.append(i)
except RuntimeError:
print(f"GPU {i} 检测到异常,已排除")
return good
available_gpus = healthy_gpus()
注意:torch.cuda.init() 只能检测 CUDA runtime 是否能建立上下文,无法识别卡上的 ECC 错误或温度过高导致的降频。要更完整地检测,需要配合 nvidia-smi 查询 Xid 错误或持续报错日志。这里只是一个基础版本。
当检测到某张卡不可用后,请求重试逻辑可以根据可用的卡来重新选择设备执行推理。伪代码如下:
def infer_with_retry(prompt, max_retries=2):
for attempt in range(max_retries + 1):
if not available_gpus:
raise RuntimeError("没有可用 GPU")
device = available_gpus[0] # 或按负载轮询
try:
output = model.sample(prompt, device=device)
return output
except RuntimeError as e:
# 如果错误是“设备端”的,例如 illegal memory access,则剔除该卡
if "CUDA error" in str(e) or "device-side assert" in str(e):
available_gpus.remove(device)
print(f"GPU {device} 被移除,剩余 {available_gpus}")
continue
raise
raise RuntimeError("重试后仍失败")
重试机制要避免无限重试。建议最多尝试 2-3 次,且每次重试前都重新调用健康检查,防止把已经恢复的卡一直排除在外。这个逻辑要放在请求线程里,不要放在模型加载阶段,因为故障可能发生在推理中途。
用实验验证负载是否分配到所有卡
配置完成后,不能只看启动日志没有报错,要实际运行一段时间观察每卡负载。最直接的方法是持续调用推理请求,同时每 2-3 秒采样一次 nvidia-smi:
watch -n 2 nvidia-smi `--query-gpu`=index,memory.used,utilization.gpu,power.draw `--format`=csv
判断负载是否均匀的标准:在生成一张图的过程中,所有参与推理的卡都应该有接近同步的显存峰值和利用率变化。比如 4 卡并行,如果某张卡在采样期间一直为 0% 利用率,说明该卡未被使用;如果某张卡显存总是比其他卡高出一倍,说明 device_map 没有把大模块拆开。
更细致的验证方式是在代码里记录每卡的事件时间:
import torch
torch.cuda.synchronize()
start_event = [torch.cuda.Event(enable_timing=True) for _ in range(torch.cuda.device_count())]
# 在推理前为每卡记录 start,推理后记录 end
for i in range(torch.cuda.device_count()):
start_event[i].record()
# ... 推理代码 ...
for i in range(torch.cuda.device_count()):
end_time = torch.cuda.default_stream(i).get_device() # 仅为示意
# 实际应使用 end_event,比较各卡耗时
这个方式的判断标准是各卡的等待时间大致接近,如果某张卡时间明显更长,说明通信或切分不均衡。实际上,更简单的办法是看 nvidia-smi 里各卡利用率曲线是否同时出现峰值。如果出现某卡在推理期间一直空闲,就先查看 device_map 是否覆盖了该卡对应的层,再看是否为数据并行下 batch 分配不均(可以调大 batch 或看框架是否支持按卡负载动态分配)。