已关闭
[Requirement|需求建议]: rnn/thnn_fused_gru_cell算子开发 #6174
huangzhiyuan创建于  10 天前关闭于  9 天前
huangzhiyuan
huangzhiyuan成员
10 天前 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

Backgroud(背景信息)

新增ThnnFusedGruCell算子

1. 功能定位

对 GRU(Gated Recurrent Unit,门控循环单元)的单个时间步执行门控融合计算:输入两路门控预激活(输入侧 + 隐层侧)、上一步隐状态和可选双 bias,一次算子调用同时产出两个结果——新隐状态 hy 和反向计算复用的中间量 storage。适用于 RNN 逐步推理/训练中的高频单步计算,等价于 PyTorch torch.nn.functional.gru_cell 的 THNN 内核语义。

  • 支持产品:Ascend 950PR / Ascend 950DT(arch35),AIV 纯矢量算子。
  • 输出为双结果融合算子:hy 供下一时间步递推,storage 供反向梯度计算复用,避免训练场景重复计算。

2. 计算公式

门控预激活沿最后一维按门序 [r, z, n] 三等分:

  • input_gates (B, 3H) → gi_r、gi_z、gi_n(各 H 列)
  • hidden_gates (B, 3H) → gh_r、gh_z、gh_n
  • 两路 bias 同步三等分为 b1_r/b1_z/b1_n 与 b2_r/b2_z/b2_n,沿 batch 维广播
  • hx (B, H) 为上一步隐状态

计算过程:

rg=11+e−(gir+ghr+b1r+b2r)rg = \frac{1}{1 + e^{-(gi_r + gh_r + b1_r + b2_r)}}

zg=11+e−(giz+ghz+b1z+b2z)zg = \frac{1}{1 + e^{-(gi_z + gh_z + b1_z + b2_z)}}

ng=tanh⁡(gin+b1n+rg×(ghn+b2n))ng = \tanh(gi_n + b1_n + rg \times (gh_n + b2_n))

hy=ng+zg×(hx−ng)hy = ng + zg \times (hx - ng)

storage=[rg∣zg∣ng∣hx∣(ghn+b2n)]storage = [rg \mid zg \mid ng \mid hx \mid (gh_n + b2_n)]

语义说明:

  • rg 重置门、zg 更新门:sigmoid 激活,取值 (0, 1)。
  • ng 候选状态:tanh 激活,重置门 rg 调制隐层侧候选分量。
  • hy:更新门 zg 在旧状态 hx 与候选状态 ng 之间做凸组合。
  • storage:反向复用中间量,五段 [rg、zg、ng、hx、gh_n+b2_n] 沿最后一维拼接,每段 H 列,故 shape 为 (B, 5H)。
  • bias 缺省(空指针)时等价于全零 bias。

3. 输入输出

参数名 输入/输出/属性 描述 数据类型 数据格式
input_gates 输入 输入侧门控预激活,对应 gi_r/gi_z/gi_n,shape 为 (B, 3H) BFLOAT16、FLOAT16、FLOAT ND
hidden_gates 输入 隐层侧门控预激活,对应 gh_r/gh_z/gh_n,shape 为 (B, 3H) BFLOAT16、FLOAT16、FLOAT ND
hx 输入 上一步隐状态,shape 为 (B, H) BFLOAT16、FLOAT16、FLOAT ND
input_bias 可选输入 输入侧 bias,对应 b1_r/b1_z/b1_n,shape 为 (3H,);缺省等价全零 BFLOAT16、FLOAT16、FLOAT ND
hidden_bias 可选输入 隐层侧 bias,对应 b2_r/b2_z/b2_n,shape 为 (3H,);缺省等价全零 BFLOAT16、FLOAT16、FLOAT ND
hy 输出 新隐状态,shape 为 (B, H) BFLOAT16、FLOAT16、FLOAT ND
storage 输出 反向复用中间量,shape 为 (B, 5H) BFLOAT16、FLOAT16、FLOAT ND

4. 精度说明

  • BFLOAT16 / FLOAT16 输入在 fp32 中间精度下完成 sigmoid、tanh 与乘加计算,结果舍回原数据类型。
  • FLOAT 输入直接计算。
  • 输出数据类型与输入一致(全张量同 dtype)。

5. 实现要点(kernel 双路径)

路径 tilingKey 适用场景 管线
常规路径(RANK=4) 0 通用场景,H 优先切分 6 相位矢量流水线:NDDMA 搬入(含 bias 随路广播)→ fp32 计算域(3-buffer:累加器/结果载体/第二操作数)→ DataCopyPad 搬出
h-major 转置路径(RANK=8 哨兵档) 1 小 H 大 B 家族 连续 2D 搬入 → UB 内 ConfusionTransposeOnly 转置 → 面板域计算 → 结果拼接转置回 → 连续 2D 搬出

tiling 侧按 UB 容量三形态切分:形态 A 沿 H 切(单行 tile)、形态 B 沿 B 切(多行 tile,多核并行)、形态 C 全量单 tile;行级 32B 对齐填充(paddedHidden)消除非对齐 H 的形态退化;微量数据(<160B)回退单核避免 launch-bound 退化。

6. 约束说明

  • 输入与输出的数据类型必须一致,仅支持 BFLOAT16、FLOAT16、FLOAT;不支持跨数据类型组合,也不支持 DOUBLE、INT64 等其它数据类型。
  • input_gates 与 hidden_gates 的 shape 必须相同且为 (B, 3H),hx 的 shape 为 (B, H),需满足 input_gates.shape[1] == 3 × hx.shape[1];input_bias 与 hidden_bias 在位时元素个数必须为 3H 且两者相同。
  • 输入与输出的数据格式仅支持 ND。
  • B=0 或 H=0(numel 为 0 的空 Tensor)为合法输入,直接返回空输出。
  • aclnn 接口中可选 bias 以空指针表达缺省,等价于全零 bias;输入与输出均支持非连续 Tensor(输入由接口层自动连续化,输出由接口层按声明的布局逐元素写回)。

7. 调用方式

调用方式 说明
aclnn API 通过 aclnnThnnFusedGruCell 接口调用,两段式:aclnnThnnFusedGruCellGetWorkspaceSize + aclnnThnnFusedGruCell
GE 图模式 通过算子 IR 定义(op_graph/thnn_fused_gru_cell_proto.h)构图调用

Origin(信息来源)

NA

Benefit / Necessity (价值/作用)

新增ThnnFusedGruCell算子

Design(设计方案)

likedislike
huangzhiyuanhuangzhiyuan成员
10 天前 添加了label:requirement
huangzhiyuanhuangzhiyuan成员
10 天前 关联了pull request:rnn/thnn_fused_gru_cell算子开发
yuning_chenyuning_chen成员
10 天前 将 h1234515 设为负责人
CANN-robotCANN-robot成员
9 天前 关闭了 issue
CANN-robotCANN-robot成员
8 天前 添加了label:resolved