接口功能:实现门控循环单元(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)。当前暂不支持该模式。
计算公式:
其中:
- 为 sigmoid 激活函数。
- 表示逐元素乘法(Hadamard积)。
- 、、 分别为重置门、更新门和新门。
- 为当前时刻的隐状态。
- 、、 为输入权重矩阵,形状为 (转置存储为 )。
- 、、 为隐状态权重矩阵,形状为 (转置存储为 )。
- 若 hasBias=true,则 、、 和 、、 分别为输入偏置和隐状态偏置,形状均为 (存储为 )。
每个算子分为,必须先调用"aclnnGRUGetWorkspaceSize"接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用"aclnnGRU"接口执行计算。
[object Object]
[object Object]
- 确定性计算:
- [object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]、[object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]:aclnnGRU默认确定性实现。
[object Object]