ARTICLE DETAIL

资讯详情

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

CUTLASS 任务调度(TS)Schedules 完全指南:@schedule、domain_loop 与持久化调度的 Python DSL 实战

CUTLASS 任务调度(TS)Schedules 完全指南:@schedule、domain_loop 与持久化调度的 Python DSL 实战 CUTLASS 任务调度TSSchedules 完全指南schedule、domain_loop 与持久化调度的 Python DSL 实战【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass导读本文以 CUTLASS Python DSL 的 Task SchedulingTS框架中的 Schedules调度文档为核心系统讲解如何用schedule装饰的 Python 函数显式记录 GPU 内核的资源调用顺序包括domain_loop、work_tile_loop、首/末/周期迭代守卫、运行时条件执行、动态域、可跳过 tile 以及三种调度形态非持久化 / 静态持久化 / CLC 动态持久化。读完本文你将能依据仓库源码schedule_builder.py、resources.py写出可被 TS 静态验证器检查的正确调度并理解每种调度结构在底层生成的 Schedule 树节点形态。TS 框架的背景资源、任务、依赖图与验证生命周期可先参考 ts_introduction.rst资源与TaskLocalVariable的声明方式见 ts_resources.rst。一、什么是 Schedule记录而非执行在 TS 框架中**调度函数schedule function**是一个用schedule装饰的 Python 函数。它的本质是录制器而不是执行器调用被装饰函数时函数体内的资源方法调用并不会真正执行 GPU 操作不会发起 TMA、MMA 或内存搬运调用发生在 Python 侧返回的是一个Schedule对象——一棵由节点组成的调度树这个Schedule再以Task(schedule...)的形式传给Task使用。从源码看schedule装饰器schedule_builder.py在调用时为每个MemoryResource参数包装一个ResourceProxy追踪代理将代理上的方法调用录制为调度条目record_step通过作用域栈ScheduleBuilder维护Schedule/ConditionalBlock/DomainLoop/WorkTileLoop的嵌套结构返回builder.finalize()得到的完整调度树。调度必须遵守两条结构性规则每个调度至多一个domain_loop域循环每个调度至多一个work_tile_loop工作瓦片循环。违反规则会由ScheduleBuilder抛出ScheduleError。例如源码中明确校验work_tile_loop()必须是调度顶层块不能嵌套在domain_loop、条件块或其他work_tile_loop之内domain_loop不能被另一个domain_loop嵌套。二、声明一个 Schedule调度函数以参与调度的资源作为参数返回None。函数体内开发者调用这些资源上的 producer/consumer 方法以及专门的同步方法从而录制调度schedule def schedule_fn(input_gmem, output_gmem) - None: ... # record resource method calls on the parameters task Task(..., scheduleschedule_fn(input_gmem_res, output_gmem_res))调用schedule_fn时传入的实参就是该任务实际操作的资源对象。TS 会把这些资源与调度树绑定str(schedule)可以漂亮地打印整棵调度树含数据流路由边便于检查。结构骨架HEAD / work_tile_loop / domain_loop / TAIL源码模块 docstring 给出了 TS 调度的标准文法schedule_builder.py可选with work_tile_loop(wq) as wtl:进入持久化的逐瓦片执行wtl是WorkTileLoopProxy其唯一方法是wtl.skippable()work_tile_loop(wq, skip_if...)启用可跳过的瓦片迭代skippable()块之外的内容对正常瓦片和跳过瓦片都会执行块之内的内容只对正常瓦片执行因此domain_loop必须放在 skippable 上下文内让被跳过的瓦片只执行非跳过的 HEAD/TAIL 簿记如 PDL 与 WorkQueue 阶段省略work_tile_loop时直接用domain_loop(start, end, step)HEAD 条目在前TAIL 条目在后。一个综合的参考形态来自 schedule_builder.py 的 docstring 示例schedule def my_schedule(gmem_qkv, smem_q, smem_kv, tmem_c, wq): with work_tile_loop(wq) as wtl: # HEAD 条目每个瓦片的预取 smem_q.try_acquire() smem_q.acquire() with domain_loop(0, k_iters, 1) as d: with d.first_iter(): coord_q, coord_k gmem_qkv.compute_coords() smem_q.q_load(coord_q) smem_q.commit() smem_kv.try_acquire() smem_kv.acquire() smem_kv.k_load(coord_k) smem_kv.commit() # TAIL 条目每个瓦片的排空 tmem_c.commit() wq.try_wait() wq.wait() wq.get_and_advance_work_tile() wq.release()三、Domainsdomain_loop 的边界与展开domain_loop(start, end, step, *, unroll1)定义了对域range(start, end, step)的循环。当前迭代索引通过stage_info.loop_offset到达工作方法循环索引从StageInfo读取从不作为数据传递。各参数含义start、end、step——循环边界。start与step分别默认0和1与 Pythonrange一致对于逐瓦片动态域可传入Task方法作为边界而不是静态整数该方法会对每个工作瓦片被调用见下文动态域一节unroll——展开提示默认1不展开传None让编译器自行选择。源码层面对domain_loop的边界做了严格校验schedule_builder.py接受 1~3 个位置参数domain_loop(end)、(start, end)、(start, end, step)step必须非零step 0直接抛ScheduleError边界既可以是静态值也可以是可调用对象可调用边界必须满足恰好两个必需位置参数(self, work_tile_coord)的签名校验_validate_domain_loop_bound_fn且必须是在类上访问的Task方法不能是实例绑定方法。下面是一个网格步长grid-stride域调度每个线程处理按网格大小间隔开的索引schedule def schedule_fn(input_gmem: InputGmemResource, output_gmem: OutputGmemResource) - None: threads_per_block num_warps * 32 start bx * threads_per_block tx step gdimx * threads_per_block with domain_loop(start, num_entries, step, unrollunroll): res input_gmem.get_item() output_gmem.set_item(datares)注意threads_per_block、bx、tx、gdimx均为追踪期的 Python 值trace-time它们在录制阶段被求值并固化进调度的DomainLoop节点其 dataclass 字段start / end / step / unroll / body见 schedule_builder.py。四、工作方法之间的数据流TaskLocalVariable 令牌值通过TaskLocalVariable令牌在工作方法之间流动声明了returns的consumer 工作方法在被调度调用时产生一个令牌tokenproducer 工作方法把该令牌作为参数消费掉。在上面的 grid-stride 示例中input_gmem.get_item()返回res令牌output_gmem.set_item(datares)消费它。循环索引不是令牌——工作方法通过stage_info.loop_offset读取它。资源如何声明这些变量见 ts_resources.rst 中Task-Local Variables一节其核心模式为dataclass(kw_onlyTrue) class InputGmemResource(MemoryResource): ... item: cutlass.Constexpr[TaskLocalVariable] TaskLocalVariable.uninitialized() def __post_init__(self) - None: self.item TaskLocalVariable( dtypecutlass.Int16, defaultcutlass.Int16(0), docsInput element loaded for the current grid-stride iteration., ) consumer_work(returnsitem) cute.jit def get_item(self, stage_info: StageInfo) - cutlass.Int16: gid stage_info.loop_offset val cutlass.Int16(0) if gid self.num_entries: val self.source_tensor[gid] return val令牌机制的底层实现在调度树中体现为Route数据流边schedule_builder.py每个被路由的输入会在产生它的Step上追加一条Route(source变量, destination消费Step, destination_argument参数名)边。ScheduleBuilder还会做以下检查_bind_routes与_validate_routed_token每个路由输入必须传令牌None或非令牌值都会报错令牌必须携带变量没有returns的管道操作如acquire/commit也会返回令牌但其variable is None不可路由传给下游会抛错令牌必须新鲜变量被再次产生后旧令牌成为 stale 句柄再路由会被拒绝is_stale按对象身份比较一个consumer_work(returns(a, b, ...))声明多个输出时调用返回按声明顺序的一个令牌元组每个变量可独立路由coord_q, coord_k gmem_qkv.compute_coords() # returns(coord_q, coord_k) smem_q.q_load(coord_q) # 路由 coord_q smem_kv.k_load(coord_k) # 路由 coord_k产生但从未路由到下游的输出会触发UserWarning_warn_unconsumedWorkQueue 的瓦片状态变量被豁免。五、首次、末次与周期迭代d.first_iter()与d.last_iter()是上下文管理器它们内部的运算只在域循环的第一次或最后一次迭代执行。周期性的工作使用d.every(period, start0)它基于从 0 开始的迭代计数start, start period, start 2 * period, ...触发与循环具体的start、step无关。这些守卫块之外的代码在每次迭代都执行。用first iteration做一次性设置例如首次 acquire用periodic 守卫做节奏性工作例如每N个瓦片推进一次元数据窗口用last iteration做排空例如最终的 commit。当循环只运行一次迭代时该迭代既是首次也是末次因此 first/last 块都会执行若计数0匹配周期节奏对应的周期守卫也会执行其验证含义见 ts_validation.rst。schedule def guarded_schedule(smem, page_offsets) - None: with domain_loop(0, num_iters, 1) as d: with d.first_iter(): smem.try_acquire() with d.every(4, start0): page_offsets.advance() smem.acquire() smem.producer_work() with d.last_iter(): smem.commit()在源码层面first_iter()/last_iter()由DomainLoopProxy提供它们各自打开一个ConditionalBlock条件来自枚举BlockConditionFirstIter/LastIter/Skippable。校验器保证FirstIter/LastIter条件只能出现在domain_loop()之内否则抛错守卫句柄绑定到产生它的具体DomainLoop一旦该with块退出包括在后续兄弟domain_loop内调用会拒绝使用_check_active。六、通用条件执行when_true / when_falsewhen_true(condition)与when_false(condition)是数据相关的运行时条件的通用块开启器。区分原则是迭代派生的条件首次、末次、周期用域循环句柄的方法让每个守卫绑定到当前活动的domain_loop()数据相关的运行时条件用普通的、以consumer_work(returns...)声明输出的工作方法产生守卫值辅助方法auxiliary work很适合只计算守卫状态的场景。验证器会把运行时条件与自动派生的(resource, method, result)键关联当两个任务必须共享同一个运行时值时可以显式提供key。from cutlass.experimental.task_scheduling import when_true schedule def conditional_schedule(page_offsets, smem) - None: with domain_loop(0, num_iters, 1) as d: smem.acquire() needs_epilogue smem.needs_epilogue() with when_true(needs_epilogue): smem.epilogue() smem.commit()TS 对运行时条件有严格约束每个运行时条件结果都必须由TaskLocalVariable槽位支撑。同一个存储令牌既驱动运行时执行也驱动穷举式静态调度验证不存在仅验证专用的条件两个任务若要对同一个共享运行时值分支必须在when_true/when_false上传同一个key或复用同一个存储令牌布尔槽直接读取整数类槽只有当存储值为零时才为假当一条运行时指令产生多个守卫值时用consumer_work(returns(...))为每个结果声明一个TaskLocalVariable槽。调度只会记录一个工作步骤来存储所有返回值每个when_true/when_false块读取各自选中的存储结果而不会再次调用产生方法。七、持久化调度Persistent Scheduling持久化调度用work_tile_loop(wq)包裹重复工作其驱动源是WorkQueue。程序员有责任确保每个参与任务在同一个逻辑边界上等待、推进并释放队列。典型的簿记序列是wq.try_wait() wq.wait() wq.get_and_advance_work_tile() wq.release()从WorkQueue的源码看它支持两种模式静态持久化调度把队列状态降低为轻量级算术get_and_advance_work_tile在消费侧直接调用StaticPersistentTileScheduler.advance_to_next_work()每个 CTA 用本地算术取下一个瓦片无需专用调度任务CLC 动态持久化调度使用拥有ClcFetchAsync管线的WorkQueue一个专用调度任务驱动其 producerfetch侧每个数据任务从其中消费工作瓦片。此时get_and_advance_work_tile不做推进而是从每阶段响应缓冲区_get_stage_response_ptr多阶段队列每个阶段一个 16 字节 Int128 响应记录读取硬件响应解码 CTA 坐标与有效性位。WorkQueue的标准管道方法在调度中通过ResourceProxy暴露其中get_and_advance_work_tile()与fetch_work_tile()是显式命名方法consumer_work/producer_work简写在 WorkQueue 上被禁止见 schedule_builder.py 的说明。八、动态域Dynamic Domain当任何域循环边界并非对每个工作瓦片都相同、必须按瓦片在运行时计算时使用动态域——最常见的是域循环的上界。做法是提供Task子类实现get_domain_size(self, tile_coord)方法返回每个瓦片的上界把该 provider 作为domain_loop的end边界传入。其余边界也允许动态函数名可以任意源码签名校验只要求恰好两个位置参数(self, work_tile_coord)。下面的例子展示一个由 offsets 数组计算边界的变长瓦片class DynamicDomainTask(Task): def __init__(self, offsets, **kwargs): super().__init__(**kwargs) self._offsets offsets cute.jit def get_domain_size(self, tile_coord): return self._offsets[tile_coord[0] 1] - self._offsets[tile_coord[0]] schedule def main_schedule(src, dst, wq) - None: with work_tile_loop(wq): with domain_loop( tx, DynamicDomainTask.get_domain_size, threads_per_block, ): val src.load() dst.store(valval) wq.try_wait() wq.wait() wq.get_and_advance_work_tile() wq.release()注意传入的是在类上访问的方法DynamicDomainTask.get_domain_size而不是实例绑定方法——源码的_validate_domain_loop_bound_fn会拒绝实例绑定方法以保证self保持为显式参数、运行时以(self, work_tile_coord)调用。最终ScheduleResult.dynamic_domain属性为True见 schedule_builder.py。九、可跳过瓦片Skippable Tileswtwl.skippable()是一个上下文管理器其内部的运算只在未跳过的瓦片上运行是否跳过由skip_if谓词决定它外部的内容在每个瓦片上运行。用它包裹数据工作区域同时把 WorkQueue 簿记留在外部这样每个启动的 CTA 仍会推进队列。与动态域的区别动态域给循环一个计算边界的逐瓦片回调固定签名get_domain_size(self, tile_coord)skip_if给工作瓦片循环一个决定瓦片是否运行可跳过工作的逐瓦片谓词。skip_if接受多种形式一个WorkQueue方法或普通函数 / lambda形参可以是(work_queue, work_tile)也可以只接收(work_tile)。下面的例子只把行复制工作标记为可跳过队列簿记保持在可跳过区域之外schedule def copy_schedule(copy_res: MemoryResource, wq: WorkQueue) - None: with work_tile_loop( wq, skip_ifOversubscribedCopyWorkQueue.skip_work_tile_if ) as wtwl: with wtwl.skippable(), domain_loop(0, num_rows, 1): copy_res.copy_tile_row() wq.try_wait() wq.wait() wq.get_and_advance_work_tile() wq.release()源码中的相关强制约束schedule_builder.pyskippable()要求外层存在带skip_if的work_tile_loop否则报错skippable()块不能嵌套带skip_if时domain_loop()必须位于wtl.skippable()之内否则被跳过的瓦片会跳过整个数据工作违背设计意图WorkQueue 操作不能出现在skippable()块内——被跳过的瓦片仍必须运行 wait / advance / release 簿记record_step中的检查skip_if若绑定到其他WorkQueue 实例或在__mro__中不属于本队列类的方法都会被拒绝_check_skip_predicate普通函数 / lambda 直接接受。WorkQueue还内置了默认谓词skip_work_tile_ifresources.py读取skip_work_tile任务局部变量。十、转发上下文信息Constexpr 调用点常量工作方法可能需要依赖调用点的上下文。调度可以把这类值作为关键字专用的cutlass.Constexpr[...]参数转发声明在工作方法上声明关键字专用的cutlass.Constexpr[...]参数允许带默认值传递在调度调用点传一个字面量例如smem.load(slot_index1)。该字面量在该次调用中被捕获并在调度被追踪时原样转发进工作方法体。schedule def schedule_fn(input_gmem: InputGmemResource, output_gmem: OutputGmemResource) - None: frag0 input_gmem.load(slot_index0) output_gmem.store(fragfrag0, slot_index1) frag1 input_gmem.load(slot_index1) output_gmem.store(fragfrag1, slot_index0)两对load/store调用的是同一个方法只有编译期slot_index字面量不同每次调用都会录制一个绑定到该值的新条目。源码实现上Constexpr 参数被分类为_WorkInfo.constexpr_names与数据流参数routed_names严格分离schedule_builder.py数据流输入必须绑定令牌argumenttokenConstexpr 输入必须绑定编译期字面量——给它传令牌会抛ScheduleError必填的 Constexpr 参数无默认值若在调用点缺失会在构建期报错而不是等到追踪期才抛TypeError每个调用点捕获的字面量存入Step.constexpr_kwargs运行时自动转发到工作方法。由于值是cutlass.Constexpr工作方法体内部可以用cutlass.const_expr(...)分支编译器会把该分支折叠掉。每个不同的调用仍然占据独立的调度槽位打印出的调度中call-idx列按顺序编号见 ts_printout.rst。十一、捕获的控制流 vs 追踪期 Python调度的运行时结构只能用with上下文管理器表达domain_loop、work_tile_loop、wtwl.skippable()、d.first_iter()、d.last_iter()。schedule函数内部的普通 Pythonfor和if语句属于追踪期元编程trace-time metaprogramming它们必须是编译期已知的它们会被展开进录制的调度unroll它们不会变成运行时循环或运行时守卫。schedule def store_schedule(tmem_c, gmem_d, wq) - None: with work_tile_loop(wq): with domain_loop(0, num_k_tiles, 1): pass for subtile_idx in cutlass.range_constexpr(subtile_cnt): t2r_rmem tmem_c.load_subtile(subtile_idxsubtile_idx) gmem_d.store(t2r_rmemt2r_rmem, subtile_idxsubtile_idx)这里的for是编译期元编程因为subtile_cnt编译期已知它会为每次迭代录制一对load_subtile/store——每一对绑定各自的编译期subtile_idx与上一节转发上下文信息一致——而不是在捕获的调度中形成一个运行时循环。这正是cutlass.range_constexpr(...)与domain_loop的本质区别前者在追踪期展开后者留下一个DomainLoop节点成为内核中的真实循环。十二、三种调度形态Scheduling ShapesTS 支持三种主要调度形态1. 非持久化Non-persistent不使用WorkQueue。启动网格直接映射到逻辑瓦片。每个 CTA 上的任务运行一次调度体domain_loop遍历该瓦片的工作域。2. 静态持久化Static persistentwork_tile_loop(wq)基于一个用本地算术分配下一个瓦片的WorkQueue。没有专用调度任务每个任务只需在每个工作瓦片迭代结束时运行标准的队列尾声epilogue——等待瓦片、推进到下一个、释放队列with work_tile_loop(wq): with domain_loop(0, num_k_tiles, 1): ... # 任务的数据工作 wq.try_wait() wq.wait() wq.get_and_advance_work_tile() wq.release()在WorkQueue._get_and_advance_work_tile_impl的实现注释中特别说明静态持久化模式下没有专用调度 warpadvance_to_next_work在消费侧调用。这看起来有些反直觉但正是这种设计让每个任务在静态与动态两种队列下都可以使用完全相同的三段式尾部模式try_wait/wait/get_and_advance_work_tile/release。3. CLC 动态持久化Dynamic persistentWorkQueue拥有ClcFetchAsync管线一个专用调度任务从硬件获取工作瓦片。典型的调度任务本身不做数据工作它只acquire队列、fetch下一个瓦片、commit然后运行相同的 wait / advance / release 尾声schedule def scheduler_schedule(wq: WorkQueue) - None: with work_tile_loop(wq): wq.try_acquire() wq.acquire() wq.fetch_work_tile() wq.commit() wq.try_wait() wq.wait() wq.get_and_advance_work_tile() wq.release()所有其他数据任务使用与静态情形相同的标准尾声——try_wait/wait/get_and_advance_work_tile/release——因此一个数据任务体在静态队列与 CLC 动态队列下完全相同区别只在于是否存在调度任务以及队列的管线类型。在 CLC 动态模式下get_and_advance_work_tile从每阶段响应缓冲区读取瓦片信息支持多阶段 WorkQueuefetch_work_tile由调度 warp 作为 producer 发出瓦片获取请求数据任务作为 consumer 仅等待管线见 resources.py 附近及PipelineType.ClcFetchAsync相关实现。十三、把 Schedule 交给 Task收尾要点调度函数返回的Schedule直接传给Taskschedule mma_schedule(smem_ab, tmem_c, work_queue) task Task( scheduleschedule, ... # src_resources, dst_resources, warp_idx, num_warps )Tasktask.py把软件任务映射到一段连续的 warp 范围warp_idx起、共num_warps个 warp 执行该调度体并通过src_resources读取的资源与dst_resources写入的资源声明数据方向。它还可以声明num_registers寄存器预算参与 warp 组寄存器验证源码校验要求其取值在 8~256 之间且为 8 的倍数。完整的使用生命周期资源类声明 → 实例化 → 依赖图 → 捕获调度 → 建 Task →TaskManager.print_and_verify()验证 → 填充工作方法体并做正确性校验见 ts_introduction.rst 的Authoring Lifecycle一节。调度的穷举式静态验证deadlock / race / barrier 检查细节可进一步阅读 ts_validation.rst而资源上声明的 SMEM/TMEM 分配与布局约束见 ts_allocators.rst、管线类型与PipelineConfig见 ts_pipelines.rst。实践要点回顾schedule录制而非执行至多一个domain_loop与至多一个work_tile_loop迭代派生的守卫first / last / every绑定到当前域循环数据相关的运行时分支用when_true/when_false并由TaskLocalVariable槽位支撑持久化任务必须确保所有参与者在同一逻辑边界 wait / advance / release追踪期的普通 Python 控制流是编译期元编程不会成为运行时循环。牢记这些规则再配合TaskManager.print_and_verify()的早期静态检查绝大多数调度顺序、barrier 与死锁错误都能在启动 GPU 内核之前暴露出来。【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表