TileLang Carver 框架实战:基于 Tile 结构的调度提示推荐引擎
2026/9/16 23:36:06 网站建设 项目流程

TileLang Carver 框架实战:基于 Tile 结构的调度提示推荐引擎

【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang

导读

Carver 是 TileLang 内置的一个轻量级调度提示推荐框架,它通过融合硬件架构信息、用户定义的 tile 结构与内置启发式规则,自动生成并排序面向矩阵乘、逐元素变换、归约类算子的 tile 配置(tiling strategy / blocking scheme / scheduling hint)。本文将以 tilelang/carver/README.md 为核心,结合仓库源码讲解 Carver 的模板 API、Hint 数据结构、后端架构抽象以及如何将推荐结果适配到 Triton 等其他编译器,帮助你在 GPU、CPU 与加速器后端上快速获得可落地的分块调度方案。


一、Carver 是什么:为编译器生成"调度提示"的框架

在 TileLang 这类面向高性能内核的领域专用语言中,循环的划分方式(tile 结构)直接决定内核在具体硬件上的性能表现。手工枚举分块参数(block 大小、warp 数量、归约步长等)既繁琐又难以覆盖不同设备的约束条件。Carver 正是为此设计:它统一生成面向多后端的 tile 候选集,并在生成时纳入硬件约束(如 CUDA 共享内存容量smem_cap、warp 大小、CPU 缓存结构、可用张量指令等),最终输出一组带排序的调度提示。

从实现上看,Carver 的核心 API 定义在 tilelang/carver/init.py 中,对外暴露了:

  • 架构抽象:CUDACDNARDNA(以及CPUMetal);
  • 操作模板:MatmulTemplateGEMVTemplateElementwiseTemplateGeneralReductionTemplateFlashAttentionTemplate
  • 底层调度生成器:roller(含policybestfitrasterization等子模块)。

典型用法只需三步:创建架构对象 → 构造模板 → 调用recommend_hints(topk)获取推荐配置。


二、快速上手:GeneralReductionTemplate 与 SSR 结构

2.1 基础用法

对于通用的循环嵌套,Carver 提供了GeneralReductionTemplate。它接受一个由S(Spatial,空间轴)与R(Reduce,归约轴)组成的结构字符串,以及对应的各维形状:

from tilelang import carver from tilelang.carver.arch import CUDA # 实例化 RTX 4090 的 CUDA 设备对象 arch = CUDA("nvidia/geforce-rtx-4090") # 描述如下循环嵌套: # for i in Spatial(1024): # for j in Spatial(1024): # for k in Reduce(1024): # ... carve_template = carver.GeneralReductionTemplate( structure="SSR", shape=[1024, 1024, 1024], dtype="float16", ).with_arch(arch) # 生成前 20 个 tile 候选(即调度提示) hints = carve_template.recommend_hints(topk=20) for hint in hints: print(hint)

输出示例(截断)如下:

{ 'block': [1, 128], 'thread': [1, 128], 'rstep': [64], ... }, { 'block': [2, 64], 'thread': [2, 64], 'rstep': [64], ... }, ... { 'block': [1, 16], 'thread': [1, 16], 'rstep': [512], 'reduce_thread': [8], ... }

2.2 SSR 结构如何映射为计算

GeneralReductionTemplateinitialize_function会按结构字符串逐轴解析:S轴进入输出空间形状,R轴被构造为te.reduce_axis,最终通过te.compute构建一个带归约的 TVMPrimFunc(见 tilelang/carver/template/general_reduce.py)。源码中的关键校验包括:

  • structureshape必须同时提供且长度一致;
  • shape各维必须为正整数;
  • 结构字符串只允许S/R(大小写均可),否则抛出ValueError

一个由 S 与 R 组成的 tile 结构可以模拟大量场景:SS表示二维逐元素操作,SSR则可以表示一次通用的矩阵乘法。因此该模板是快速验证调度思想的通用入口。

2.3 with_arch 与自动推断架构

with_arch(arch)由基类BaseTemplate实现(见 tilelang/carver/template/base.py),它把架构写入模板的_arch字段并返回自身以支持链式调用。值得注意的是,BaseTemplate_arch字段默认通过auto_infer_current_arch自动推断,也就是说即使不显式调用with_arch,模板也会尝试探测当前运行环境的设备架构。recommend_hints(topk)本质上是对get_hardware_aware_configs(self._arch, topk)的封装,后者由各子类实现并调用统一的get_roller_hints_from_func进入 roller 调度生成管线。


三、MatmulTemplate:矩阵乘的专用模板

针对C = A * B这类标准矩阵乘法,Carver 提供了更精细的MatmulTemplate,可自动推断线程块、warp 划分以及是否启用 Tensor Core 等策略:

from tilelang import carver from tilelang.carver.arch import CUDA arch = CUDA("nvidia/geforce-rtx-4090") carve_template = carver.MatmulTemplate( M=1024, N=1024, K=1024, in_dtype="float16", accum_dtype="float16", out_dtype="float16", ).with_arch(arch) # 获取描述该矩阵乘的(符号化)函数 func = carve_template.equivalent_function() print("Equivalent Function:\n", func) # 生成提示 hints = carve_template.recommend_hints(topk=20) for hint in hints: print(hint)

输出示例:

{ 'block': [32, 64], 'warp': [16, 32], 'rstep': [128], 'use_tc': True, ... }, { 'block': [64, 32], 'warp': [32, 16], 'rstep': [128], 'use_tc': True, ... }, ... { 'block': [256, 32], 'warp': [128, 16], 'rstep': [32], 'use_tc': True, ... }

3.1 模板参数说明

根据 tilelang/carver/template/matmul.py 的源码,MatmulTemplate的完整参数如下:

参数类型默认值含义
Mint必填矩阵 A 与 C 的行数
Nint必填矩阵 B 与 C 的列数
Kint必填A 的列数 / B 的行数(归约维)
trans_AboolFalse乘法前是否转置 A
trans_BboolTrue乘法前是否转置 B
in_dtypestr"float16"输入矩阵数据类型
out_dtypestr"float16"输出矩阵数据类型
accum_dtypestr"float16"累加中间结果数据类型
with_biasboolFalse是否叠加偏置项

initialize_function内部要求MNK均为正整数(否则断言失败),并根据trans_A/trans_B计算输入与权重的实际形状:input_shape = (M, K)(A 转置时为(K, M)),weight_shape = (K, N)(B 转置时为(N, K))。计算完成后,若with_bias=True会追加C[i, j] + Bias[j],若out_dtype != accum_dtype则会插入一次类型转换节点。

3.2 equivalent_function:拿到可调度的符号函数

equivalent_function()返回模板构建的PrimFunc,即后续可供 roller 策略分析、也可直接交给调度器使用的符号化计算函数。这一能力让模板不仅是"提示生成器",同时还能作为内核的原型定义,便于在 TileLang / TVM 生态内继续做进一步的调度变换。


四、Hint:调度提示的数据结构

无论是GeneralReductionTemplate还是MatmulTemplate,返回的每个 hint 都是一个Hint对象(见 tilelang/carver/roller/hint.py),其核心字段如下:

字段含义
block线程块(block)各空间轴的 tile 大小
thread不使用 Tensor Core 时各轴的线程划分
warp使用 Tensor Core(MFMA/MMA)时各轴的 warp 划分
rstep归约轴的步长(每次加载多少 K 数据)
reduce_thread归约轴上额外分配的线程数
use_tc是否启用 Tensor Core
vectorize各张量加载时的向量化宽度(如{'A_reindex': 8, 'B_reindex': 8}
pipeline_stage软件流水线级数(默认 1)
split_k_factorSplit-K 因子,用于 SM 浪费优化(TileLang 专属)
rasterization_plan光栅化(block 映射)策略
output_strides输出张量的 stride 信息

Hint.to_dict()在输出时会做精简:use_tc为真时输出warp否则输出thread;只有reduce_thread的乘积大于 1、vectorize非空、pipeline_stage != 1等条件下才会带上相应字段。这也解释了为什么不同 hint 打印出来的键并不完全一致。


五、支持的架构与扩展方式

5.1 开箱即用的后端

Carver 目前为以下后端提供开箱即用支持:

  • CUDA:如arch = CUDA("nvidia/geforce-rtx-4090")
  • CDNA(AMD GPU 类后端);
  • CPU
  • 另有RDNAMetal架构类位于 tilelang/carver/arch/ 目录下。

新增一种架构,只需实现TileDevice的一个子类(或提供自定义 target),描述清楚以下约束即可:

  • 共享/本地内存容量(smem_capmax_smem_usage);
  • warp(或向量)大小(warp_size);
  • 缓存大小(l2_cache_size_bytes等);
  • 可用的张量指令(available_tensor_instructions)。

TileDevice基类(见 tilelang/carver/arch/arch_base.py)统一定义了这些字段,并声明了必须实现的get_avaliable_tensorintrin_shapes

5.2 CUDA 后端内部结构

以下是 CUDA 后端的示意性代码(节选自 tilelang/carver/arch/cuda.py):

class CUDA(TileDevice): def __init__(self, target: Union[tvm.target.Target, str]): ... self.platform = "CUDA" # 设备约束 self.smem_cap = device.max_shared_memory_per_block self.compute_max_core = device.multi_processor_count self.warp_size = device.warp_size ... self.transaction_size = [32, 128] # 字节 self.bandwidth = [750, 12080] # MB/s,近似值 self.available_tensor_instructions = None def get_avaliable_tensorintrin_shapes(self): self.available_tensor_instructions = ( TensorInstruction("mma", [16, 16]), TensorInstruction("wmma", [16, 16]), ) return [t.shape for t in self.available_tensor_instructions] def __repr__(self): return f"CUDA({self.target})"

在实际实现中,CUDA构造器还会通过正则解析 SM 架构字符串(如sm_90sm_90a均解析为计算能力 90),并据此暴露一系列能力判定函数:is_volta_arch(sm 70–79)、is_ampere_arch(sm 80–88)、is_ada_arch(sm 89)、is_hopper_arch(sm 90)、has_mma_support(sm >= 80)。每个代际支持的张量核心精度矩阵也被硬编码在源码中,例如:

  • Volta:(float16, float32)(float16, float16)
  • Ampere:在 Volta 基础上新增(bfloat16, float32)(int8, int32)(int4, int32)等;
  • Ada:进一步加入(float8_e5m2, float32)(float8_e4m3, float32)
  • Hopper:与 Ada 一致。

is_tensorcore_supported_precision(in_dtype, accum_dtype, arch)正是依据这些矩阵判断某组输入/累加精度是否支持 Tensor Core——这解释了MatmulTemplate输出中use_tc字段是如何被决定的。

5.3 CDNA 与 CPU

  • CDNA(tilelang/carver/arch/cdna.py)通过tvm.runtime.rocm(0)获取设备信息,并针对 gfx950(CDNA4 / MI350)做了 160 KB LDS 的特殊处理(若驱动报告的默认值小于 163840 字节则覆盖之);
  • CPU(tilelang/carver/arch/cpu.py)实现较为轻量,注释指出 LLVM 后端本身无需精细调优,仅保持接口一致性。

六、将 Hint 适配到其他编译器(以 Triton 为例)

Carver 推荐结果的一大价值在于跨编译器适配。假设拿到如下 hint:

{ 'block': [32, 64], 'warp': [16, 32], 'rstep': [128], 'use_tc': True, 'vectorize': {'A_reindex': 8, 'B_reindex': 8} }

Triton中可以这样解读:

  • block_m = 32, block_n = 64, block_k = 128
  • 潜在 warp 划分warp_m = 16, warp_n = 32
  • vectorize:加载数据时使用向量宽度 8;
  • use_tc为真,在支持的情况下优先使用 Triton 的TensorOps(Tensor Core)。

这样即可快速测试多组配置,而无需手工猜测参数组合。同样的思路可以推广到 TVM、TileLang 或其他领域专用编译器——Carver 输出的 Hint 是一份"与后端无关"的调度蓝图,各后端只需定义自己的映射规则。


七、支持的模板一览

Carver 通过模板抽象了常见的循环模式,当前内置模板包括:

模板适用场景核心构造参数
GeneralReductionTemplate通用Spatial-Spatial-Reduce(SSR)等结构structureshapedtype
FlashAttentionTemplate类 Attention 操作(带 flash 内存访问模式)batch_sizenum_headshead_dimseq_lengthseq_kv_lengthis_causal、各 dtype
MatmulTemplate标准矩阵乘C = A * BMNKtrans_Atrans_B、三种 dtype、with_bias
GEMVTemplatey = Axy = xA类操作NKtrans_B、三种 dtype、with_bias
ElementwiseTemplate逐元素 / 逐点变换shapedtype

补充说明(依据对应源码):

  • ElementwiseTemplate(tilelang/carver/template/elementwise.py)内部以B = A + 1构造计算图,用于探索纯逐元素内核的最优分块与向量化;
  • GEMVTemplate(tilelang/carver/template/gemv.py)固定M = 1,描述一次矩阵向量乘,输出形状为(N,)
  • FlashAttentionTemplate(tilelang/carver/template/flashattention.py)将 QK^T 与 SV 两次矩阵乘建模为两个PrimFuncNode,并通过边(Edge)串联成计算图,再经TensorCorePolicy从输出节点统一发射配置;
  • 仓库中还提供了卷积模板 tilelang/carver/template/conv.py,可用于探索卷积类的调度提示。

如果你的算子有独特的循环结构或约束,完全可以仿照以上模板自定义新的专用模板,例如针对卷积、Flash Attention 的变体等——只需实现initialize_functionget_hardware_aware_configs两个抽象方法(见基类 tilelang/carver/template/base.py)。


八、底层原理:roller 策略管线

recommend_hints的最终执行路径收敛到 tilelang/carver/utils.py 的get_roller_hints_from_func:它首先将模板持有的PrimFunc交给DefaultPolicy(启发式最小化内存流量、最大化并行度),随后尝试通过get_tensorized_func_and_tags识别可张量化(Tensor Core)的子图;一旦识别成功,便改用TensorCorePolicy发射含use_tcwarp等字段的配置。

DefaultPolicy.emit_config(见 tilelang/carver/roller/policy/default.py)会经过"计算基础 tile → 分配归约步长rstep→ DFS 枚举共享内存 tile 候选 → 依据内存流量、共享内存占用、每 SM 的 block 数、wave 数等指标排序"的完整流程,最终截取topk个最优结果。每个Hint还附带trafficsmem_costblock_per_SMnum_wavegrid_size等评估量(存于TileDict),为后续筛选提供量化依据。此外,roller 还内置了rasterization(block 到 SM 的映射,如 panel 光栅化)与bestfit(best-fit 分配)机制,进一步细化 hint 的完整度。


九、路线图:与 TileLang 的端到端集成

Carver 的官方 TODO 清单中明确列出了一项计划:

  • 适配 tile language:为 TileLang 提供现成的调度调用或封装器(wrapper),打通端到端集成。

当前 Carver 已经作为 tilelang/carver/ 模块随 TileLang 源码发布,但其产出仍以"通用 Hint 字典"为主;未来若补齐面向 TileLang 自身调度语法的直接映射(例如把block/warp/rstep/split_k_factor/pipeline_stage等字段翻译成 TileLang 的T.parallelT.Pipelined、软件流水线等调度原语),即可实现"一次模板描述、自动产出可编译内核"的完整链路。对于开发者而言,当前阶段可以先将recommend_hints的结果作为搜索空间初始点,配合 TileLang 的自动调优(autotuning)工具进一步精化。


十、总结

Carver 以"模板描述算子结构 + 硬件约束建模 + 策略化搜索"的方式,把调度提示的生成从经验试错变成了可复现的工程流程:

  1. 统一 APIwith_arch(...)+recommend_hints(topk)即可跨后端获取排序后的分块方案;
  2. 硬件感知TileDevice子类携带共享内存、warp、缓存、带宽与张量指令信息,保证提示贴合实际设备;
  3. 结构可表达SSR等结构字符串与MatmulTemplate等专用模板覆盖从逐元素到注意力的大多数常见算子;
  4. 跨编译器可移植:Hint 是无后端依赖的中间表示,可映射到 Triton、TVM 与未来的 TileLang 调度封装。

想深入阅读源码,推荐从 tilelang/carver/template/base.py、tilelang/carver/roller/hint.py 与 tilelang/carver/utils.py 三个文件入手,它们构成了模板层、数据结构层与策略层的完整闭环。

【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询