Skip to content

[Bug] VST_VLD 未保证 vector store → auxiliary scalar load 的可见性 #1593

Description

@KurrinQu

Description

问题概述

在同一个 vector kernel 中,向量单元先通过 vsts 写入 UB,随后 auxiliary scalar unit 通过 pto.load 读取同一地址。PTOAS 在相关位置插入了:

pto.mem_bar "VST_VLD"

但设备结果表明,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 输入运行;设备数值结果见本文的验证环境和设备现象。

分析与修复建议

  1. 在 PTOAS 依赖分析中识别 vsts/vstore → pto.load 的跨执行单元 RAW 依赖。
  2. 确认 VST_VLD 的覆盖范围;当前结果表明它不能覆盖 vector store 到 auxiliary scalar load。
  3. 对该依赖插入 VST_LD,或使用语义等价且经过硬件确认的 barrier kind。
  4. 对 pto.store → pto.load 的 scalar RAW 依赖明确对应的同步规则。
  5. 保持 barrier 位于前序 store 之后、consumer scalar load 之前。
  6. 修复后重新生成 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 选择和插入位置。本文未对这些组合逐一提供设备复现证据。

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

bugSomething isn't working

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions