开发者
下载
[object Object][object Object][object Object]undefined
[object Object]
  • 算子功能:GRU的反向传播,计算正向输入input、权重params、初始状态hx的梯度。
  • 正向计算公式:
    • 重置门 rt=σ(Wirxt+bir+Whrht1+bhr)r_t= \sigma(W_{ir}x_t + b_{ir} + W_{hr}h_{t-1} + b_{hr})
    • 更新门 zt=σ(Wizxt+biz+Whzht1+bhz)z_t = \sigma(W_{iz}x_t + b_{iz} + W_{hz}h_{t-1} + b_{hz})
    • 候选隐藏状态 nt=tanh(Winxt+bin+rt(Whnht1+bhn))n_t = \tanh(W_{in}x_t + b_{in} + r_t \odot (W_{hn}h_{t-1} + b_{hn}))
    • 隐藏状态 ht=(1z)nt+ztht1h_t = (1-z) \odot n_t + z_t \odot h_{t-1}
  • 反向计算公式:
    • 上游总梯度:dht=dyt+dhnext\text{上游总梯度:} \quad dh_t = dy_t + dh_{next} \quad(dytdy_t为上层梯度,dhnextdh_{next}为t+1时刻传回的梯度)
    • 更新门梯度:dzt=dht(ht1h~t)zt(1zt)\text{更新门梯度:} \quad dz_t = dh_t * (h_{t-1} - \tilde{h}_t) * z_t * (1 - z_t)
    • 候选态梯度:dhh~t=dht(1zt)(1h~t2)\text{候选态梯度:} \quad dh_{\tilde{h}t} = dh_t * (1 - z_t) * (1 - \tilde{h}_t^2)
    • 重置门梯度:drt=dhh~tlinhh[2hidden_size:3hidden_size]rt(1rt)\text{重置门梯度:} \quad dr_t = dh_{\tilde{h}t} * lin_{hh}[2*hidden\_size:3*hidden\_size] * r_t * (1 - r_t)
    • 线性变换梯度拆分:\text{线性变换梯度拆分:} dlinih=[dzt;drt;dhh~t],dlinhh=[dzt;drt;dhh~trt]\quad dlin_{ih} = [dz_t; dr_t; dh_{\tilde{h}t}], \quad dlin_{hh} = [dz_t; dr_t; dh_{\tilde{h}t} * r_t]
    • 输入梯度(传给下层):dxt=WihT@dlinih\text{输入梯度(传给下层):} \quad dx_t = W_{ih}^T @ dlin_{ih}
    • 前一时刻隐藏态梯度(传给t-1):dhprev=WhhT@dlinhh+dhtzt\text{前一时刻隐藏态梯度(传给t-1):} \quad dh_{prev} = W_{hh}^T @ dlin_{hh} + dh_t * z_t
    • 权重/偏置梯度累加:\text{权重/偏置梯度累加:} dWih+=dlinih@xtT,dWhh+=dlinhh@ht1T\quad dW_{ih} += dlin_{ih} @ x_t^T, \quad dW_{hh} += dlin_{hh} @ h_{t-1}^T dbih+=dlinih.sum(dim=1),dbhh+=dlinhh.sum(dim=1)\quad db_{ih} += dlin_{ih}.sum(dim=1), \quad db_{hh} += dlin_{hh}.sum(dim=1)
[object Object]

每个算子分为两段式接口,必须先调用“aclnnGRUBackwardGetWorkspaceSize”接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用“aclnnGRUBackward”接口执行计算。

[object Object]
[object Object]
[object Object]
  • 参数说明:
[object Object]
  • 返回值:

aclnnStatus: 返回状态码,具体参见[aclnn返回码]。

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

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

    [object Object]
  • 返回值:

    aclnnStatus: 返回状态码,具体参见[aclnn返回码]。

[object Object]
  • 确定性计算:

    • aclnnGRUBackward默认确定性实现。
    • 支持FP16/FP32,所有输入的数据类型需保持一致
[object Object]

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

[object Object]