Skip to content

Compare

产品支持情况

  • Ascend 950PR/Ascend 950DT:支持
  • Atlas A3 训练系列产品/Atlas A3 推理系列产品:不支持
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品:不支持
  • Atlas 200I/500 A2 推理产品:不支持
  • Atlas 推理系列产品 AI Core:不支持
  • Atlas 推理系列产品 Vector Core:不支持
  • Atlas 训练系列产品:不支持

功能说明

头文件路径为:"basic_api/reg_compute/kernel_reg_compute_vec_cmpsel_intf.h"

逐元素比较两个RegTensor大小,根据模板参数指定的比较模式,将比较结果写入目的操作数MaskReg中对应比特位,如果比较后的结果为真,则输出结果的对应比特位为1,否则为0。

其中cmp由模板参数mode指定比较模式,取值为EQ/GT/LT/GE/NE/LE。

函数原型

C++
template <typename T = DefaultType, CMPMODE mode = CMPMODE::EQ, typename U>
__simd_callee__ inline void Compare(MaskReg& dst, U& srcReg0, U& srcReg1, MaskReg& mask)

参数说明

表1 模板参数说明

参数名描述
T操作数数据类型。
mode比较模式。支持如下取值:
• LT:小于
• GT:大于
• GE:大于或等于
• EQ:等于
• NE:不等于
• LE:小于或等于
U源操作数的RegTensor类型,例如RegTensor<half>,由编译器自动推导,用户不需要填写。

表2 参数说明

参数名输入/输出描述
dst输出MaskReg类型,目的操作数。
srcReg0、srcReg1输入源操作数。
类型为RegTensor
两个源操作数的数据类型需要与目的操作数保持一致。
mask输入源操作数元素操作的有效指示,详细说明请参考MaskReg

返回值说明

数据类型

表3 数据类型组合

srcReg0和srcReg1数据类型dst数据类型
int8_tMaskReg
uint8_tMaskReg
int16_tMaskReg
uint16_tMaskReg
halfMaskReg
bfloat16_tMaskReg
int32_tMaskReg
uint32_tMaskReg
floatMaskReg
int64_tMaskReg
uint64_tMaskReg

约束说明

  • 通过mask参数控制的未选中元素在目的操作数中被置零。
  • 操作数重叠约束:srcReg0和srcReg1可以是同一个RegTensor

调用示例

C++
// 示例:逐元素比较src0Addr和src1Addr的数据(EQ模式),再按比较结果掩码选取元素写入dstAddr
// 输入:src0Addr指向源操作数0(如half类型[1.0, 4.0, 3.0, 2.0, 5.0, 6.0, 7.0, 8.0, 1.0, 2.0, 9.0, 3.0, 4.0, 6.0, 5.0, 7.0, ...])
//       src1Addr指向源操作数1(如half类型[1.0, 3.0, 3.0, 5.0, 2.0, 6.0, 9.0, 4.0, 2.0, 5.0, 9.0, 1.0, 8.0, 3.0, 5.0, 6.0, ...])
//       对应位置关系:pos0等于(1.0==1.0)、pos1大于(4.0>3.0)、pos2等于(3.0==3.0)、pos3小于(2.0<5.0)...
// 输出:dstMask为比较结果掩码(EQ模式下[1, 0, 1, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, ...],等于位置为1,其余为0)
//       dstAddr存储按掩码Select选取的结果
template<typename T>
__simd_vf__ inline void CompareVF(__ubuf__ T* dstAddr, __ubuf__ T* src0Addr, __ubuf__ T* src1Addr, uint32_t count, uint32_t oneRepeatSize, uint16_t repeatTimes)
{
    AscendC::Reg::RegTensor<T> srcReg0;   // 源操作数0
    AscendC::Reg::RegTensor<T> srcReg1;   // 源操作数1
    AscendC::Reg::RegTensor<T> dstReg;    // 目的操作数,存储Select选取的结果
    AscendC::Reg::MaskReg mask;
    AscendC::Reg::MaskReg dstMask;        // 比较结果掩码,相等位置为1
    for (uint16_t i = 0; i < repeatTimes; i++) {
        mask = AscendC::Reg::UpdateMask<T>(count);  // 根据实际元素个数更新mask
        AscendC::Reg::LoadAlign(srcReg0, src0Addr + i * oneRepeatSize);  // 加载源操作数0
        AscendC::Reg::LoadAlign(srcReg1, src1Addr + i * oneRepeatSize);  // 加载源操作数1
        AscendC::Reg::Compare<T, AscendC::CMPMODE::EQ>(dstMask, srcReg0, srcReg1, mask);  // 逐元素比较是否相等,结果写入dstMask
        AscendC::Reg::Select(dstReg, srcReg0, srcReg1, dstMask);  // 按比较结果掩码选取元素
        AscendC::Reg::StoreAlign(dstAddr + i * oneRepeatSize, dstReg, mask);  // 将选取结果存回dstAddr
    }
}

免责声明:本站内容由 asc-devkit 仓 master 分支自动编译生成,属于持续开发版本,可能存在缺陷,仅供预览与参考。如需稳定及商用资料,请查阅官方 昇腾社区