文件最后提交记录最后更新时间
17 天前
24 天前
24 天前
24 天前
24 天前
24 天前
24 天前
18 天前
README

ScatterPaKvCacheWithKScale

产品支持情况

产品 是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品 ×
Atlas A2 训练系列产品/Atlas A2 推理系列产品 ×
Atlas 200I/500 A2 推理产品 ×
Atlas 推理系列产品 ×
Atlas 训练系列产品 ×

功能说明

  • 算子功能:更新KvCache中指定位置的key和value,同时更新key的scale值。

  • 输入输出支持以下场景:

    • 场景一:

      key:[batch * seq_len, num_head, head_size]
      value:[batch * seq_len, num_head, head_size]
      key_cache:[num_blocks, num_head, block_size, head_size]
      value_cache:[num_blocks, num_head, block_size, head_size]
      slot_mapping:[batch * seq_len]
      key_scale:[batch * seq_len, num_head]
      key_scale_cache:[num_blocks, num_head, block_size, 1]
      cache_layout:"BNBD"
      

      其中key和value的dtype为FLOAT8_E5M2或FLOAT8_E4M3FN,key_scale和key_scale_cache的dtype为FLOAT。

      计算公式:

      对于每个token(i ∈ [0, num_tokens))和每个头(j ∈ [0, num_head)):

      block_idx = slot_mapping[i] // block_size
      block_offset = slot_mapping[i] % block_size
      
      key_cache[block_idx][j][block_offset][:] = key[i][j][:]
      value_cache[block_idx][j][block_offset][:] = value[i][j][:]
      key_scale_cache[block_idx][j][block_offset][0] = key_scale[i][j]
      

      其中:

      • num_tokens = batch * seq_len
      • block_idx:slot_mapping映射到的block索引
      • block_offset:block内的偏移量
  • Ascend 950PR/Ascend 950DT:仅支持场景一。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
key 输入 待更新的key值,当前step多个token的key。 FLOAT8_E5M2、FLOAT8_E4M3FN ND
value 输入 待更新的value值,当前step多个token的value。 FLOAT8_E5M2、FLOAT8_E4M3FN ND
key_cache 输入/输出 需要更新的key cache,当前layer的key cache。 与key保持一致 ND
value_cache 输入/输出 需要更新的value cache,当前layer的value cache。 与value保持一致 ND
slot_mapping 输入 每个token key或value在cache中的存储偏移。 INT32、INT64 ND
key_scale 输入 待更新的key scale值,当前step多个token的key scale。 FLOAT ND
key_scale_cache 输入/输出 需要更新的key scale cache,当前layer的key scale cache。 FLOAT ND
cache_layout 属性 表示key_cache和value_cache的内存排布格式。当传空指针或"BNBD"时,表示格式为[num_blocks, num_head, block_size, head_size]。 STRING -

约束说明

  • 确定性计算:
    • 默认确定性实现。
  • key、value、key_cache、value_cache的数据类型必须一致;
  • slot_mapping的取值范围[0, num_blocks*block_size-1],且slot_mapping内的元素值保证不重复,重复时不保证正确性;
  • key和value的前两维shape必须相同;
  • key_scale是两维tensor,shape为[batch * seq_len, num_head],尾轴可以不连续;
  • key_scale_cache是四维tensor,shape为[num_blocks, num_head, block_size, 1],最后一维必须为1,尾轴必须连续。

调用说明

调用方式 样例代码 说明
aclnn接口 test_aclnn_ScatterPaKvCacheWithKScale 通过aclnnScatterPaKvCacheWithKScale调用ScatterPaKvCacheWithKScale算子
图模式 test_geir_ScatterPaKvCacheWithKScale 通过算子IR调用ScatterPaKvCacheWithKScale算子