十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

TileLang Host-Side Tensor Checks 复现指南:用 10 个最小示例定位 Kernel 调用错误

TileLang Host-Side Tensor Checks 复现指南:用 10 个最小示例定位 Kernel 调用错误 TileLang Host-Side Tensor Checks 复现指南用 10 个最小示例定位 Kernel 调用错误【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelangTileLang 在编译生成 kernel 时会向宿主端host stub自动插入一套针对torch.Tensor/ DLPack 兼容对象的参数校验逻辑覆盖参数个数、指针类型、dtype、shape、strides、设备信息等维度从而在靠近调用点处提前、精准地暴露错误。本篇文章以 maint/host_checks/ 目录下的 10 个独立复现脚本为主线结合 tensor_checks.md 中完整校验规则带你掌握如何用最小化示例复现各类 host-side 校验错误、如何读懂错误消息以及如何通过get_host_source()快速定位并修复问题。什么是 Host-Side Tensor ChecksTileLang 的 kernel 入口基于 TVM FFI DLPack 协议构建当你把torch.Tensor或任意 DLPack 兼容对象传给编译后的函数时宿主端 stub 会在调用设备 kernel 之前自动校验参数。这套机制写在 docs/compiler_internals/tensor_checks.md 中主要动机有三点ABI 稳定性入口统一接收张量与标量类型信息不依赖 Python 层的动态检查更低开销把校验从 Python 下沉到 C 层避免解释器属性访问开销整体调用开销低于基于 pybind 的等价方案精准报错断言在靠近调用点处抛出消息会明确指出哪个字段校验失败。校验范围覆盖参数个数num_args与指针类型、每个张量的可空性nullability、秩ndim、dtype、shape、strides、byte_offset、设备类型与设备号以及标量参数的类型。复现脚本的组织方式maint/host_checks/ 目录包含 11 个文件10 个形如01_num_args_mismatch.py的独立脚本外加一个批量运行器run_all.py和公共工具common.py。前置条件具备 CUDA 能力的环境绝大多数脚本会编译一个 CUDA 目标 kernelPython 包torch与tilelang至少一块 CUDA 设备。其中08_device_id_mismatch.py需要两块 GPU单卡环境下脚本会打印[SKIP]提示并跳过。运行方式逐个运行例如python 01_num_args_mismatch.py python 02_pointer_type_error.py # ... 直至 python 10_scalar_type_mismatch.py或一次性运行并输出汇总python run_all.pyrun_all.py会把每个脚本的标准输出/标准错误保存到logs/目录文件名为script.out/script.err。从源码run_all.py可以看到其判定逻辑非零退出码视为PASS即错误被成功复现输出中包含[SKIP]视为SKIP退出码为 0 且无跳过标记则视为FAIL即未观察到预期错误最后打印汇总并统计各状态数量若存在FAIL脚本以非零码退出。公共工具common.py所有复现脚本共享 common.py 中的两个构造函数build_matmul_kernel(M, N, K, targetcuda)构造一个A[M,K] × B[K,N] → C[M,N]的 FP16 矩阵乘 kernel并通过out_idx[2]将第 3 个参数 C 标记为输出因此编译后的函数只接收(A, B)两个输入若目标为 CUDA 但环境无可用设备会抛出RuntimeError。build_scalar_check_kernel(targetcuda)构造一个形如scalar_check(x: T.int32, flag: T.bool())的标量校验 kernel。kernel 本体是标准的 TileLang GEMM 实现T.Kernel划分线程块、T.alloc_shared分配共享内存、T.Pipelined流水线搬运、T.gemm完成矩阵乘这与文档 tensor_checks.md 中的 Matmul ReLU 参考示例结构一致。需要说明的是run_all.py与各脚本均以当前目录为工作目录运行cwdstr(root)因此可以直接使用from common import ...这种同目录导入方式。十个复现脚本逐一解析下面按照脚本编号逐一说明每个脚本复现的错误类型、触发方式与预期报错信息。若无特别说明均假设目标为 CUDAdevice_type2张量约定为A: float16 [M, K]、B: float16 [K, N]、C: float16 [M, N]。01参数个数不匹配num_args mismatch01_num_args_mismatch.py 构造MNK256的 kernel 后只传一个参数fn(a)缺少b。由于out_idx[2]让适配器adapter只期待两个输入参数个数错误会在进入 host stub 之前由适配器层抛出ValueError报错信息会包含预期输入数与实际输入数的对比。脚本注释明确说明适配器层会在 host stub 之前抛出部分错误如输入个数错误错误消息已尽可能与 host 校验对齐——这一点也写在了目录 README 的 Notes 中。02期望指针却传入标量pointer type error02_pointer_type_error.py 把整数a 1传给期望张量的位置。该值经由适配器转发到 host 层而 host 层对每个参数的 FFI 类型要求是指针类型DLTensor/handle或合法的标量类型于是触发类似Expect buffer A_handle to be pointer or tensor的报错具体名字取决于 kernel 参数名。03秩不匹配ndim mismatch03_ndim_mismatch.py 传入a torch.empty((M, K, 1), ...)让 A 的运行时秩变为 3而编译期秩为 2。预期报错形如kernel.A_handle.ndim is expected to equal 2, but got mismatched ndim04dtype 不匹配04_dtype_mismatch.py 让 A 使用torch.float32而预期是float16。该脚本在调用前还特意执行了print(fn.get_host_source())用于展示如何直接检视自动生成的 host 源码含全部断言与最终设备 kernel 调用。预期报错kernel.A_handle.dtype is expected to be float16, but got incompatible dtype修复方式A A.to(torch.float16)或直接用正确 dtype 构造。05shape 约束不满足05_shape_mismatch.py 把 A 的第二维设为K1破坏编译期 shape 绑定。预期报错Argument kernel.A_handle.shape[i] has an unsatisfied constraint: ... expected当 shape 为符号symbolic时host 会在运行时绑定并在满足单个检查点只有一个未知量的前提下即时求解线性关系。例如文档 tensor_checks.md 给出的示例对A: T.Tensor((m,))、B: T.Tensor((mn,))、C: T.Tensor((n*k,))运行时可以强制len(B) m n、len(C) n * k这类跨张量约束。06strides 校验失败非连续张量06_strides_mismatch.py 通过a.t()转置得到非连续张量传入。host 层对 strides 的检查规则是若buffer_type AutoBroadcast允许strides NULL并由shape推导否则逐维检查strides NULL时由shape推导并比较连续张量需满足strides[-1] 1、strides[-2] shape[-1]等。转置后的 A 违反该约束预期报错Argument kernel.A_handle.strides[1] has an unsatisfied constraint: ... 1修复方式传入A_nc.contiguous()或在 kernel 中调整对布局的假设。07设备类型不匹配07_device_type_mismatch.py 把 CPU 张量传给 CUDA 目标 kernel。host 会断言device_type 目标后端报错消息包含 DLPack 代码对照表。预期报错kernel.A_handle.device_type mismatch [expected: 2 (cuda)] ...DLPack 设备类型编码错误消息中引用1CPU, 2CUDA, 7Vulkan, 8Metal, 10ROCM, 14OneAPI, 15WebGPU。修复方式把张量移动到 CUDA 设备。08设备号不匹配多 GPU 场景08_device_id_mismatch.py 在torch.cuda.device_count() 2时打印[SKIP] Need at least 2 CUDA devices to reproduce device_id mismatch.并直接返回只有双卡及以上环境才真正构造cuda:0与cuda:1上的张量。当多个张量参与时host 会断言它们的device_id一致预期报错Argument kernel.B_handle.device_id has an unsatisfied constraint: ... ...09NULL 数据指针09_null_data_pointer.py 演示的是高级场景直接向函数传None。脚本注释说明真正的 NULL 数据指针通常来自手工构造的 DLTensor/NDArray或外部框架传入未分配/已释放的存储常规torch.Tensor分配极少触发。由于 FFI 处理方式不同传None可能触发指针类型断言如Expect buffer name to be pointer or tensor也可能触发 host 侧的非 NULL 指针检查kernel.name is expected to have non-NULL data pointer, but got NULL标量参数校验与可空性规则10标量类型不匹配10_scalar_type_mismatch.py 使用build_scalar_check_kernel构造scalar_check(x: T.int32, flag: T.bool())然后依次触发两类错误fn(1.0, True) # x 是 float - Expect arg[0] to be int fn(1, 2.5) # flag 是 float - Expect arg[1] to be booleanhost 对标量校验的规则是T.int*系列要求整数报错Expect arg[i] to be intT.bool要求布尔值报错Expect arg[i] to be boolean。修复方式即传入正确类型scalar_check(1, True)。可空性nullability规则张量是否允许为 NULL取决于静态分析中该张量是否被函数体实际使用相关规则与示例完整收录在 tensor_checks.md 的 Nullability Rules and Examples 一节必须非 NULL被使用A[0] 1这类直接访问传None会报main.A_handle is expected to have non-NULL pointer常量真分支仍必须非 NULLsome_cond: bool True且分支内使用 A静态分析认为可达可空常量假分支静态不可达some_cond: bool False且分支内使用 A此时 A 静态不可达允许 NULL运行时条件仍必须非 NULLsome_cond是T.bool参数时运行时才知道条件静态分析无法证明 A 未被使用因此 A 不可空。此外host 还会检查byte_offset必须为 0非零即报错以保证寻址简单且对齐数据指针在张量被要求非空时必须非 NULL。dtype 容差规则dtype 校验按(code, bits, lanes)三元组匹配并带有容差float8_e4m3接受e4m3、e4m3fn、e4m3fnuzfloat8_e5m2接受e5m2、e5m2fnuzbool接受int8/uint8bits8lanes 相同、kDLBool(code6, bits1 或 8)以及任意bitwidth1lanes 必须匹配对打包位宽 dtype如Int(1)、Int(4)、UInt(4)跳过严格 dtype 检查。批量运行与结果判定使用run_all.py时需要注意它的判定语义脚本预期失败因此非零退出码被判定为 PASS错误成功复现退出码为 0 反而被视为 FAIL。输出中的[SKIP]对应 08 号脚本在单卡环境下的跳过。日志统一落在logs/目录便于在 CI 或本地排查时回溯每个脚本的实际输出。快速排错清单run_all.py 之外文档 tensor_checks.md 还提供了一张速查表这里整理为对照清单#校验项触发方式典型错误消息1参数个数缺参/多参num_args should be N; expected: num_args, got: N2指针类型标量传给张量参数Expect arg[i] to be pointer3秩ndim运行时秩 ≠ 编译期秩ndim is expected to equal R, but got mismatched ndim4dtype不匹配且不在容差集内dtype is expected to be dtype, but got incompatible dtype5shape破坏常量/符号绑定shape[i] has an unsatisfied constraint: ... expected6strides布局不匹配如转置/切片strides[j] has an unsatisfied constraint: ... expected7设备类型错误后端设备device_type mismatch [expected: code (name)]8设备号张量分散在不同 GPUdevice_id has an unsatisfied constraint: ... ...9数据指针要求非空却为 NULLis expected to have non-NULL data pointer, but got NULL10标量类型错误标量类型Expect arg[i] to be int/boolean排错技巧与 FAQ打印 host 源码print(fn.get_host_source())是排查的首选手段可以看到精确的断言以及期望值与实际值的字段对比也能借此确认符号 shape 的绑定/求解顺序tensor_checks.md 的 Troubleshooting Tips 一节。strides 问题对非连续张量调用.contiguous()或避免生成转置/切片布局破坏假设。设备对齐确保所有参与张量共享相同的device_type与device_id。dtype 对齐用.to(dtype)或直接以正确 dtype 构造张量注意float8与bool的容差。动态 shape确保跨张量线性关系在检查点可唯一确定同一时刻只有一个未知量。能否关闭校验文档 FAQ 明确不建议且通常不支持关闭。校验在 host 侧完成以保持 ABI 稳定并能在靠近设备调用的位置提前失败。开销是否明显校验本身只是分支与字段读取相比 Python 侧校验更快主导成本仍是 Python→C 边界总体上比在 Python 中等价检查更便宜。最后无论遇到哪类错误都可以遵循文档 Closing Notes 的建议对照 kernel 签名交叉检查 shape / strides / device / dtype 四项即可高效定位问题涉及复杂符号关系时先打印 host 源码确认绑定与求解顺序再相应调整运行时 shape 与布局。将 maint/host_checks/ 的 10 个脚本作为最小复现基线配合 tensor_checks.md 的规则说明即可把每次 kernel 调用失败快速收敛到具体字段形成可复用的排错流程。【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表