Batchnormgrad
=================


逐元素计算加法梯度

计算批标准化 (Batch Normalization) 的梯度。

该算子计算损失函数 L 分别对输入 `x`、缩放因子 `scale` (γ) 和偏置 `bias` (β) 的梯度。其中 `bias` 的梯度为 `dbias`。

.. math::

    dscale(\gamma) &= \sum_{i=1}^{m} dy_i \cdot \hat{x}_i \\
    dbias(\beta) &= \sum_{i=1}^{m} dy_i

.. math::

    dx_i = \frac{\gamma}{m\sqrt{\sigma^2 + \epsilon}} \left[ m \cdot dy_i - \sum_{j=1}^{m}dy_j - \hat{x}_i \sum_{j=1}^{m}dy_j \hat{x}_j \right]

其中 :math:`m` 是批处理大小 (batch)，:math:`\hat{x}` 是归一化后的 :math:`x`。

输入：
    - **x** - 前向传播时的输入张量。
    - **dy** - 来自后一层的上游梯度。
    - **params** - 其他参数打包成数组。
    - **core_mask** - 核掩码。

输出：
    - **dx** - 对输入 `x` 的梯度。
    - **dbias** - 对偏置 `bias` (β) 的梯度。
    - **dscale** - 对缩放因子 `scale` (γ) 的梯度。

支持平台：
    ``FT78NE``
    ``MT7004``

.. note::
    - FT78NE 支持fp32
    - MT7004 支持fp16, fp32

**参数数组结构：**

.. code-block:: c
    :linenos:

    long long params[12];
    params[0] = (long long)mean; 前向传播时计算的均值。
    params[1] = (long long)invar; 前向传播时计算的逆方差 (1 / sqrt(variance + epsilon))。
    params[2] = (long long)scale; 前向传播时使用的缩放因子 (gamma, γ)。
    params[3] = (long long)dbias; 对偏置 `bias` (β) 的梯度。
    params[4] = (long long)dscale; 对缩放因子 `scale` (γ) 的梯度。
    params[5] = (long long)batch; 批处理大小。 
    params[6] = (long long)channel; 通道数。
    params[7] = (long long)is_train; 是否为训练模式。

**共享存储版本:**

.. c:function:: void fp_batch_norm_grad_s(float* x, float* dy, float* dx, int core_mask)
.. c:function:: void hp_batch_norm_grad_s(half* x, half* dy, half* dx, int core_mask)

**C调用示例：**

.. code-block:: c
    :linenos:
    :emphasize-lines: 20

    //FT78NE示例
    #include <stdio.h>
    #include <batchnormgrad.h>
    int main(int argc, char* argv[]) {
        float *x = (float *)0xA0000000;          // forward input x
        float *dy = (float *)0xB0000000;         // upstream gradient dy
        float *mean = (float *)0xC0000000;       // forward mean
        float *invar = (float *)0xD0000000;      // forward inverse variance
        float *scale = (float *)0xE0000000;      // forward scale (gamma)
        
        float *dx = (float *)0xA1000000;         // output gradient dx
        float *dbias = (float *)0xB1000000;      // output gradient dbias
        float *dscale = (float *)0xC1000000;     // output gradient dscale

        int batch = 4;
        int channel = 64;
        int is_train = true;
        int core_mask = 0xff;
        
        long long params[12];
        params[0] = (long long)mean;
        params[1] = (long long)invar;
        params[2] = (long long)scale;
        params[3] = (long long)dbias;
        params[4] = (long long)dscale;
        params[5] = (long long)batch;
        params[6] = (long long)channel;
        params[7] = (long long)is_train;
        fp_batch_norm_grad_s(x, dy, dx, core_mask);
        return 0;
    }


**私有存储版本:**

.. c:function:: void fp_batch_norm_grad_p(float* x, float* dy, float* dx, long long *params)
.. c:function:: void hp_batch_norm_grad_p(half* x, half* dy, half* dx, long long *params)

   
**C调用示例：**

.. code-block:: c
    :linenos:
    :emphasize-lines: 19

    //FT78NE示例
    #include <stdio.h>
    #include <batchnormgrad.h>
    int main(int argc, char* argv[]) {
        float *x = (float *)0x10000000;          // forward input x in L2 space
        float *dy = (float *)0x10100000;         // upstream gradient dy
        float *mean = (float *)0x10200000;       // forward mean
        float *invar = (float *)0x10300000;      // forward inverse variance
        float *scale = (float *)0x10400000;      // forward scale (gamma)
        
        float *dx = (float *)0x10500000;         // output gradient dx
        float *dbias = (float *)0x10600000;      // output gradient dbias
        float *dscale = (float *)0x10700000;     // output gradient dscale

        int batch = 4;
        int channel = 32;
        int is_train = true;
        
        long long params[12];
        params[0] = (long long)mean;
        params[1] = (long long)invar;
        params[2] = (long long)scale;
        params[3] = (long long)dbias;
        params[4] = (long long)dscale;
        params[5] = (long long)batch;
        params[6] = (long long)channel;
        params[7] = (long long)is_train;

        fp_batch_norm_grad_p(x, dy, dx, params);
        return 0;
    }