如何在PyTorch中实现菱形继承的模型并解决初始化参数报错问题?
如何在PyTorch中实现菱形继承的模型并解决初始化参数报错问题?
这个问题其实是Python多继承里*MRO(方法解析顺序)*导致的典型问题,咱们一步步来拆解解决,同时满足你要求的B、C能单独实例化的需求~
先分析报错原因
你写的代码里,类D的继承顺序是 D → B → C → A → nn.Module,当B的__init__里调用super().__init__(ratio=4)时,实际上调用的是C的__init__,而不是A的!这时候C的__init__要求必须传c_args参数,但B的调用里没传,自然就触发了参数缺失的报错。
解决方案:统一参数传递逻辑
核心思路是让每个类的__init__都正确传递参数:自己需要的参数单独接收,把不需要的参数通过**kwargs传给下一个父类,这样不管是单独实例化类,还是作为多继承的父类,都能正常工作。
下面是修改后的完整代码:
import torch import torch.nn as nn import torch.nn.functional as F import torch.utils.model_zoo as modelzoo import torch.distributed as dist class A(nn.Module): def __init__(self, ratio=4, *args, **kwargs): super().__init__(*args, **kwargs) # 把多余参数传给nn.Module(好习惯) self.conv_base = nn.Conv2d(3, 3 * ratio, 3, 1, 1) class B(A): def __init__(self, b_args, **kwargs): # 传递ratio参数和剩余的kwargs给下一个父类 super().__init__(ratio=4, **kwargs) self.conv1 = nn.Conv2d(4, 3, 1, 1, 0) class C(A): def __init__(self, c_args, **kwargs): # 同理传递参数 super().__init__(ratio=4, **kwargs) self.conv2 = nn.Conv2d(4, 3, 1, 1, 0) class D(B, C): def __init__(self, b_args, c_args): # 把两个参数都传给super(),让MRO链上的父类按需获取 super().__init__(b_args=b_args, c_args=c_args) self.conv3 = nn.Conv2d(4, 3, 1, 1, 0) # 测试单独实例化B、C b_args = dict(a=1) c_args = dict(b=2) b_model = B(b_args=b_args) c_model = C(c_args=c_args) # 测试实例化D model = D(b_args, c_args) print(model)
代码说明
- 类A:保留
ratio参数,同时接收*args, **kwargs并传给父类nn.Module,避免后续扩展出问题。 - 类B/C:单独接收自己需要的
b_args/c_args,然后把ratio和剩余的kwargs传给super():- 当单独实例化B/C时,
super()会调用A的__init__,kwargs是空的,完全没问题; - 当作为D的父类时,B的
super()会调用C的__init__,此时kwargs里包含c_args,刚好满足C的参数需求。
- 当单独实例化B/C时,
- 类D:把
b_args和c_args都传给super(),参数会沿着MRO链依次传递,B先拿到b_args,剩下的c_args传给C,最后C把参数传给A,完美解决菱形继承的参数传递问题。
这样修改后,不管是单独使用B、C,还是使用组合后的D,都能正常初始化运行啦~
备注:内容来源于stack exchange,提问作者coin cheung
相关产品推荐
相关产品推荐

