【CANN训练营】+ 开源之星 + optimized_transducer算子开发分享
收藏回复举报
【CANN训练营】+ 开源之星 + optimized_transducer算子开发分享
发表于2026-01-22 13:26:06
0 查看

optimized_transducer算子背景介绍

算子功能:计算的是序列到序列映射的负对数似然损失,功能上完全对其torchaudio.functional.rnnt。

算子输入

参数名输入/输出/属性描述数据shape数据类型数据格式
logits输入输入序列(B,T,U,V)FLOAT32、FLOAT16ND
targets输入目标序列(B,U-1)INT32ND
logit_lengths输入每个输入序列的长度(B)INT32ND
logit_lengths输入每个目标序列的长度(B)INT32ND
loss输出损失结果(B)FLOAT32、FLOAT16ND
grad输出输入logits的梯度结果(B,T,U,V)FLOAT32、FLOAT16ND
blank属性blank的索引位置,默认为-1INT-
clamp属性梯度的范围,默认为-1表示不限制梯度DOUBLE-
fused_log_softmax属性输入是否经过log_softmax,默认为trueBOOL-
requires_grad属性是否需要计算梯度,默认为trueBOOL-

算子逻辑

参考https://github.com/csukuangfj/optimized_transducer 链接中的算法,算子逻辑包含4个阶段:概率计算ComputeLogProbs、正向推导ComputeAlpha、反向推导ComputeBeta、梯度计算ComputeGrad。

ComputeLogProbs实现逻辑

初始化:
    batch_size = targets.size(0)
    V = 词汇表大小
    
遍历每个批次 b (0 到 batch_size-1):
    T = 当前批次的输入长度 (logit_len_arr[b])
    U = 当前批次的目标长度 (target_len_arr[b])
    U_p1 = U + 1  # 添加空白前缀
    
    遍历每个时间步 t (0 到 T-1):
        遍历每个扩展位置 u (0 到 U_p1-1):
            # 当前在logits中的位置对应: (b, t, u, :)
            # 当前在denominator中的位置对应: (b, t, u)
            # 计算空白标签的概率
            blank_logit = logits[b, t, u, blank]
            denominator_value = denominator[b, t, u]
            blank_log_prob = blank_logit - denominator_value
            # 计算非空白标签的概率
            if u < U:  # 不是空白后缀位置
                target_label = targets[b, u]
                sym_logit = logits[b, t, u, target_label]
                sym_log_prob = sym_logit - denominator_value
            else:  # u == U, 空白后缀位置
                sym_log_prob = -inf  # 无效
            存储到 log_probs:
                log_probs[累计索引, 0] = blank_log_prob
                log_probs[累计索引, 1] = sym_log_prob

ComputeAlpha实现逻辑(ComputeBeta实现类似)

起始点: alpha(0, 0) = 0
第一列 (u=0): alpha(t, 0) = alpha(t-1, 0) + log_probs(t-1, 0).blank
第一行 (t=0): alpha(0, u) = alpha(0, u-1) + log_probs(0, u-1).symbol
内部点 (t > 0 u > 0):
alpha(t, u) = log_sum_exp(
    alpha(t-1, u) + log_probs(t-1, u).blank,      # 路径1: 接受空白标签
    alpha(t, u-1) + log_probs(t, u-1).symbol      # 路径2: 接受目标符号
)

ComputeGrad实现逻辑

定义:
    g = logits[t,u,v] + alpha[t,u] - denominator[t,u] - beta[0]
    beta_cur = beta[t,u]
情况1(最后一个空白):
    grad = exp(g + beta_cur) - exp(g)
情况2(普通空白):
    grad = exp(g + beta_cur) - exp(g + beta[t+1,u])
情况3(目标符号):
    grad = exp(g + beta_cur) - exp(g + beta[t,u+1])
情况4(其他):
    grad = exp(g + beta_cur)

AscendC实现和性能优化

Host侧设计

  1. tiling划分

从上面的算法不难看出对于每个batch而言,上面的4个计算batch内部是相互依赖的,而batch之间是相互独立的,一个batch的计算流程必须放在同一个core上面,所以分核的tiling策略是按照输入shape的batch大小就行划分。

  1. workspace申请

由于中间计算的结果,包括ComputeLogProbs中的logsum结果denominator和概率计算结果logprobs,ComputeAlpha中的alpha,ComputeBeta中的beta。在一般情况下这些中间结果不能保存在core内部的UB上,所以需要在Host侧预分配workspace的空间。这里需要注意由于每个core都会对自己的workspace进行SetValue操作,而多核的SetValue很容易出错,需要保证每个core操作的workspace的起始地址64B对齐。

Kernel侧设计及优化

Kernel侧的计算完全对其上面的4个计算步骤,其中core内部计算exp和log必须使用vector核的Exp和Log接口。

通过性能分析发现如果按照上面的执行顺序计算alpha和beta,那么每一次循环内部在计算log_sum_exp操作都需要调用一次Exp和Log接口,很显然这不能充分利用vector核的计算能力。

在计算alpha(t, u)时会依赖于alpha(t, u - 1)和alpha(t - 1, u),可以发现沿着对角线计算可以提高vector核的计算,对角线y = t + u上的alpha完全依赖于对角线y = t + u - 1上的alpha,这样优化将vector核接口调用次数从T * U次减少到T + U次,性能得到明显提升,beta计算完全同理。

我要发帖子