Developers
Download
[object Object][object Object][object Object]undefined
[object Object]
  • 接口功能:实现门控循环单元(Gated Recurrent Unit, GRU)计算,支持多层堆叠、双向、定长序列和不定长序列(PackedSequence)两种输入模式。训练模式下可输出各门控中间结果(r、z、n、hn、h),用于反向传播。

  • 输入输出shape:

    • 定长模式(batchSizes为nullptr):如果输入张量 input 的shape为(T, B, I)(batchFirst=false)或(B, T, I)(batchFirst=true),则输出张量 output 的shape为(T, B, D H)(batchFirst=false)或(B, T, D H)(batchFirst=true),最终隐状态 hy 的shape为(L * D, B, H)。
    • 不定长模式(batchSizes非nullptr):输入张量 input 的shape为(sum(batch_size), I),输出张量 output 的shape为(sum(batch_size), D H),最终隐状态 hy 的shape为(L D, B, H)。当前暂不支持该模式。
  • 计算公式:

    rt=σ(Wirxt+bir+Whrh(t1)+bhr)r_t = \sigma(W_{ir} x_t + b_{ir} + W_{hr} h_{(t-1)} + b_{hr}) zt=σ(Wizxt+biz+Whzh(t1)+bhz)z_t = \sigma(W_{iz} x_t + b_{iz} + W_{hz} h_{(t-1)} + b_{hz}) nt=tanh(Winxt+bin+rt(Whnh(t1)+bhn))n_t = \tanh(W_{in} x_t + b_{in} + r_t \odot (W_{hn} h_{(t-1)} + b_{hn})) ht=(1zt)nt+zth(t1)h_t = (1 - z_t) \odot n_t + z_t \odot h_{(t-1)}

    其中:

    • σ\sigma 为 sigmoid 激活函数。
    • \odot 表示逐元素乘法(Hadamard积)。
    • rtr_tztz_tntn_t 分别为重置门、更新门和新门。
    • hth_t 为当前时刻的隐状态。
    • WirW_{ir}WizW_{iz}WinW_{in} 为输入权重矩阵,形状为 [H,I][H, I](转置存储为 [3H,I][3H, I])。
    • WhrW_{hr}WhzW_{hz}WhnW_{hn} 为隐状态权重矩阵,形状为 [H,H][H, H](转置存储为 [3H,H][3H, H])。
    • 若 hasBias=true,则 birb_{ir}bizb_{iz}binb_{in}bhrb_{hr}bhzb_{hz}bhnb_{hn} 分别为输入偏置和隐状态偏置,形状均为 [H][H](存储为 [3H][3H])。
[object Object]

每个算子分为,必须先调用"aclnnGRUGetWorkspaceSize"接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用"aclnnGRU"接口执行计算。

[object Object]
[object Object]
[object Object]
  • 参数说明

    [object Object]
  • 返回值

    aclnnStatus:返回状态码,具体参见

    第一段接口完成入参校验,出现以下场景时报错:

    [object Object]
[object Object]
  • 参数说明

    [object Object]
  • 返回值

    aclnnStatus:返回状态码,具体参见

[object Object]
  • 确定性计算:
    • [object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]、[object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]:aclnnGRU默认确定性实现。
[object Object]

示例代码如下,仅供参考,具体编译和执行过程请参考

[object Object]