0x00 概要

Fuser/compose_end_to_end.py 是 Fuser 管道中的最后一个关键步骤,它将分散的、针对特定子图优化的 Triton 内核无缝地整合成一个单一的、高性能的端到端 Triton 内核,同时确保其功能与原始 PyTorch 实现的数值等价性。

Composer的架构图如下,其功能概括是一句话:把所有验证通过的子图内核+原问题喂给LLM拼成单文件。也会把错误日志让LLM修,最多重试 max_iters轮,并自动 patch 常见 Triton 陷阱。

composer

0x01 核心功能

1.1 核心作用

Composer的核心功能如下:

  • 子图内核整合:将 Fuser 流程中拆分的子图及其对应的、已验证的 Triton 内核,重新组合为一个端到端的 Triton 实现,替代原始 PyTorch 代码的前向传播逻辑。
  • LLM 驱动的代码生成:以原始问题代码、子图信息、子图 Triton 内核为输入,通过定制化 Prompt 调用 LLM 生成完整的 Triton 内核代码。
  • 功能验证与迭代优化:支持自动验证生成的内核(对比 PyTorch 参考结果),若验证失败则基于错误信息迭代调用 LLM 修正代码,直到通过验证或达到最大迭代次数。
  • 约束保障:通过严格的代码规范(如必须包含 kernel_function、禁止 PyTorch 计算逻辑)和数值校验,确保生成的 Triton 内核可用、正确。

1.2 核心特色

特色方向 具体说明
严格的代码约束 强制要求生成的代码包含 kernel_function 顶层函数(与原始模型输入一致)、@triton.jit 内核、自测函数(输出 PASS/0 退出码);禁止内核中使用 PyTorch 计算逻辑(仅允许自测时对比)。
智能错误迭代 捕获编译 / 运行错误(stderr/stdout),构建精细化 Prompt 让 LLM 定位并修正问题(如 Triton 常见的 tl.broadcast 误用)。
自动补丁修复 内置 Triton 常见问题的文本补丁(如替换 tl.broadcast(0.0) 为 0.0),减少无意义的 LLM 迭代。
完整的日志留存 保存每一轮的 Prompt、生成的代码、验证结果,便于调试和追溯生成过程。
数值等价性保障 要求自测函数使用 allclose 校验数值(fp32: rtol≤1e-3/atol≤1e-3;fp16/bf16: ≤2e-2),确保 Triton 实现与 PyTorch 结果一致。

1.3 流程图

核心逻辑关系图

compose_end_to_end.py.逻辑关系图

完整执行流程图

compose_end_to_end.py.流程图

0x02 详细功能

2.1 使用

compose_end_to_end.py 会合成(compose)一个端到端的 Triton 内核,用来解决原始 KernelBench 问题。

compose_end_to_end.py 会将原始问题文件、子图分解信息 + 各张量形状(来自 subgraphs.json)、以及已生成的子图 Triton 内核作为输入,构建提示(prompt)发送给 LLM。然后LLM 会把这些碎片拼成一个语义与原始问题完全一致的完整内核,最终返回一个 Python 文件(/composed_kernel.py),里面提供:

  • 一个或多个使用 @triton.jit 装饰的 Triton 内核。
  • 一个名为 kernel_function(...) 的顶层 Python 包装函数,它接受与原始模型相同的输入张量,并协调 Triton 内核的执行,返回最终输出。
  • 一个自测函数(如 test_kernel 或 run_tests),该函数比较 Triton 实现的结果与原始 PyTorch 问题代码的参考结果,并在成功时打印 'PASS' 并退出。
  • 生成过程的元数据和验证结果会被记录在一个 JSON 格式的摘要文件中(如 composition_summary.json)。

compose_end_to_end.py 用法如下:

python -m Fuser.compose_end_to_end \
    --problem /abs/path/to/kernelbench_problem.py \
    --subgraphs /abs/path/to/subgraphs.json \
    --kernels-summary /abs/path/to/kernels_out/summary.json \
    [--model gpt-5] \
    [--out-dir ./compose_out] \
    [--verify]

可以通过 --verify 标志启用自动验证。在此模式下,每个 LLM 生成的组合尝试(无论是初始的还是经过修正的)都会被传递给 Fuser/runner.py 中的 run_candidate 函数来执行。

验证的成功与否取决于执行是否正常退出(exit code 0)并且输出中包含 'PASS' 字符串或 ALL_TESTS_PASSED 字符串。

compose_end_to_end_new

2.2 错误处理与迭代

如果 LLM 生成的第一个组合内核无法通过验证(运行 / 编译失败),该脚本会捕获错误信息(stderrstdout)。它会构建一个新的提示(_build_refinement_prompt),将错误信息作为上下文提供给 LLM,要求其修正代码。这个过程可以重复多次(max_iters 参数控制最大迭代次数),直到生成的内核通过验证或达到最大迭代次数。

2.3 _load_kernels_from_summary

_load_kernels_from_summary 为构建 prompt 提供 「代码素材」(各子图的有效 Triton 内核代码),具体而言,_load_kernels_from_summary从调度阶段生成的内核汇总 JSON 文件中,过滤并加载所有成功生成的有效子图 Triton 内核,校验数据格式与文件有效性,封装为标准化 KernelItem 对象列表,为 LLM 组合内核提供可直接复用的有效代码素材,过滤失败、无效的内核产物,避免无效素材干扰后续组合流程。

其特殊如下:

  • 多维度有效内核过滤:依次过滤「非列表格式汇总数据、非字典格式子项、标记为失败的内核、无 ID / 无内核路径的子项、内核文件不存在的子项」,仅保留全量校验通过的成功内核,从源头保证代码素材的有效性;

  • 关键字段强制校验:子图 ID(sid)和内核文件路径(kernel_path)为必选字段,缺失任一则直接过滤,确保每个有效内核都能关联到唯一子图且存在实际代码文件;

  • 标准化对象封装:将内核的子图 ID、文件路径、代码内容封装为 KernelItem 对象,而非原始字典 / 字符串,提升后续代码处理的可读性与可维护性;

  • 无有效内核直接终止:若汇总文件中无任何有效内核,直接抛出 SystemExit 异常终止流程,避免后续流程因无有效素材而无意义执行;

  • 兼容调度阶段输出格式:严格适配 dispatch 步骤生成的 summary.json 格式,实现上下游流程的无缝衔接。

2.4 Prompt

以下三个函数是 PyTorch KernelAgent 中LLM 生成端到端 Triton 内核的 Prompt 构建核心模块,为 LLM 提供标准化、高指导性、场景适配的精准输入提示,是连接「分散的子图 / 内核 / 问题数据」与「LLM 可理解的生成指令」的关键枢纽。其中

  • _summarize_subgraphs_for_prompt 为基础数据处理函数,负责将子图信息格式化;
  • _build_composition_prompt 和 _build_refinement_prompt 为双 Prompt 构建主函数,分别支撑首次端到端内核组合生成基于错误的迭代精修两大核心场景,通过严格的指令约束、完整的上下文信息、针对性的优化指导,确保 LLM 生成符合工程要求、硬件适配、可直接运行的 Triton 内核代码。
核心作用
  1. _summarize_subgraphs_for_prompt:作为基础支撑函数,将模型分解后的子图信息列表结构化、简洁化转换为文本摘要,提取子图 ID、类型、数据布局、数据类型、输入输出形状、核心算子等关键约束信息,按统一格式拼接为易读字符串,为两个 Prompt 构建函数提供标准化的子图信息描述,让 LLM 快速理解各子图的功能与计算约束。
  2. _build_composition_prompt:为 LLM 首次生成构建全量上下文、强约束的组合型 Prompt,整合原始 PyTorch 问题代码、子图信息摘要、各子图有效 Triton 内核代码、目标硬件平台配置四大核心信息,明确 LLM 的核心任务是融合子图内核生成端到端 Triton 实现,并制定严格的工程要求、硬件适配规则、Triton 开发规范,指导 LLM 完成从「分散子图」到「一体化内核」的组合与优化。
  3. _build_refinement_prompt:为 LLM 迭代优化构建错误导向、针对性的精修型 Prompt,在保留核心基础信息的前提下,新增前次生成的错误日志(stdout/stderr)、上一轮失败的代码实现,明确 LLM 的核心任务是基于错误信息定位问题并修正代码,同时追加更具体的错误修复要求,确保精修后的代码能解决编译 / 运行问题,且不违背原有工程规范。
_build_composition_prompt 专属特色
  1. 四大核心信息全量整合:完整融入「原始问题代码(需求基准)、子图摘要(逻辑约束)、子图内核(代码素材)、平台配置(硬件规则)」,让 LLM 既理解「要做什么(PyTorch 模型功能)」,又知道「有什么可用(子图内核)」,还明确「该怎么适配(硬件平台)」;
  2. 核心理念明确:融合与融合:明确要求 LLM 优先将多子图融合为尽可能少的 Triton 内核启动,在保证数值语义准确的前提下提升运行效率,贴合 Triton 内核「大粒度融合、减少内核启动开销」的优化核心理念;
  3. 超详细的硬性要求:制定 10 余项不可违背的硬规则,覆盖代码输出格式、设备张量管理、函数命名与封装、计算接口限制(禁用 PyTorch 计算)、数据格式与算子顺序、数值验证要求、导入与行为限制等,从源头规范代码生成;
  4. 实用的 Triton 开发指导:提供针对性的 Triton 实现技巧,包括子图融合的形状匹配、常量权重优化、内存访问规范(tl.load/tl.store 带掩码)、网格与分块设计,同时明确列出 Triton 常见开发陷阱及规避方法,降低 LLM 生成错误代码的概率;
  5. 数值等价性强制要求:明确要求生成的代码必须包含自测试函数,通过与 PyTorch 参考结果的对比验证数值正确性,并制定严格的误差容忍度(fp32/_fp16/bf16 区分),确保 Triton 内核的功能正确性与数值准确性。
_build_refinement_prompt 专属特色
  1. 错误导向的精准精修:将前次生成的 stderr/stdout 错误日志作为核心参考,让 LLM 聚焦于「定位问题→修复问题」,避免无目的的重生成,大幅提升迭代优化的效率与针对性;
  2. 保留原有约束,追加修复要求:明确「原有所有要求保持不变」,仅针对错误场景追加更具体的修复规则(如禁止 tl.broadcast 滥用标量),确保精修后的代码不违背原有工程规范,同时解决具体问题;
  3. 全量代码重生成,拒绝差分输出:强制要求 LLM 返回完整的修正后代码,而非代码差分或修改建议,避免后续代码拼接、整合的额外工作,确保输出可直接替换使用;
  4. 关键标识强制保留:明确要求保留顶层 kernel_function 函数名、自测试函数及「PASS」打印 / 退出码规则,确保精修后的代码能无缝对接后续的自动化验证流程,无需修改验证逻辑;
  5. 失败代码参考,定位问题更高效:将上一轮的失败代码完整融入 Prompt,让 LLM 能直接对比错误日志与代码实现,快速定位问题所在(如编译错误的行号、运行错误的逻辑),提升修复的准确性。
_summarize_subgraphs_for_prompt 专属特色

_summarize_subgraphs_for_prompt 提供 「逻辑约束」(各子图的功能、形状、布局等约束信息)。具体而言,_summarize_subgraphs_for_prompt:对模型分解后的子图信息列表进行结构化、简洁化的文本汇总,提取子图 ID、类型、数据布局、数据类型、输入输出形状、核心算子等关键信息,按统一格式拼接为易读的文本字符串,为构建 LLM 组合 Prompt 提供标准化的子图信息描述,让 LLM 快速理解各子图的功能、形状约束与计算要求。

  1. 关键信息精准提取,剔除冗余:仅保留 LLM 组合 / 精修内核所需的核心约束信息(ID、类型、布局、dtype、输入输出形状、核心算子),剔除无关冗余信息,减少 Prompt 的 Token 占用,提升 LLM 处理效率;
  2. 合理默认值兜底,保证鲁棒性:对数据布局(默认 NCHW)、数据类型(默认 float32)、算子列表(默认空列表)等易缺失字段设置合理默认值,避免因字段缺失导致 Prompt 构建失败,提升流程的容错性;
  3. 层级化紧凑格式,易读易解析:采用「一级行标注子图基础属性 + 二级行标注核心算子」的层级格式,既保证信息紧凑(适配 Prompt 长度限制),又结构清晰,让 LLM 能快速关联子图 ID 与对应的功能、约束;
  4. 算子信息长度控制,避免超限:将算子列表序列化后截取前 400 字符,避免因算子过多导致 Prompt 过长超出 LLM 上下文窗口,同时兼容 JSON 序列化失败的情况(降级为直接字符串截取),保证信息完整性;
  5. 形状信息灵活适配:优先使用 inputs 字段描述输入,无 inputs 则降级为 input_shape,适配不同子图分解工具的输出格式差异,保证输入输出形状信息的有效传递。
协同工作关系

三个函数形成 「基础数据处理→首次生成 Prompt 构建→迭代精修 Prompt 构建」 的层级支撑关系,为 LLM 生成端到端 Triton 内核提供全流程的 Prompt 支撑:

  1. 基础层_summarize_subgraphs_for_prompt 对原始子图信息做统一格式化处理,生成标准化子图摘要,为上层两个 Prompt 构建函数提供一致、易解析的子图约束信息,实现数据的一次处理、多次复用;
  2. 首次生成层_build_composition_prompt 基于标准化子图摘要,整合问题代码、子图内核、平台配置,构建全量约束的组合 Prompt,指导 LLM 完成首次端到端 Triton 内核生成,是「从无到有」的核心指导;
  3. 迭代精修层_build_refinement_prompt 复用标准化子图摘要与核心基础信息,新增错误日志和前次失败代码,构建错误导向的精修 Prompt,指导 LLM 完成「从错到对」的迭代优化,是提升代码有效性的关键;
  4. 闭环支撑:三个函数的输出共同支撑 KernelAgent 的「生成→验证→精修→再验证」闭环流程,确保每一轮 LLM 生成都有明确的指导、充足的依据、严格的约束,大幅提升端到端 Triton 内核的生成成功率与工程质量。
代码
def _build_composition_prompt(
    problem_code: str,
    subgraphs: list[dict[str, Any]],
    kernel_items: list[KernelItem],
    target_platform: PlatformConfig,
) -> str:
    """Create a single user message to instruct composition by the LLM.
    构建初始合成Prompt:
    为LLM生成结构化指令,引导其基于子图和参考内核,合成端到端的Triton算子代码
    参数:
    - problem_code: 原始KernelBench问题代码(PyTorch)
    - subgraphs: 子图信息列表(JSON格式)
    - kernel_items: 参考内核列表(包含子图ID和对应的Triton代码)
    - target_platform: 目标平台配置(如CUDA/XPU)
    返回值:完整的LLM用户指令字符串
    """
    # 第一步:生成子图摘要(压缩子图信息,控制Token消耗,便于LLM快速理解核心特征)
    sg_summary = _summarize_subgraphs_for_prompt(subgraphs)

    # 第二步:构建参考内核代码区块(仅保留核心代码,避免Token溢出)
    # 注释说明:暂时保留完整文件内容,调用方可根据模型窗口限制进一步裁剪
    # 初始化内核区块的文本片段列表
    kernels_section_parts: list[str] = []
    # 遍历每个参考内核项
    for ki in kernel_items:
        # 为每个子图内核构建带格式的代码片段(Markdown Python代码块)
        kernels_section_parts.append(
            f"### Subgraph {ki.subgraph_id}\n```python\n" + ki.code + "\n```\n"
        )
    # 拼接所有内核片段,形成完整的参考内核区块
    kernels_section = "\n".join(kernels_section_parts)
    
    # 第三步:获取平台专属的指导规则(如CUDA的内存访问规则、XPU的编译要求)
    platform_guidance = target_platform.guidance_block

    # 第四步:构建核心指导语(包含任务背景、平台信息、硬性要求、实现技巧)
    # 使用textwrap.dedent去除缩进,保证Prompt格式整洁
    guidance = textwrap.dedent(
        f"""
        You are given:
        - The original problem file (PyTorch module and helpers).
        - A decomposition of the model into fusable subgraphs with exact shapes.
        - Working Triton kernels generated for some subgraphs.

        TARGET PLATFORM: {target_platform.name}
        DEVICE STRING: {target_platform.device_string}
        {platform_guidance}

        Task:
        - Compose an end-to-end Triton implementation that matches the original
          model's forward pass for the provided shapes. You may inline, adapt,
          or reuse the given subgraph kernels. Prefer fusing into as few kernel
          launches as possible while preserving exact numerical semantics.

        Hard requirements:
        - Return ONE complete Python file only, fenced as a single ```python block.
        - Allocate inputs, weights, intermediates, and outputs on device='{target_platform.device_string}' and keep them there throughout forward/verification.
        - CPU is acceptable only for metadata, scalars, and export serialization—avoid `.cpu()` or `.to('cpu')` on compute tensors.
        - Provide at least one @triton.jit kernel and a top-level Python wrapper
          named kernel_function(...). This wrapper must accept the same primary
          input tensor(s) as the model and any required weights/biases with shapes
          implied by the problem; it should orchestrate Triton kernel(s) and
          return the final output tensor.
        - No PyTorch math path: kernel_function MUST compute the final outputs
          using your Triton kernels only. Do NOT implement or fall back to
          torch.nn / torch.nn.functional / torch.* ops
          sigmoid, etc.) for producing the final result. Using PyTorch for
          reference comparisons is allowed only inside the self-test.
        - Use the data layout and dtype semantics indicated by subgraphs, defaulting
          to NCHW + float32 if unspecified. Respect stride/padding/dilation/groups,
          and exact op order.
        - Numerical equivalence: include a self-test (test_kernel or run_tests)
          that compares your Triton-based result to a PyTorch reference computed
          from the original problem code below (use get_init_inputs() and
          get_inputs() if present to instantiate the Model). The test must print
          'PASS' on success and exit with code 0. Use allclose with rtol<=1e-3,
          atol<=1e-3 for fp32; for fp16/bf16 allow up to 2e-2.
        - No imports beyond torch, triton, triton.language as tl, and stdlib. No I/O.
        - Do NOT monkey-patch PyTorch device functions or torch.cuda.is_available()
        - Do NOT manipulate TRITON_BACKENDS environment variable
        - Do NOT disable or mock XPU/CUDA drivers

        Implementation tips:
        - If merging multiple subgraphs, ensure intermediate tensor shapes match.
        - Hoist constant weights or parameters to avoid reloading per block.
        - Use tl.load/tl.store with masks for boundary conditions.
        - Favor coalesced memory access; tile by blocks; compute grid from shape.
        - Common Triton pitfalls to avoid:
          * Do NOT call tl.broadcast on Python scalars; tl.maximum(x, 0.0) works.
          * Prefer scalar constants directly in elementwise ops (no explicit broadcast needed).
          * Keep BLOCK_SIZE power-of-two; mask stores at tail.
        """
    ).strip()  # 去除首尾空白字符

    # 第五步:拼接完整的用户指令(按逻辑组织各部分内容)
    user_lines: list[str] = []
    user_lines.append(guidance)  # 核心指导语
    user_lines.append("")  # 空行分隔
    user_lines.append("SUBGRAPHS (summary):")  # 子图摘要标题
    user_lines.append(sg_summary)  # 子图摘要内容
    user_lines.append("")  # 空行分隔
    user_lines.append("ORIGINAL PROBLEM FILE:")  # 原始问题代码标题
    user_lines.append("```python")  # Python代码块开始标记
    user_lines.append(problem_code)  # 原始问题代码内容
    user_lines.append("```")  # Python代码块结束标记
    user_lines.append("")  # 空行分隔
    user_lines.append("SUBGRAPH KERNELS (reference implementations):")  # 参考内核标题
    user_lines.append(kernels_section)  # 参考内核代码内容
    user_lines.append("")  # 空行分隔
    # 最终要求:仅返回一个包含完整代码的Python代码块
    user_lines.append(
        "Return only one fenced Python code block with your final composed implementation."
    )
    # 拼接所有行,形成完整的Prompt
    return "\n".join(user_lines)

def _build_refinement_prompt(
    problem_code: str,
    subgraphs: list[dict[str, Any]],
    kernel_items: list[KernelItem],
    previous_code: str,
    error_info: dict[str, str],
    target_platform: PlatformConfig,
) -> str:
    """Prompt the LLM to refine the previously produced code based on errors.
    构建迭代优化Prompt:
    基于上一轮代码的错误信息,引导LLM修复Triton算子代码中的编译/运
Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐