接口功能:对序列执行因果一维卷积,沿序列维度使用缓存数据(长度为卷积核宽减1)对各序列头部进行padding,确保输出依赖当前及历史输入;卷积完成后,将当前序列部分数据更新到缓存;在因果一维卷积输出的基础上,将原始输入加到输出上以实现残差连接。支持APC(Automatic Prefix Caching)、MTP(投机解码)、残差连接、原地更新等特性。相较于标准causal_conv1d算子,本算子新增APC缓存复用、PD混部、残差连接可选等功能。[object Object]
支持以下场景:
场景一(prefill场景):
[object Object]其中cu_seq_len为batch内所有变长序列拼接后的总长度。
场景二(prefill和decode混合场景):
[object Object]其中cu_seq_len为batch内所有变长序列拼接后的总长度。
场景三(decode场景 - 变长序列):
[object Object]其中state_len必须大于所有batch中最大的token个数加1。
场景四(decode场景 - 固定batch):
[object Object]
计算公式:
K是卷积核宽度(固定为3),L是原始序列长度,dim是特征维度。
缓存读取
缓存行索引:
Case 1:首次计算(numComputedTokens[batchId] == 0)
Case 2:投机解码模式(numAcceptedTokens存在)
Case 3:默认模式
缓存拼接
缓存更新
Offset裁剪
APC缓存填充(可选,APC模式下)
对每个chunk = 0, 1, ..., nBlockToFill - 1:
因果1维卷积
零填充重置(可选,当convMode == 1并且numComputedTokens不为空时)
残差连接(可选)
原地更新
每个算子分为,必须先调用[object Object]接口获取入参并计算所需workspace大小以及包含了算子计算流程的执行器,再调用[object Object]接口执行计算。
确定性计算:
- aclnnInplaceFusedCausalConv1d默认确定性实现。
输入shape限制:
- prefill场景:
- x支持2维[cu_seq_len, dim]。
- weight必须是2维[K, dim],其中K固定为3。
- conv_states必须是3维[..., K-1, dim],第0维大小不固定且大于等于batch,同时大于等于cache_indices总维度大小。
- query_start_loc必须存在。
- cache_indices为1维[batch, ]或2维[batch, maxNumBlocks],其中1维表示未开启APC,2维表示开启APC。
- cu_seq_len范围[batch, 1024 1024],dim范围[64, 16384]且是16的倍数,且两者乘积需满足[64 batch, 4G]。
- batch范围[1, 256],maxNumBlocks范围[1, 1024]。
- max_query_len > 8。
- prefill和decode混合场景:
- x支持2维[cu_seq_len, dim]。
- weight必须是2维[K, dim],其中K固定为3。
- conv_states必须是3维[..., K-1+m, dim],第0维大小不固定且大于等于batch,同时大于等于cache_indices总维度大小。
- query_start_loc必须存在。
- cache_indices为1维[batch, ]或2维[batch, maxNumBlocks],其中1维表示未开启APC,2维表示开启APC。
- cu_seq_len范围[batch, 1024 1024],dim范围[64, 16384]且是16的倍数,且两者乘积需满足[64 batch, 4G]。
- batch范围[1, 256],maxNumBlocks范围[1, 1024]。
- max_query_len > 8。
- decode场景(变长序列):
- x支持2维[cu_seq_len, dim]。
- weight必须是2维[K, dim],其中K固定为3。
- conv_states必须是3维[..., k-1+m, dim],第0维大小不固定且大于等于batch,同时大于等于cache_indices总维度大小。
- query_start_loc必须存在。
- cache_indices为1维[batch, ]或2维[batch, maxNumBlocks],其中1维表示未开启APC,2维表示开启APC。
- cu_seq_len范围[batch, batch*8],每个batch的seq_len范围为[1, 8]。dim范围[64, 16384]且是16的倍数,batch范围[1, 256],maxNumBlocks范围[1, 1024]。
- max_query_len范围[1, 8]。
- decode场景(固定batch):
- x支持3维[batch, m+1, dim]。
- weight必须是2维[K, dim],其中K固定为3。
- conv_states必须是3维[..., K-1+m, dim],第0维大小不固定且大于等于batch,同时大于等于cache_indices总维度大小。
- cache_indices为1维[batch, ]或2维[batch, maxNumBlocks],其中1维表示未开启APC,2维表示开启APC。
- m范围[0, 7],dim范围[64, 16384]且是16的倍数,batch范围[1, 256],maxNumBlocks范围[1, 1024]。
- max_query_len范围[1, 8],可为-1。
- prefill场景:
输入值域限制:
- query_start_loc是累计偏移量,取值范围[0, cu_seq_len],长度为batch+1,query_start_loc[i]表示第i个序列的起始偏移,query_start_loc[batch+1]表示最后一个序列的结束位置。
- blockSize 必须大于等于 2。
- blockIdxFirstScheduledToken、blockIdxLastScheduledToken、initialStateIdx、num_computed_tokens和cache_indices均存在时表示APC开启,且满足以下条件(i为batch的索引):
- cache_indices为2维
- initialStateIdx[i] <= blockIdxFirstScheduledToken[i]+1
- initialStateIdx[i] <= blockIdxLastScheduledToken[i]
- blockIdxFirstScheduledToken[i] <= blockIdxLastScheduledToken[i]
- blockIdxLastScheduledToken[i] < maxNumBlocks
- num_accepted_tokens分为None和非None,非None情况下长度为batch,每个元素取值不超过当前batch的seq_len-1且大于0。
- num_computed_tokens中每个元素取值大于等于0。
- cache_indices的取值范围为[0, conv_states.dim[0]-1],且元素均不能相等。
- max_query_len = batch中的最大seq_len。
- Pangu V2 模式(conv_mode = 1)下,num_computed_tokens不能为 None。
- 算子入参与中间计算结果,在对应数据类型(float16/bfloat16)下,数值均不会超出该类型值域范围。
- 算子输入不支持有±inf和nan的情况。