如何使用PyTorch和Hugging Face实现GPU内存渐进式分配
实现PyTorch渐进式GPU内存分配的解决方案
你当前使用的Hugging Face Trainer默认启动多卡并行时,会按GPU数量均匀拆分总批量,每块卡的实际运行batch size都等于你设置的per_device_train_batch_size,因此每卡内存占用基本一致,不会触发优先用满单卡再用下一张的逻辑。你可以参考以下两种方案实现需求:
方案1:基于Hugging Face Trainer的最小改动实现
该方案依赖Hugging Face内置的设备映射能力,不需要修改核心训练逻辑即可满足需求:
- 首先调整模型加载逻辑,删除手动指定设备的代码,启用自动设备映射
device_map="auto"的内置分配规则正好是优先占满排序靠前的GPU可用显存,再依次使用后续GPU,完全匹配你的要求
import os import torch os.environ["CUDA_DEVICE_ORDER"]="PCI_BUS_ID" os.environ["CUDA_VISIBLE_DEVICES"]='1,2,3,4' max_length = 64 # 注意:删除原来的.to("cuda"),新增device_map参数 model = BertForSequenceClassification.from_pretrained( model_name, num_labels=2, device_map="auto" ) train_encodings = tokenizer(train_texts, truncation=True, padding=True, max_length=max_length) training_args = TrainingArguments( per_device_train_batch_size=64, # 关闭自动批量大小调整,避免覆盖你的自定义配置 auto_find_batch_size=False, ... ) trainer = Trainer( args=training_args, model=model, ... )
- 额外适配:如果希望数据也按相同逻辑分配显存,可在训练参数中新增
include_inputs_for_metrics=True,Trainer会自动把输入数据放到对应层所在的GPU上,不需要手动处理。
方案2:自定义训练循环实现(完全可控)
如果你需要更灵活的分配规则,可手动检测显存后自定义分配逻辑:
- 先检测所有可用GPU的空闲显存,计算单batch对应的显存占用量
# 按你指定的卡顺序获取空闲显存(单位GB) cuda_ids = [1,2,3,4] free_mem_per_gpu = [] for cuda_id in cuda_ids: free_mem = torch.cuda.mem_get_info(device=cuda_id)[0]/1024**3 free_mem_per_gpu.append((cuda_id, free_mem)) # 自行测试单batch=64对应的显存占用,替换为实际测量值 per_batch_mem_usage = 5.8 # 计算每张卡可承载的最大batch size batch_alloc = [] total_needed_batch = 256 # 替换为你需要的总批量大小 for cuda_id, free_mem in free_mem_per_gpu: if total_needed_batch <=0: break max_batch = int(free_mem / per_batch_mem_usage) * 64 alloc = min(max_batch, total_needed_batch) batch_alloc.append((f"cuda:{cuda_id}", alloc)) total_needed_batch -= alloc
- 手动拆分模型层到不同GPU,优先把模型层放到顺序靠前的卡,直到显存占满再拆分到下一张卡,数据加载时按上面计算的
batch_alloc给对应卡分配对应大小的batch即可。
注意事项
- 不要手动调用
model.to("cuda"),否则会覆盖device_map的分配逻辑,导致模型全部放到默认第一张卡触发显存溢出 - 渐进式显存分配会牺牲一定的多卡并行训练速度,默认的均匀分配是计算效率最高的方案,你需要在显存利用率和训练速度之间做权衡
- 不要搭配默认的DDP分布式训练使用,DDP要求每张卡的batch size一致,无法适配差异化的内存分配逻辑
内容的提问来源于stack exchange,提问作者Ondrej Sotolar
相关产品推荐
相关产品推荐

