@jit_kernel(name, params)
def kernel(src_ptr: pto.ptr(dtype, "gm"), out_ptr: pto.ptr(dtype, "gm")):
src_ub = pto.castptr(pto.i64(0), pto.ptr(dtype, "ub"))
out_ub = pto.castptr(pto.i64(nbytes), pto.ptr(dtype, "ub"))
pto.mte_gm_ub(src_ptr, src_ub, 0, nbytes, nburst=(1, 0, 0))
pto.set_flag("MTE2", "V", event_id=0)
pto.wait_flag("MTE2", "V", event_id=0)
source = pto.vmi.vload(src_ub, pto.const(0, dtype=pto.index), size=vl)
if masked:
if mask_kind == "group_prefix":
predicate = pto.vmi.create_mask(active_lanes, size=vl, group=group)
else:
predicate = pto.vmi.create_mask(active_lanes, size=vl)
result = pto.vmi.vabs(source, predicate, pmode=pmode)
else:
result = pto.vmi.vabs(source)
pto.vmi.vstore(result, out_ub, pto.const(0, dtype=pto.index))
pto.set_flag("V", "MTE3", event_id=0)
pto.wait_flag("V", "MTE3", event_id=0)
pto.mte_ub_gm(out_ub, out_ptr, nbytes, nburst=(1, 0, 0))
测试参数如下:
dtype = f16
VL = 256
mask = all-off
active_lanes = 0
pmode = zero
all zero.
Component
PTO Dialect / ODS (include/PTO/IR)
Description
pto.vmi.vabs(source, mask, pmode="zero") 在 PTOAS lowering 过程中丢失了输入的 mask。
根据 VMI 语义,pmode="zero" 的预期行为为:
active lane:
result[i] = abs(source[i])
inactive lane:
result[i] = 0
但是当前 lowering 后的 VABS 实际表现为所有 lane 都执行了 abs:
result[i] = abs(source[i])
因此,inactive lane 没有被置零。
Reproduction (minimal)
@jit_kernel(name, params) def kernel(src_ptr: pto.ptr(dtype, "gm"), out_ptr: pto.ptr(dtype, "gm")): src_ub = pto.castptr(pto.i64(0), pto.ptr(dtype, "ub")) out_ub = pto.castptr(pto.i64(nbytes), pto.ptr(dtype, "ub")) pto.mte_gm_ub(src_ptr, src_ub, 0, nbytes, nburst=(1, 0, 0)) pto.set_flag("MTE2", "V", event_id=0) pto.wait_flag("MTE2", "V", event_id=0) source = pto.vmi.vload(src_ub, pto.const(0, dtype=pto.index), size=vl) if masked: if mask_kind == "group_prefix": predicate = pto.vmi.create_mask(active_lanes, size=vl, group=group) else: predicate = pto.vmi.create_mask(active_lanes, size=vl) result = pto.vmi.vabs(source, predicate, pmode=pmode) else: result = pto.vmi.vabs(source) pto.vmi.vstore(result, out_ub, pto.const(0, dtype=pto.index)) pto.set_flag("V", "MTE3", event_id=0) pto.wait_flag("V", "MTE3", event_id=0) pto.mte_ub_gm(out_ub, out_ptr, nbytes, nburst=(1, 0, 0)) 测试参数如下: dtype = f16 VL = 256 mask = all-off active_lanes = 0 pmode = zeroExpected behavior
all zero.
Actual behavior / error logs
Git commit
e32488c
Host platform
A5 Simulator
Target Ascend arch (if relevant)
None
PTOAS build level (if relevant)
None