Description
问题概述
在同一个 vector kernel 中,向量单元先通过 vsts 写入 UB,随后 auxiliary scalar unit 通过 pto.load 读取同一地址。PTOAS 在相关位置插入了:
但设备结果表明,VST_VLD 不能保证向量 store 对 auxiliary scalar load 可见:scalar load 可能读到旧值、上一轮值或未初始化值。
需要确认并修复 PTOAS 的 barrier 选择或依赖分析,使以下 RAW 依赖得到正确保证:
vector vsts/vstore -> auxiliary scalar pto.load
从执行单元划分看,该场景需要覆盖 vector store 到 scalar load 的同步语义。建议评估 VST_LD,或使用 PTOAS 中语义等价的 barrier kind。
验证环境
- 芯片:Ascend950PR
- CANN:9.2.0
- PTOAS VMI:0.67
- PTOAS 验证分支:
tmp/vecscope-membar-verify
- PTOAS 验证提交:
e9c79dd67 + da07dfa97
- 设备:0
用例一:vector store 后的 scalar load
PTODSL(完整源码)
from ptodsl import pto
from tilelang.contrib.ptodsl.dcache_bypass import (
pto_read_gm_bypass_dcache as _tl_pto_read_gm_bypass_dcache,
pto_write_gm_bypass_dcache as _tl_pto_write_gm_bypass_dcache,
)
from tilelang.contrib.ptodsl.simt import (
scalar_div as _tl_scalar_div,
scalar_rsqrt as _tl_scalar_rsqrt,
simt_allreduce_max as _tl_simt_allreduce_max,
simt_allreduce_min as _tl_simt_allreduce_min,
simt_allreduce_sum as _tl_simt_allreduce_sum,
vectorize_binary_f32x2 as _tl_vectorize_binary_f32x2,
vectorize_unary_f32x2 as _tl_vectorize_unary_f32x2,
)
from ptodsl._ops import _coerce_i64 as _tl_coerce_i64
from ptodsl._surface_values import wrap_surface_value as _tl_wrap_surface_value
@pto.jit(name="main_kernel", kernel_kind="vector", target="a5", mode="explicit")
def main_kernel(A: pto.ptr(pto.f32, "gm"), B: pto.ptr(pto.f32, "gm")):
buf_dyn_shmem = pto.castptr(pto.const(0, dtype=pto.i64), pto.ptr(pto.ui8, "ub"))
pto.mte_gm_ub(A, pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 0, 600, nburst=(1, 600, 600))
pto.set_flag("MTE2", "V", event_id=0)
pto.wait_flag("MTE2", "V", event_id=0)
with pto.vecscope():
remaining = pto.const(0, dtype=pto.int64)
remaining = _tl_wrap_surface_value(_tl_coerce_i64(150, context="PTO local.var store"))
for q in range(0, 2, 1):
mask = pto.vmi.create_mask(remaining, size=128)
vload_0 = pto.vmi.vload(pto.addptr(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), q * 128), 0, size=128)
brc_1 = pto.vmi.vbrc(pto.f32(float.fromhex('0x1p+1')), size=128)
vmul_2 = pto.vmi.vmul(vload_0, brc_1, mask)
pto.vmi.vstore(vmul_2, pto.addptr(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), q * 128), 0, mask)
remaining = _tl_wrap_surface_value(_tl_coerce_i64(remaining - 128, context="PTO local.var store"))
mask_3 = pto.vmi.create_mask(128, size=128)
vld_red_4 = pto.vmi.vload(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 0, size=128)
vld_red_5 = pto.vmi.vload(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 128, size=128)
mask_6 = pto.vmi.create_mask(22, size=128)
vmax_7 = pto.vmi.vmax(vld_red_4, vld_red_5, mask_6)
vsel_8 = pto.vmi.vsel(mask_6, vmax_7, vld_red_4)
vcmax_9 = pto.vmi.vcmax(vsel_8, mask_3)
vbrc_red_10 = pto.vmi.vbrc(vcmax_9, size=128)
mask_11 = pto.vmi.create_mask(1, size=128)
pto.vmi.vstore(vbrc_red_10, pto.addptr(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 152), 0, mask_11)
remaining_1 = pto.const(0, dtype=pto.int64)
remaining_1 = _tl_wrap_surface_value(_tl_coerce_i64(150, context="PTO local.var store"))
for q_1 in range(0, 2, 1):
mask_1 = pto.vmi.create_mask(remaining_1, size=128)
vload_0_1 = pto.vmi.vload(pto.addptr(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), q_1 * 128), 0, size=128)
sload_1 = pto.load(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 152)
brc_2 = pto.vmi.vbrc(sload_1, size=128)
vmul_3 = pto.vmi.vmul(vload_0_1, brc_2, mask_1)
pto.vmi.vstore(vmul_3, pto.addptr(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), q_1 * 128), 0, mask_1)
remaining_1 = _tl_wrap_surface_value(_tl_coerce_i64(remaining_1 - 128, context="PTO local.var store"))
pto.set_flag("V", "MTE3", event_id=0)
pto.wait_flag("V", "MTE3", event_id=0)
pto.mte_ub_gm(pto.castptr(pto.const(0, dtype=pto.int64), pto.ptr(pto.f32, "ub")), B, 600, nburst=(1, 600, 600), l2_cache="naci")
VPTO IR
VPTO IR 关键片段(完整 IR 较长,此处保留相关部分)
%17 = pto.vdup %16, %3 {position = "LOWEST"} : !pto.vreg<64xf32>, !pto.mask<b32> -> !pto.vreg<64xf32>
pto.vsts %17, %10[%c0], %9 : !pto.vreg<64xf32>, !pto.ptr<f32, ub>, !pto.mask<b32>
pto.vsts %17, %10[%c64], %8 : !pto.vreg<64xf32>, !pto.ptr<f32, ub>, !pto.mask<b32>
pto.mem_bar "VST_VLD"
%23 = pto.load %1[%c152] : !pto.ptr<f32, ub> -> f32
设备现象
- 用例:
e150-l128-max-consume
- 随机输入:
max_diff=41.28215026855469
- 常量输入
1.0 时,期望输出为 4.0
- 重复运行可观察到未初始化值或上一轮结果
用例二:vector store 后读取标量并继续写回
PTODSL(完整源码)
from ptodsl import pto
from tilelang.contrib.ptodsl.dcache_bypass import (
pto_read_gm_bypass_dcache as _tl_pto_read_gm_bypass_dcache,
pto_write_gm_bypass_dcache as _tl_pto_write_gm_bypass_dcache,
)
from tilelang.contrib.ptodsl.simt import (
scalar_div as _tl_scalar_div,
scalar_rsqrt as _tl_scalar_rsqrt,
simt_allreduce_max as _tl_simt_allreduce_max,
simt_allreduce_min as _tl_simt_allreduce_min,
simt_allreduce_sum as _tl_simt_allreduce_sum,
vectorize_binary_f32x2 as _tl_vectorize_binary_f32x2,
vectorize_unary_f32x2 as _tl_vectorize_unary_f32x2,
)
from ptodsl._ops import _coerce_i64 as _tl_coerce_i64
from ptodsl._surface_values import wrap_surface_value as _tl_wrap_surface_value
@pto.jit(name="main_kernel", kernel_kind="vector", target="a5", mode="explicit")
def main_kernel(A: pto.ptr(pto.f32, "gm"), B: pto.ptr(pto.f32, "gm")):
buf_dyn_shmem = pto.castptr(pto.const(0, dtype=pto.i64), pto.ptr(pto.ui8, "ub"))
pto.mte_gm_ub(A, pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 0, 1024, nburst=(1, 1024, 1024))
pto.set_flag("MTE2", "V", event_id=0)
pto.wait_flag("MTE2", "V", event_id=0)
with pto.vecscope():
mask = pto.vmi.create_mask(128, size=128)
for q in range(0, 2, 1):
vload_0 = pto.vmi.vload(pto.addptr(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), q * 128), 0, size=128)
brc_1 = pto.vmi.vbrc(pto.f32(float.fromhex('0x1p+1')), size=128)
vmul_2 = pto.vmi.vmul(vload_0, brc_1, mask)
pto.vmi.vstore(vmul_2, pto.addptr(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), q * 128), 0, mask)
mask_3 = pto.vmi.create_mask(128, size=128)
vld_red_4 = pto.vmi.vload(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 0, size=128)
vld_red_5 = pto.vmi.vload(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 128, size=128)
mask_6 = pto.vmi.create_mask(128, size=128)
vmax_7 = pto.vmi.vmax(vld_red_4, vld_red_5, mask_6)
vsel_8 = pto.vmi.vsel(mask_6, vmax_7, vld_red_4)
vcmax_9 = pto.vmi.vcmax(vsel_8, mask_3)
vbrc_red_10 = pto.vmi.vbrc(vcmax_9, size=128)
mask_11 = pto.vmi.create_mask(1, size=128)
pto.vmi.vstore(vbrc_red_10, pto.addptr(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 256), 0, mask_11)
pto.store(pto.load(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 0), pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 256)
for q_1 in range(0, 2, 1):
vload_0_1 = pto.vmi.vload(pto.addptr(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), q_1 * 128), 0, size=128)
sload_1 = pto.load(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), 256)
brc_2 = pto.vmi.vbrc(sload_1, size=128)
vadd_3 = pto.vmi.vadd(vload_0_1, brc_2, mask)
pto.vmi.vstore(vadd_3, pto.addptr(pto.castptr(buf_dyn_shmem, pto.ptr(pto.f32, "ub")), q_1 * 128), 0, mask)
pto.set_flag("V", "MTE3", event_id=0)
pto.wait_flag("V", "MTE3", event_id=0)
pto.mte_ub_gm(pto.castptr(pto.const(0, dtype=pto.int64), pto.ptr(pto.f32, "ub")), B, 1024, nburst=(1, 1024, 1024), l2_cache="naci")
VPTO IR
VPTO IR 关键片段(完整 IR 较长,此处保留相关部分)
pto.vsts %16, %10[%c0], %8 : !pto.vreg<64xf32>, !pto.ptr<f32, ub>, !pto.mask<b32>
pto.vsts %16, %10[%c64], %9 : !pto.vreg<64xf32>, !pto.ptr<f32, ub>, !pto.mask<b32>
%11 = pto.load %1[%c0] : !pto.ptr<f32, ub> -> f32
pto.store %11, %1[%c256] : !pto.ptr<f32, ub>, f32
pto.mem_bar "VST_VLD"
%18 = pto.load %1[%c256] : !pto.ptr<f32, ub> -> f32
这里 %11 位于前一组 vsts 和后续 scalar 写回之间,设备读取到向量 store 之前的旧值。
- 用例:
e256-l128-max-writethenread
- 随机输入:
max_diff=1.125840187072754
- 常量输入
1.0 时,期望输出为 4.0,实测为 3.0
复现命令
ptoas --pto-arch=a5 --pto-backend=vpto --emit-vpto \
plainload-consume.pto -o plainload-consume.vpto.mlir
ptoas --pto-arch=a5 --pto-backend=vpto --emit-vpto \
plainload-writethenread.pto -o plainload-writethenread.vpto.mlir
上述命令分别对对应的 PTO 输入运行;设备数值结果见本文的验证环境和设备现象。
分析与修复建议
- 在 PTOAS 依赖分析中识别
vsts/vstore → pto.load 的跨执行单元 RAW 依赖。
- 确认
VST_VLD 的覆盖范围;当前结果表明它不能覆盖 vector store 到 auxiliary scalar load。
- 对该依赖插入
VST_LD,或使用语义等价且经过硬件确认的 barrier kind。
- 对
pto.store → pto.load 的 scalar RAW 依赖明确对应的同步规则。
- 保持 barrier 位于前序 store 之后、consumer scalar load 之前。
- 修复后重新生成 VPTO IR,并使用随机输入、常量输入和重复运行验证不再出现旧值、未初始化值或上一轮值。
本文以 vector store → auxiliary scalar load 作为复现实例,对应规范表格中的 VS 家族 VST_LD,只覆盖一种向量→标量依赖组合。建议 PTOAS 修复时按规范表格同时梳理其它指令/依赖组合,包括 VV 家族的 VV_ALL、VST_VLD、VLD_VST、VST_VST,VS 家族的 VS_ALL、VLD_ST、VST_ST,以及 SV 家族的 SV_ALL、ST_VLD、LD_VST、ST_VST,确认各组合的依赖分析、barrier 选择和插入位置。本文未对这些组合逐一提供设备复现证据。
Description
问题概述
在同一个 vector kernel 中,向量单元先通过
vsts写入 UB,随后 auxiliary scalar unit 通过pto.load读取同一地址。PTOAS 在相关位置插入了:但设备结果表明,
VST_VLD不能保证向量 store 对 auxiliary scalar load 可见:scalar load 可能读到旧值、上一轮值或未初始化值。需要确认并修复 PTOAS 的 barrier 选择或依赖分析,使以下 RAW 依赖得到正确保证:
从执行单元划分看,该场景需要覆盖 vector store 到 scalar load 的同步语义。建议评估
VST_LD,或使用 PTOAS 中语义等价的 barrier kind。验证环境
tmp/vecscope-membar-verifye9c79dd67+da07dfa97用例一:vector store 后的 scalar load
PTODSL(完整源码)
VPTO IR
VPTO IR 关键片段(完整 IR 较长,此处保留相关部分)
设备现象
e150-l128-max-consumemax_diff=41.282150268554691.0时,期望输出为4.0用例二:vector store 后读取标量并继续写回
PTODSL(完整源码)
VPTO IR
VPTO IR 关键片段(完整 IR 较长,此处保留相关部分)
这里
%11位于前一组vsts和后续 scalar 写回之间,设备读取到向量 store 之前的旧值。e256-l128-max-writethenreadmax_diff=1.1258401870727541.0时,期望输出为4.0,实测为3.0复现命令
上述命令分别对对应的 PTO 输入运行;设备数值结果见本文的验证环境和设备现象。
分析与修复建议
vsts/vstore → pto.load的跨执行单元 RAW 依赖。VST_VLD的覆盖范围;当前结果表明它不能覆盖 vector store 到 auxiliary scalar load。VST_LD,或使用语义等价且经过硬件确认的 barrier kind。pto.store → pto.load的 scalar RAW 依赖明确对应的同步规则。本文以
vector store → auxiliary scalar load作为复现实例,对应规范表格中的 VS 家族VST_LD,只覆盖一种向量→标量依赖组合。建议 PTOAS 修复时按规范表格同时梳理其它指令/依赖组合,包括 VV 家族的VV_ALL、VST_VLD、VLD_VST、VST_VST,VS 家族的VS_ALL、VLD_ST、VST_ST,以及 SV 家族的SV_ALL、ST_VLD、LD_VST、ST_VST,确认各组合的依赖分析、barrier 选择和插入位置。本文未对这些组合逐一提供设备复现证据。