如何高效获取PyTorch nn.Module自身参数(不含子模块参数)
获取PyTorch nn.Module自身参数(不含子模块)的高效方法
要获取模块自身的参数(排除子模块的参数),最直接高效的方式是直接访问模块的_parameters属性,这是PyTorch内部存储模块自身参数的字典,无需遍历子模块,性能最优。
方法1:直接访问_parameters(推荐)
_parameters是一个OrderedDict,存储了所有通过nn.Parameter或register_parameter()直接注册到当前模块的参数,不包含子模块的参数。只需过滤掉其中的None值即可:
# 获取带名称的自身参数列表 own_named_params = [(name, param) for name, param in self._parameters.items() if param is not None] # 仅获取参数对象列表 own_params = [param for param in self._parameters.values() if param is not None]
方法2:通过named_parameters()过滤
如果需要使用named_parameters()的接口,可以通过检查参数的所属模块来过滤。每个参数的_owner属性指向其所属的模块,只需判断是否为当前模块即可:
own_named_params = [] for name, param in self.named_parameters(): if param._owner is self: own_named_params.append((name, param))
说明
- 方法1的效率更高,因为它直接访问模块内部存储的参数集合,无需遍历所有子模块的参数,适合大型模型场景。
_parameters中可能存在None值(例如动态设置参数为None的情况),因此需要额外过滤。
内容的提问来源于stack exchange,提问作者Roy
相关产品推荐
相关产品推荐

