PyTorch中nn.DataParallel场景下var.to(device)的替代方案
在nn.DataParallel场景下迁移数据的正确写法
设备(device)的设置
你需要将device设为nn.DataParallel指定的device_ids中的主设备——也就是列表的第一个元素(如果device_ids是[0,1],那主设备就是cuda:0,这也是默认的主设备)。
数据迁移的替代写法
写法和原语句基本一致,只需要确保数据迁移到主设备即可:
# 定义主设备 device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") # 迁移数据到主设备 X_batch, y_batch = X_batch.to(device), y_batch.to(device)
关键说明
nn.DataParallel会自动帮你处理多GPU的数据分发与结果汇总,不需要手动拆分数据到各个GPU:
- 它会把主设备上的
X_batch、y_batch拆分后,分发到device_ids列表中的所有GPU; - 在各个GPU上并行执行模型前向传播;
- 最后把所有GPU的计算结果收集回主设备。
如果初始化nn.DataParallel时手动指定了output_device,则需要将数据迁移到output_device对应的设备,但默认情况下output_device就是device_ids[0],所以不用额外修改。
内容的提问来源于stack exchange,提问作者Adnan Ali
相关产品推荐
相关产品推荐

