PyTorch 1.2.0中Weight Normalization产生NaN问题的修改方法咨询
1. 为什么修改symbolic_opset9.py里的_weight_norm没有效果
你定位到的torch/onnx/symbolic_opset9.py下的_weight_norm是ONNX模型导出专用的符号映射函数,仅在调用torch.onnx.export()导出模型时才会触发调用,正常训练、推理的计算流不会走该函数,因此添加打印不会有输出,修改这里也无法影响实际运行的权重归一化逻辑。
2. 加eps避免NaN的正确修改方式
你最开始定位的torch/nn/utils/weight_norm.py就是正确的修改路径,核心逻辑在compute_weight函数中,PyTorch 1.2.0版本的该函数默认实现等价于:
def compute_weight(module, name): weight_g = getattr(module, name + '_g') weight_v = getattr(module, name + '_v') return weight_v * (weight_g / torch.norm(weight_v, dim=module.dim, keepdim=True))
只需要在范数计算的分母位置添加eps即可,修改后代码如下:
def compute_weight(module, name): weight_g = getattr(module, name + '_g') weight_v = getattr(module, name + '_v') eps = 1e-6 return weight_v * (weight_g / (torch.norm(weight_v, dim=module.dim, keepdim=True) + eps))
修改后保存文件即可生效,不需要改动其他文件。
3. 关于替换1.9.0版本函数的说明
不建议直接替换本地1.2.0版本的_weight_norm相关文件,1.2和1.9版本间隔较大,内部依赖的底层API存在差异,直接替换大概率会引发兼容性报错。如果你需要可配置的eps参数,可以按以下步骤修改本地的weight_norm.py:
- 第一步:修改
WeightNorm类的__init__方法,新增eps入参,默认值设为1e-6,并将其保存为类实例属性 - 第二步:修改
compute_weight函数,从module实例中读取eps属性,添加到范数分母 - 第三步:修改对外暴露的
weight_norm接口,新增eps参数透传给WeightNorm类,调用时就可以自定义eps取值
更稳妥的无侵入修改方案
如果不想改动PyTorch原生源码,可以自己实现带eps的权重归一化封装,替换原有调用即可,示例代码如下:
import torch from torch.nn.utils import weight_norm as torch_original_weight_norm def weight_norm_with_eps(module, name='weight', dim=0, eps=1e-6): # 调用原生接口完成参数注册 module = torch_original_weight_norm(module, name, dim) # 定义带eps的前向预处理钩子 def custom_compute_weight(module, _): g = getattr(module, f"{name}_g") v = getattr(module, f"{name}_v") norm_v = torch.norm(v, dim=dim, keepdim=True) setattr(module, name, v * (g / (norm_v + eps))) # 注册钩子覆盖原有计算逻辑 module.register_forward_pre_hook(custom_compute_weight) return module
使用时直接将原有代码中torch.nn.utils.weight_norm(层实例)替换为weight_norm_with_eps(层实例)即可,不会影响环境中其他项目的PyTorch使用。
内容的提问来源于stack exchange,提问作者Mohit Lamba
相关产品推荐
相关产品推荐

