开发者
下载
[object Object]

[object Object][object Object]undefined
[object Object]
  • 算子功能:该算子为AICPU算子,是aclnnSparseLightningIndexerKLLossGrad算子的前置算子。根据aclnnSparseLightningIndexerKLLossGrad算子的输入shape、layout、mask和压缩比例信息,计算并输出分核切分metadata。输出结果可作为aclnnSparseLightningIndexerKLLossGrad算子的metadataOptional输入,减少主算子tiling阶段对host array的访问。

    该算子不建议单独使用,建议与aclnnSparseLightningIndexerKLLossGrad算子配合使用,形成完整的工作流。

    1. 接收主算子的shape信息,包括batchSize、maxSeqLenQ、maxSeqLenK、numHeadsQ、numHeadsK、headDim、topk、layout和mask信息。
    2. 根据每个query对应的有效sparse长度估算负载,并将B/S1合轴后的任务均衡切分到可用AIC核上。
    3. 输出metadata后,后续作为aclnnSparseLightningIndexerKLLossGrad算子的metadataOptional输入使用。
[object Object]

每个算子分为,必须先调用"aclnnSparseLightningIndexerKLLossGradMetadataGetWorkspaceSize"获取workspace大小,再调用"aclnnSparseLightningIndexerKLLossGradMetadata"执行计算。

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

    [object Object][object Object]
  • 返回值:

    返回aclnnStatus状态码,具体参见

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

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

    [object Object]
  • 返回值:

    返回aclnnStatus状态码,具体参见

[object Object]
  • aclnnSparseLightningIndexerKLLossGradMetadata为确定性实现,确定性计算配置不会改变其输出规则。
  • B(Batch)表示输入样本批量大小。
  • BSND场景
    • 必传batchSize、maxSeqLenQ、maxSeqLenK和topk参数,以获取shape信息。
  • TND场景
    • 必传cuSeqLensQOptional、cuSeqLensKOptional和topk参数,以获取正确shape信息。
    • 当batchSize为0时,通过cuSeqLensQOptional的shape推导batch。
[object Object][object Object]
  • Batch取值规则
    • 如果batchSize大于0,优先使用batchSize。
    • 如果batchSize小于等于0,且layoutQOptional为TND,则通过cuSeqLensQOptional的shape推导batch。
    • 如果batchSize小于等于0,且layoutQOptional为BSND,则报错。
  • Seqlen取值规则
    • TND场景下,通过cuSeqLensQOptional和cuSeqLensKOptional计算每个batch的实际q/k长度。
    • BSND场景下,通过maxSeqLenQ和maxSeqLenK获取q/k长度。
  • layout约束
    • layoutQOptional必须为BSND或TND。
    • layoutKOptional支持BSND和TND,建议与layoutQOptional保持一致。
  • head约束
    • numHeadsQ、numHeadsK和headDim必须为正数。
    • numHeadsQ必须能被numHeadsK整除。
  • sparse约束
    • topk必须为正数。
    • cmpRatio取值范围为[0, 128]。
    • maskMode当前仅支持0和3。
[object Object][object Object]

metadata输出为INT32 Tensor,当前shape固定为(64,),字段布局如下。

[object Object][object Object][object Object]

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

[object Object]
[object Object][object Object]undefined