PyTorch DDP双GPU训练耗时超单GPU,求解决方案
双GPU DDP训练比单GPU更慢,疑似未并行计算
我在同一台机器上使用两块RTX3090 GPU,通过PyTorch的DistributedDataParallel(DDP)训练模型,发现每个epoch训练耗时约133秒,反而比单GPU训练的105秒更长。数据已正确加载到两块GPU内存中,但双GPU似乎未执行并行计算,而是重复运行相同任务。以下是核心代码、启动命令及训练日志截图:
核心代码
import argparse import os import pickle from torch.utils.data.distributed import DistributedSampler import ruamel.yaml as yaml import time from pathlib import Path import torch import torch.nn as nn from dataset.dataset_DPT import VQAFeatureDataset from torch.utils.data import DataLoader from m3ae.modules import M3AETransformerSS import torch.distributed as dist import numpy as np from dataset.tools import Logger, create_dir, set_schedule import torch.nn.functional as F @torch.no_grad() def evaluation(model, data_loader, device): # test score = 0 open_ended = 0 closed_ended = 0 open_score = 0 close_score = 0 stage1_dict = [] # model.load_state_dict(torch.load('model_mlm.pth')) model.eval() header = 'Generate VQA test result:' with torch.no_grad(): for index, batch in enumerate(data_loader): batch['image'][0] = batch['image'][0].to(device) batch['text_ids'] = batch['text_ids'].to(device) batch['text_labels'] = batch['text_labels'].to(device) batch['text_ids_mlm'] = batch['text_ids_mlm'].to(device) batch['text_labels_mlm'] = batch['text_labels_mlm'].to(device) batch['text_masks'] = batch['text_masks'].to(device) target_list = [] for l in range(len(batch['vqa_labels'])): target_list.append(batch['vqa_labels'][l].unsqueeze(0)) targets = torch.cat(target_list, dim=0).to(device) target = torch.argmax(targets, dim=1) ans_type = batch['answer_types'] phrase_type = batch['phrase_type'] logits = model(batch) values, indices = torch.topk(logits, k=8, dim=-1) values = values.detach().cpu().numpy() indices = indices.detach().cpu().numpy() for i in range(values.shape[0]): stage1_dict.append((indices[i].astype(np.int16), values[i].astype(np.float16))) pred_score = torch.argmax(logits, dim=1) for i in range(len(ans_type)): if ans_type[i] == 'OPEN': open_ended += 1 if target[i] == pred_score[i]: open_score += 1 elif ans_type[i] == 'CLOSED': closed_ended += 1 if target[i] == pred_score[i]: close_score += 1 score += open_score + close_score # with open('data/vqa/data_RAD/stage1_train.pkl', 'wb') as fp: # pickle.dump(stage1_dict, fp) score = (score / (open_ended + closed_ended)) open_score = (open_score / open_ended) close_score = (close_score / closed_ended) return score, open_score, close_score def main(args, config): create_dir('outputs') logger = Logger('outputs/log.txt') logger.write(args.__repr__()) # device = torch.device(args.device) if args.local_rank != -1: torch.cuda.set_device(args.local_rank) device = torch.device("cuda", args.local_rank) torch.distributed.init_process_group(backend="nccl") trainset = VQAFeatureDataset('train') testset = VQAFeatureDataset('test') # tokenizer = BertTokenizer.from_pretrained(args.text_encoder) #### Creating Model #### print("Creating model") config['load_path'] = '/home/liyong/PythonWorkspace/M3AE-vqa/checkpoints/m3ae.ckpt' # config['load_path'] = '' model = M3AETransformerSS(config) model = model.to(device) num_gpus = torch.cuda.device_count() if num_gpus > 1: print('use {} gpus!'.format(num_gpus)) model = nn.parallel.DistributedDataParallel(model, device_ids=[args.local_rank], output_device=args.local_rank, find_unused_parameters=True) word_size = dist.get_world_size() train_sampler = DistributedSampler(trainset, num_replicas=word_size, rank=args.local_rank) train_loader = DataLoader(trainset, batch_size=config['batch_size'], num_workers=4, sampler=train_sampler, collate_fn=trainset.collote) test_loader = DataLoader(testset, batch_size=config['batch_size'], num_workers=4, collate_fn=testset.collote, shuffle=False) optimizer, scheduler = set_schedule(model, config, len(trainset.entries)) # print(model) # score, open_score, close_score = evaluation(model, test_loader, device) best_score = 0 best_epoch = 0 for epoch in range(config['max_epoch']): train_sampler.set_epoch(epoch) strat_time = time.time() print(f"Start running basic DDP example on rank {args.local_rank}.") total_loss = 0 for index, batch in enumerate(train_loader): batch['image'][0] = batch['image'][0].to(device) batch['text_ids'] = batch['text_ids'].to(device) batch['text_labels'] = batch['text_labels'].to(device) batch['text_ids_mlm'] = batch['text_ids_mlm'].to(device) batch['text_labels_mlm'] = batch['text_labels_mlm'].to(device) batch['text_masks'] = batch['text_masks'].to(device) target_list = [] for l in range(len(batch['vqa_labels'])): target_list.append(batch['vqa_labels'][l].unsqueeze(0)) targets = torch.cat(target_list, dim=0).to(device) logits = model(batch) # loss = criterion(logits.float(), targets) loss = F.binary_cross_entropy_with_logits(logits.float(), targets) total_loss += loss optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step() end_time = time.time() if args.local_rank == 0: score, open_score, close_score = evaluation(model, test_loader, device) logger.write('epoch ========== %d' % (epoch)) logger.write('overall: %.4f, open: %.4f, close: %.4f, loss: %.4f, lr: %.6f, time: %.4f' % (score, open_score, close_score, total_loss, optimizer.state_dict()['param_groups'][0]['lr'], end_time-strat_time)) if score > best_score: torch.save(model.module.state_dict(), '/home/liyong/PythonWorkspace/M3AE-vqa/model_mlm_ddp.pth') best_epoch = epoch best_score = score # print("best_score ======= " + str(round(best_score, 4)) + " best_epoch ======= " + str(best_epoch)) logger.write('best_score: %.4f, best_epoch: %d' % (best_score, best_epoch)) if __name__ == '__main__': parser = argparse.ArgumentParser() parser.add_argument('--config', default='/home/liyong/PythonWorkspace/M3AE-vqa/configs/RAD_M3AE.yaml') # parser.add_argument('--checkpoint', default='./pretrain/2022-09-11/med_pretrain_29.pth') parser.add_argument('--checkpoint', default=None) parser.add_argument('--output_dir', default='/home/liyong/PythonWorkspace/M3AE-vqa/output/rad') parser.add_argument('--evaluate', action='store_true') parser.add_argument('--text_encoder', default='bert-base-uncased') parser.add_argument('--text_decoder', default='bert-base-uncased') parser.add_argument('--device', type=int, default=0) parser.add_argument('--seed', default=42, type=int) parser.add_argument('--world_size', default=1, type=int, help='number of distributed processes') parser.add_argument('--dist_url', default='env://', help='url used to set up distributed training') parser.add_argument('--distributed', default=False, type=bool) parser.add_argument("--local_rank", default=os.getenv('LOCAL_RANK', -1), type=int) args = parser.parse_args() config = yaml.load(open(args.config, 'r'), Loader=yaml.Loader) config['image_size'] = 224 config['tokenizer'] = 'roberta-base' args.result_dir = os.path.join(args.output_dir, 'result') Path(args.output_dir).mkdir(parents=True, exist_ok=True) Path(args.result_dir).mkdir(parents=True, exist_ok=True) # yaml.dump(config, open(os.path.join(args.output_dir, 'config.yaml'), 'w')) print("config: ", config) print("args: ", args) main(args, config)
启动命令
python -m torch.distributed.launch --nproc_per_node=2 DPT_MLM_ddp2.py
训练日志截图

问题排查与修复方案
1. 分布式初始化逻辑错误
当前代码中dist.get_world_size()和DistributedSampler的创建在DDP初始化之前,且未判断分布式是否初始化成功,导致采样器无法正确拆分数据集。
修复代码:
def main(args, config): create_dir('outputs') logger = Logger('outputs/log.txt') logger.write(args.__repr__()) # 优先初始化分布式环境 if args.local_rank != -1: torch.cuda.set_device(args.local_rank) device = torch.device("cuda", args.local_rank) torch.distributed.init_process_group(backend="nccl") world_size = dist.get_world_size() else: device = torch.device("cuda", args.device) world_size = 1 trainset = VQAFeatureDataset('train') testset = VQAFeatureDataset('test') print("Creating model") config['load_path'] = '/home/liyong/PythonWorkspace/M3AE-vqa/checkpoints/m3ae.ckpt' model = M3AETransformerSS(config) model = model.to(device) # 仅在分布式环境下初始化DDP if args.local_rank != -1: print(f'use {world_size} gpus!') model = nn.parallel.DistributedDataParallel(model, device_ids=[args.local_rank], output_device=args.local_rank, find_unused_parameters=False) # 根据环境选择采样器 if args.local_rank != -1: train_sampler = DistributedSampler(trainset, num_replicas=world_size, rank=args.local_rank) train_loader = DataLoader(trainset, batch_size=config['batch_size'], num_workers=4, sampler=train_sampler, collate_fn=trainset.collote) else: train_loader = DataLoader(trainset, batch_size=config['batch_size'], num_workers=4, shuffle=True, collate_fn=trainset.collote)
2. 批处理大小未适配分布式
DDP中每个GPU处理的是原batch_size的子集,若保持原batch_size不变,总batch_size会变为batch_size * world_size,但单GPU计算量未减少,反而增加通信开销。
修复方式:
- 若要保持总batch_size与单GPU一致,将config中的
batch_size改为原来的1/2(双GPU场景); - 若要提升训练速度,保持原batch_size不变,总batch_size翻倍,同时调整学习率(乘以world_size)。
3. 冗余进程操作优化
- 去掉
num_gpus = torch.cuda.device_count()的判断,改用world_size确认GPU数量,且仅让rank0进程输出GPU使用信息; - 训练循环中的打印语句仅让rank0执行,避免多进程重复输出。
4. 性能细节优化
- 关闭
find_unused_parameters=True(确认模型所有参数都参与计算后),减少额外通信开销; - 适当增加
num_workers(如改为8),避免数据加载成为瓶颈; - 计时逻辑仅在rank0进程中执行,避免多进程重复计时。
修复后的epoch循环核心代码:
for epoch in range(config['max_epoch']): if args.local_rank != -1: train_sampler.set_epoch(epoch) # 仅rank0进程计时 if args.local_rank == 0: start_time = time.time() print(f"Start epoch {epoch}") total_loss = 0 for index, batch in enumerate(train_loader): # 数据加载到对应device batch['image'][0] = batch['image'][0].to(device) batch['text_ids'] = batch['text_ids'].to(device) batch['text_labels'] = batch['text_labels'].to(device) batch['text_ids_mlm'] = batch['text_ids_mlm'].to(device) batch['text_labels_mlm'] = batch['text_labels_mlm'].to(device) batch['text_masks'] = batch['text_masks'].to(device) target_list = [] for l in range(len(batch['vqa_labels'])): target_list.append(batch['vqa_labels'][l].unsqueeze(0)) targets = torch.cat(target_list, dim=0).to(device) logits = model(batch) loss = F.binary_cross_entropy_with_logits(logits.float(), targets) total_loss += loss.item() optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step() # 仅rank0进程记录日志 if args.local_rank == 0: end_time = time.time() epoch_time = end_time - start_time score, open_score, close_score = evaluation(model, test_loader, device) logger.write('epoch ========== %d' % (epoch)) logger.write('overall: %.4f, open: %.4f, close: %.4f, loss: %.4f, lr: %.6f, time: %.4f' % (score, open_score, close_score, total_loss/len(train_loader), optimizer.state_dict()['param_groups'][0]['lr'], epoch_time)) if score > best_score: torch.save(model.module.state_dict(), '/home/liyong/PythonWorkspace/M3AE-vqa/model_mlm_ddp.pth') best_epoch = epoch best_score = score logger.write('best_score: %.4f, best_epoch: %d' % (best_score, best_epoch))
内容的提问来源于stack exchange,提问作者lycutter
相关产品推荐
相关产品推荐

