使用MindSpore的LayerNorm报错ValueError: For 'LayerNorm', gamma or beta shape must match input shape.

1 系统环境

硬件环境(Ascend/GPU/CPU): Ascend/GPU/CPU
MindSpore版本: mindspore=2.0
执行模式(PyNative/ Graph):不限
Python版本: Python=3.7
操作系统平台: 不限

2 报错信息

2.1 问题描述

begin_norm_axis=1, begin_params_axis=1固定,并未实现与PyTorch完全一致的的功能。

ValueError                                Traceback (most recent call last)
Cell In[12], line 3
      1 add_norm = AddNorm([3, 4], 0.5)
      2 add_norm.set_train(False)
----> 3 add_norm(ops.ones((2, 3, 4)), ops.ones((2, 3, 4))).shape

File ~/anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/nn/cell.py:664, in Cell.__call__(self, *args, **kwargs)
    662 except Exception as err:
    663     _pynative_executor.clear_res()
--> 664     raise err
    666 if isinstance(output, Parameter):
    667     output = output.data

File ~/anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/nn/cell.py:661, in Cell.__call__(self, *args, **kwargs)
    659     _pynative_executor.new_graph(self, *args, **kwargs)
    660     output = self._run_construct(args, kwargs)
--> 661     _pynative_executor.end_graph(self, output, *args, **kwargs)
    662 except Exception as err:
    663     _pynative_executor.clear_res()

File ~/anaconda3/envs/MindSpore/lib/python3.9/site-packages/mindspore/common/api.py:1304, in _PyNativeExecutor.end_graph(self, obj, output, *args, **kwargs)
   1291 def end_graph(self, obj, output, *args, **kwargs):
   1292     """
   1293     Clean resources after building forward and backward graph.
   1294 
   (...)
   1302         None.
   1303     """
-> 1304     self._executor.end_graph(obj, output, *args, *(kwargs.values()))

ValueError: For 'LayerNorm', gamma or beta shape must match input shape, but got input shape: [const vector][2, 3, 4], gamma shape: [const vector][3, 4], beta shape: [const vector][3, 4].

----------------------------------------------------
- C++ Call Stack: (For framework developers)
----------------------------------------------------
mindspore/core/ops/layer_norm.cc:111 InferShape复制

2.2 脚本代码(代码格式,可上传附件)

class AddNorm(nn.Cell):
    """残差连接后进行层规范化"""
    def __init__(self, normalized_shape, dropout, **kwargs):
        super(AddNorm, self).__init__(**kwargs)
        self.dropout = nn.Dropout(p=dropout)
        self.ln = nn.LayerNorm(normalized_shape) #begin_norm_axis=1, begin_params_axis=1需要增加

    def construct(self, X, Y):
        return self.ln(self.dropout(Y) + X)

add_norm = AddNorm([3, 4], 0.5)
add_norm.set_train(False)
add_norm(ops.ones((2, 3, 4)), ops.ones((2, 3, 4))).shape

完整的代码可以在这里找到:
https://openi.pcl.ac.cn/kewei/d2lkewei-ms/src/branch/master/chapter10_attention-mechanisms/7-transformer-ms.ipynb

3 根因分析

问题出在nn.LayerNorm这个算子上,单独把算子抽出来看。
如下代码就能复现问题,normalized_shape取[3, 4], 输入shape[2, 3, 4]

import mindspore as ms  
import numpy as np  
x = ms.Tensor(np.ones([2, 3, 4]), ms.float32)  
m = ms.nn.LayerNorm([3, 4])  
output = m(x).shape  
print(output)


报错是一样的。
API文档

class mindspore.nn.LayerNorm(normalized_shape , begin_norm_axis=- 1 , begin_params_axis=- 1 , gamma_init=‘ones’ , beta_init=‘zeros’ , epsilon=1e-07)


input_shape[begin_norm_axis:] 等于 normalized_shape
问题就出在这里,begin_norm_axis的默认值是 -1
normalized_shape是(3, 4)
所以就不一致了。
查看nn.LayerNorm的代码

self.gamma = Parameter(initializer(  
    gamma_init, normalized_shape), name="gamma")  
self.beta = Parameter(initializer(  
    beta_init, normalized_shape), name="beta")  
self.layer_norm = P.LayerNorm(begin_norm_axis=self.begin_norm_axis,  
                                begin_params_axis=self.begin_params_axis,  
                                epsilon=self.epsilon)  

def construct(self, input_x):  
y, _, _ = self.layer_norm(input_x, self.gamma.astype(input_x.dtype), self.beta.astype(input_x.dtype))  
return y  

gamma和beta的shape都等于normalized_shape,最终调用的是算子ops.LayerNorm。所以报错信息才会有For ‘LayerNorm’, gamma or beta shape must match input shape
这样的错误信息。

4 解决方案

nn.LayerNorm的参数 begin_norm_axis设置为1, 同时begin_params_axis也设置为1

import mindspore as ms  
import numpy as np  
x = ms.Tensor(np.ones([2, 3, 4]), ms.float32)  
m = ms.nn.LayerNorm([3, 4],begin_norm_axis =1,begin_params_axis=1)  
output = m(x).shape  
print(output)

cke_109263.png