pypto_pro.language.simt.atomic_sub#

产品支持情况#

  • Ascend 950PR/Ascend 950DT:支持

  • Atlas A3 训练系列产品/Atlas A3 推理系列产品:不支持

  • Atlas A2 训练系列产品/Atlas A2 推理系列产品:不支持

功能说明#

以原子方式从目的操作数target中减去源操作数value。

函数原型#

pypto_pro.language.simt.atomic_sub(
    target: Scalar,
    value: Scalar,
) -> Scalar

参数说明#

参数

输入/输出

说明

target

输入

目的操作数,Scalar类型。必须直接传入Tile或Tensor的单元素下标访问表达式,例如ub_tile[0, 0]或gm_tensor[0, 0]。
- UB Tile:必须位于UB,使用ND,支持DT_INT32、DT_UINT32、DT_FP32。
- GM Tensor:必须为ND,支持DT_INT32、DT_UINT32、DT_FP32、DT_INT64、DT_UINT64。

value

输入

源操作数,Scalar类型,表示减数。数据类型必须与target一致;数值字面量按target的数据类型处理,整数目的操作数不接受浮点字面量。

约束说明#

只能在由@pypto_pro.language.vector_function(mode=”simt”)定义的SIMT入口函数或辅助函数中调用。

返回值说明#

返回更新前的target值,返回值类型与target一致。

调用示例#

import pypto_pro.language as pl

@pl.vector_function(mode="simt", max_threads=256)
def consume_quota(
    remaining: pl.Tensor[[1, 1], pl.DT_UINT32],
    requests: pl.Tensor[[1, 256], pl.DT_UINT32],
    old_remaining: pl.Tensor[[1, 256], pl.DT_UINT32],
) -> None:
    tid = pl.simt.linear_thread_idx()
    old_remaining[0, tid] = pl.simt.atomic_sub(remaining[0, 0], requests[0, tid])


@pl.jit()
def atomic_sub_kernel(
    remaining: pl.Tensor[[1, 1], pl.DT_UINT32],
    requests: pl.Tensor[[1, 256], pl.DT_UINT32],
    old_remaining: pl.Tensor[[1, 256], pl.DT_UINT32],
):
    with pl.section_vector():
        consume_quota[256](remaining, requests, old_remaining)