如何调用被torch.nn.DataParallel包装的Module类非forward函数并保留并行功能?
问题:DataParallel包装Module后,如何并行调用自定义函数?
我定义了一个继承自torch.nn.Module的类A,并用torch.nn.DataParallel包装它。调用forward函数(即a())运行正常,但希望调用类A的其他自定义函数(如示例中的func1)时,也能保留DataParallel的数据并行功能。请问这是否可行?还是只能通过forward函数实现?
最小示例代码:
class A(torch.nn.Module): def __init__(self): # 初始化网络结构 super().__init__() # ... 省略具体实现 def forward(self, some_arguments): # forward逻辑 # ... 省略具体实现 def func1(self, some_arguments): # 自定义函数逻辑 # ... 省略具体实现 a = A() a = torch.nn.DataParallel(a, device_ids=[0, 1]) # 调用forward函数,运行正常 outputs = a(inputs) # 直接调用func1,无法运行 # outputs1 = a.func1(inputs) # 调用原始module的func1,可运行但无并行效果 outputs1 = a.module.func1(inputs)
回答
可行,不需要仅依赖forward方法实现并行。DataParallel默认只自动包装forward方法,所以直接调用a.func1会失败;而a.module.func1确实是单GPU运行,因为跳过了DataParallel的并行分发逻辑。要让自定义函数支持并行,可以通过以下几种方式实现:
方法1:手动模拟DataParallel的并行流程
直接复用PyTorch并行模块的工具函数,模仿DataParallel处理forward的步骤:
- 将输入数据拆分到各个GPU(
scatter) - 在每个GPU上复制模型副本(
replicate) - 并行调用自定义函数(
parallel_apply) - 收集各个GPU的结果(
gather)
示例代码:
from torch.nn.parallel import scatter, replicate, parallel_apply, gather def parallel_call_func1(data_parallel_model, inputs, device_ids): # 拆分输入到指定设备 inputs_scattered = scatter(inputs, device_ids) # 在各个设备上复制模型 model_replicas = replicate(data_parallel_model.module, device_ids) # 并行调用func1 outputs = parallel_apply(model_replicas, inputs_scattered, None, device_ids) # 收集结果到主设备 return gather(outputs, device_ids[0]) # 调用方式 outputs1 = parallel_call_func1(a, inputs, [0, 1])
方法2:扩展DataParallel类,添加自定义函数的包装
继承DataParallel类,为自定义函数实现并行逻辑,这样可以直接通过a.func1调用:
import torch.nn as nn class CustomDataParallel(nn.DataParallel): def func1(self, inputs): # 复用DataParallel内部的分发逻辑 inputs = self.scatter(inputs, self.device_ids) replicas = self.replicate(self.module, self.device_ids[:len(inputs)]) outputs = self.parallel_apply(replicas, inputs, None) return self.gather(outputs, self.output_device) # 使用自定义的DataParallel包装模型 a = CustomDataParallel(A(), device_ids=[0, 1]) # 直接调用func1即可获得并行效果 outputs1 = a.func1(inputs)
方法3:将自定义函数逻辑整合到forward(简易方案)
如果自定义函数逻辑不复杂,可以通过在forward中添加参数分支,控制执行不同逻辑:
class A(torch.nn.Module): def __init__(self): super().__init__() # ... 初始化逻辑 def forward(self, some_arguments, mode="forward"): if mode == "forward": # 原forward逻辑 pass elif mode == "func1": # 原func1逻辑 pass # ... 其他分支 # 调用时通过mode参数指定执行func1逻辑 outputs1 = a(inputs, mode="func1")
内容的提问来源于stack exchange,提问作者Nagabhushan S N
相关产品推荐
相关产品推荐

