Accelerate多GPU批量推理:数据拆分与结果顺序异常问询
Accelerate多进程推理结果顺序不符问题排查
问题描述
基于相关代码实现批量推理,分别以process=1和process=4开展实验,发现多进程模式下结果展平后顺序与单进程不一致,无法和ground truth对应。例如:
- 数据长度为5、批次大小为3时,单进程结果展平后为
[1,2,3,4,5] - 4进程模式下结果展平后顺序混乱
注:已通过zip(text,label)将数据传入进程实现映射,此非问题核心
相关代码
import random import os import numpy as np import torch from tqdm import tqdm from accelerate import Accelerator, notebook_launcher from transformers import set_seed, AutoTokenizer # 假设以下变量已提前定义 NUM_LABELS = 2 MAX_LENGTH = 512 zipped_text_label = [("text1", 0), ("text2", 1), ("text3", 0), ("text4", 1), ("text5", 0)] tokenizer = AutoTokenizer.from_pretrained("my_model_path") def seed_everything(seed=13): random.seed(seed) os.environ['PYTHONHASHSEED'] = str(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed(seed) set_seed(seed) torch.backends.cudnn.deterministic = True seed_everything(seed = 13) def load_model(model, lora, device, num_labels, merge_unload): # 假设此函数已实现模型加载逻辑 from peft import PeftModel from transformers import AutoModelForSequenceClassification base_model = AutoModelForSequenceClassification.from_pretrained(model, num_labels=num_labels) model = PeftModel.from_pretrained(base_model, lora) return model.to(device[""]) def test(): accelerator = Accelerator() accelerator.wait_for_everyone() seed_everything(seed = 13) model = load_model(model = "my_model_path", lora = "./my_lora_checkpoint/checkpoint-8200", device = {"": accelerator.process_index}, num_labels = NUM_LABELS, merge_unload = False) with accelerator.split_between_processes(zipped_text_label) as prompts: res = {"pred_probs": [], "pred_labels": []} BATCH_SIZE = 10 BATCHES = [prompts[i:i + BATCH_SIZE] for i in range(0, len(prompts), BATCH_SIZE)] for batch in tqdm(BATCHES): text_batch = [i[0] for i in batch] with torch.no_grad(): inputs = tokenizer(text_batch, truncation=True, max_length=MAX_LENGTH, padding="max_length", return_tensors="pt").to(model.device) logits = model(**inputs).logits.cpu().to(torch.float32) probs = torch.softmax(logits, dim=1).numpy() res["pred_probs"].append(probs.tolist()) res["pred_labels"].append(probs.argmax(axis=1).tolist()) result = accelerator.gather_object(res) if accelerator.is_main_process: print(result) notebook_launcher(test, num_processes=1)
问题根源
gather_object输入格式错误:代码中将res包装为[res]后传入gather_object,导致主进程收集到的结果结构嵌套异常(变为[[process0_res], [process1_res], ...]),后续合并时容易打乱顺序。- 结果合并逻辑缺失:未按进程分片的原始顺序,将各进程内的批次结果逐层展平拼接,直接打印原始收集结果会导致结构混乱,无法对应原始数据顺序。
解决方案
1. 修正gather_object输入
移除对res的列表包装,直接传入字典对象:
# 替换原代码中的 res = [res] 和 result = gather_object(res) result = accelerator.gather_object(res)
2. 按原始顺序合并结果
在主进程中,按进程顺序展平并拼接所有结果:
if accelerator.is_main_process: all_pred_probs = [] all_pred_labels = [] # 按进程顺序拼接结果,保证与原始数据顺序一致 for proc_res in result: # 展平当前进程内的批次结果 for batch_probs in proc_res["pred_probs"]: all_pred_probs.extend(batch_probs) for batch_labels in proc_res["pred_labels"]: all_pred_labels.extend(batch_labels) # 此时all_pred_labels顺序与原始zipped_text_label完全匹配 print("合并后的预测标签顺序:", all_pred_labels)
3. 验证分片顺序(可选)
在每个进程中打印分到的数据,确认拆分逻辑符合预期:
with accelerator.split_between_processes(zipped_text_label) as prompts: print(f"进程{accelerator.process_index}分到的数据:", prompts) # 后续处理逻辑...
内容的提问来源于stack exchange,提问作者Deshwal
相关产品推荐
相关产品推荐

