diff --git a/docs/fix-validation-flow.md b/docs/fix-validation-flow.md
new file mode 100644
index 0000000..121c645
--- /dev/null
+++ b/docs/fix-validation-flow.md
@@ -0,0 +1,106 @@
+# 编译与测试修复代码的评价和测试机制
+
+## 总体流程
+
+```mermaid
+flowchart TD
+ subgraph BUILD["编译修复循环 (pipeline/build.py)"]
+ B1["Build TA"] --> B2{编译通过?}
+ B2 -->|是| B_DONE["build_passed=true"]
+ B2 -->|否| B3["AI 修复 (ai_fix)"]
+ B3 --> B4["代码评价 (validate_fix)"]
+ B4 --> B5{校验通过?}
+ B5 -->|否| B6["回退修改 + 写入拒绝原因"]
+ B6 --> B7["不消耗 attempt 次数"]
+ B7 --> B3
+ B5 -->|是| B8["记录本次 fix attempt"]
+ B8 --> B9{attempt <= max?}
+ B9 -->|是| B1
+ B9 -->|否| B_FAIL["build_passed=false"]
+ end
+
+ subgraph TEST["测试修复循环 (pipeline/test.py)"]
+ T1["Run pytest"] --> T2{测试通过?}
+ T2 -->|是| T_DONE["test_passed=true"]
+ T2 -->|否| T3{"OOM 检测?"}
+ T3 -->|是| T4["降并发重跑 (最多5次)"]
+ T4 --> T5{OOM 消失?}
+ T5 -->|是| T6["继续正常修复"]
+ T5 -->|否| T_FAIL["test_passed=false"]
+ T3 -->|否| T6
+ T6 --> T7["AI 修复 (ai_fix)"]
+ T7 --> T8["代码评价 (validate_fix)"]
+ T8 --> T9{校验通过?}
+ T9 -->|否| T10["回退修改 + 写入拒绝原因"]
+ T10 --> T11["不消耗 attempt 次数"]
+ T11 --> T7
+ T9 -->|是| T12["Rebuild TA"]
+ T12 --> T13{编译通过?}
+ T13 -->|否| T14["attempt += 1"]
+ T14 --> T15{attempt <= max?}
+ T15 -->|是| T7
+ T15 -->|否| T_FAIL
+ T13 -->|是| T1
+ end
+
+ B_DONE --> NEXT["下一步"]
+ B_FAIL --> TERMINATE["流程终止"]
+ T_DONE --> NEXT
+ T_FAIL --> TERMINATE
+```
+
+## 代码评价机制 (validate_fix)
+
+```mermaid
+flowchart TD
+ A["AI 修复完成"] --> B["获取 modified_files 列表"]
+ B --> C{"有修改文件?"}
+ C -->|否| REJECT["❌ 拒绝: No files modified"]
+ C -->|是| D["逐个检查文件路径"]
+ D --> E{"文件在 third_party/ascend/ 下?"}
+ E -->|全部是| PASS["✅ 校验通过
记录 attempt"]
+ E -->|有文件在外面| F["记录非法文件列表"]
+ F --> G["git checkout -- ."]
+ G --> H["git clean -fd"]
+ H --> I["写入 fix_rejection.txt
包含拒绝原因 + 允许的路径"]
+ I --> J["追加到 fix_errors 列表
(AI 下次修复会看到)"]
+ J --> K["❌ 拒绝: 不消耗 attempt
continue 重新修复"]
+
+ style PASS fill:#4a9,stroke:#333
+ style REJECT fill:#d73,stroke:#333
+ style K fill:#d73,stroke:#333
+```
+
+## AI 自检机制 (prompt.md step 6)
+
+```mermaid
+flowchart LR
+ subgraph AI_SELF["AI 修复时自检 (prompt.md)"]
+ S1["Step 6: SELF-REVIEW before returning"]
+ S2["列出每个修改的文件"]
+ S3{"文件在 third_party/ascend/ 下?
(测试修复也可在 python/triton_ascend/)"}
+ S3 -->|否| S4["在返回前 REVERT 该修改"]
+ S4 --> S3
+ S3 -->|是| S5{"修复针对根因
而非掩盖错误?"}
+ S5 -->|否| S4
+ S5 -->|是| S6["返回修改"]
+ end
+
+ style S4 fill:#d73,stroke:#333
+ style S6 fill:#4a9,stroke:#333
+```
+
+## 双层防护总结
+
+| 层级 | 位置 | 机制 | 失败处理 |
+|------|------|------|----------|
+| **第1层** | AI 自检 (prompt.md) | AI 提交前自查文件路径 + 根因 | 自行回退后重新修复 |
+| **第2层** | 代码校验 (fix.py) | 硬检查 modified_files 路径 | git revert + 反馈文件 + 不消耗 attempt |
+| **第3层** | 编译/测试验证 | 实际 build/test 结果 | 编译/测试失败 → 继续 fix loop |
+
+## 关键设计决策
+
+1. **拒绝不消耗 attempt**:修复被拒后 `continue` 在同一 attempt 重新修复(while 循环),确保 AI 有完整的 max_retries 次有效尝试
+2. **拒绝后回退代码**:`git checkout -- .` + `git clean -fd` 清除所有修改,下一轮从干净状态开始
+3. **拒绝反馈传递给 AI**:`fix_rejection.txt` 追加到 `fix_errors` 列表,AI 下次修复时在 error_logs 中看到
+4. **OOM 单独处理**:OOM 不调用 AI,自动降并发重跑(test_procs 减半),与代码修复完全独立
diff --git a/docs/guide.md b/docs/guide.md
index f675996..d9d1364 100644
--- a/docs/guide.md
+++ b/docs/guide.md
@@ -1,752 +1,302 @@
-# TA Main2Main Auto-Sync — 工作流完整指南
+# TA Main2Main Auto-Sync -- 工作流完整指南
## 概述
TA_main2main_workflow 是一个自动化流水线,用于将上游 Triton 的更新同步到
-triton-ascend(Triton 的 Ascend NPU 适配版)。它基于 CrewAI Flow 框架,通过
-git merge + AI 辅助的方式,自动完成从检测更新、合并代码、解决冲突、编译构建、
-运行测试到提交 PR 的全流程。
+triton-ascend(Triton 的 Ascend NPU 适配版)。通过 **git merge + AI 辅助**,
+自动完成从检测更新、合并代码、解决冲突、编译构建、运行测试到提交 PR 的全流程。
-### 核心思路
+### 架构
-triton-ascend 是 Triton 的 **fork**。上游 Triton 的 `main` 分支每次推进,
-triton-ascend 都需要跟上。但 triton-ascend 中包含大量 Ascend NPU 的适配代码
-(新增文件 和 对上游文件的修改),简单的 `git merge` 经常会产生冲突,合并后
-也可能出现编译错误或测试失败。
-
-本工作流通过 **AI + 确定性脚本** 的组合来系统化地解决这个问题:
-确定性脚本负责 git 操作、编译、测试等可重复操作,AI 负责解决冲突和修复测试失败。
-
-### 渐进式步骤合并(Progressive Step Merge)
-
-当上游 Triton 的 `main` 分支有较多 commit 推进时(例如 50+ 个 commit),
-一次性 `git merge` 全部变更会带来几个问题:
-
-1. **冲突范围过大**:大量变更同时合并,冲突文件多,AI 解决冲突的难度指数级上升
-2. **修复定位困难**:编译或测试失败时,很难确定是哪个上游 commit 引入的问题
-3. **回滚成本高**:合并到一半失败,所有工作丢失,需要从头再来
-
-为解决这些问题,TA_main2main_workflow 引入了 **渐进式步骤合并**
-(Progressive Step Merge)机制,自动将上游 commit 按代码变更量切分为多个
-"步骤"(step),每个步骤分别合并、验证、修复。这参考了
-[vllm-ascend 的 main2main_flow](https://github.com/triton-lang/triton-ascend)
-中的 `plan_steps` 设计。
-
-#### 切分算法(plan_steps)
-
-`scripts/plan_steps.py` 中的确定性算法:
-
-1. **列出所有 upstream commit**:`git log --reverse merge_base..target_commit`
-2. **逐 commit 统计源码行变更**:对每个 commit 运行 `git diff-tree --numstat`,
- 只统计关键源码目录(`python/triton/`、`lib/`、`include/`)的增删行数
-3. **跳过无关 commit**:不涉及源码目录的 commit(如 CI 配置、文档修改)会被跳过,
- 不占用步骤配额
-4. **按预算累积分组**:
- - **行数预算**(`TA_LINE_BUDGET`):默认 1000 行。一个步骤内所有 commit
- 的源码变更行数之和不超过此值
- - **commit 数预算**(`commit_count_budget`):由公式 `max(1, round(TA_COMMIT_BUDGET * sqrt(TA_LINE_BUDGET / 1000)))` 计算。
- 默认 TA_COMMIT_BUDGET=5、TA_LINE_BUDGET=1000 → 每步最多 **5 个 commit**
- - commit 逐个累积,直到超过任一预算,则当前步骤结束,开始新步骤
-5. **超大 commit 单独成步**:单个 commit 的源码变更超过行数预算时,
- 该 commit 独占一个步骤
-
-#### 渐进式合并流程
-
-在 progressive mode 下(默认开启),每个步骤的合并流程:
+工作流采用模块化管道架构:
```
-Step 1: git merge step-1.end_commit → resolve conflicts → build → test → fix → commit
-Step 2: git merge step-2.end_commit → resolve conflicts → build → test → fix → commit
-Step 3: git merge step-3.end_commit → resolve conflicts → build → test → fix → commit
-...
-Final: 生成累积 patch 和 summary → 切回原始分支
+ta-kickoff (main.py)
+ └── TA_Main2MainFlow (flow.py) — 143 workflow编排器
+ ├── utils/ — TAConfig, WorkflowContext, TALogger, run_git, timed
+ ├── pipeline/ — 13 个独立管道模块
+ ├── agent/ — AI 适配器 + prompt 模板
+ └── reference/ — AI 参考知识库
```
-关键设计点:
-- 第一个步骤会从 `triton-lang/triton-ascend` 的 `main` 分支创建新的 work branch
-- 后续步骤在同一个 work branch 上累积合并(使用 `git merge --no-ff`)
-- 每个步骤成功后自动 `git commit`,失败时保留 work branch 的中间状态
-- 最终生成的 `final_target.patch` 是**所有步骤的累积 diff**
+每个管道模块遵循统一签名:`def step(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext`
-#### 控制开关
+### 渐进式步骤合并
+
+当上游有多 commit 时,自动按代码变更量切分为多个步骤,每步分别合并、验证、修复:
+
+1. 逐 commit 统计源码行变更
+2. 按行数预算(`TA_LINE_BUDGET`,默认 1000 行)分组
+3. LLVM hash 变更的 commit 独占一步
+4. 超大 commit 独占一步
| 环境变量 | 说明 | 默认值 |
|---------|------|--------|
-| `TA_PROGRESSIVE_MERGE` | 是否启用渐进式步骤合并 | `true` |
| `TA_LINE_BUDGET` | 每步骤最大源码变更行数 | `1000` |
-```bash
-# 关闭渐进式合并,使用传统的单步合并(一次性合并全部 commit)
-TA_PROGRESSIVE_MERGE=false ta-kickoff ...
-
-# 调整每步骤的行数预算(越大每步包含的 commit 越多,执行时间越长)
-TA_LINE_BUDGET=2000 ta-kickoff ...
-
-# 更细粒度控制(每步最多 500 行变更)
-TA_LINE_BUDGET=500 ta-kickoff ...
-
-# 控制 commit 数粒度(每步最多 ~3 个 commit,生成更多步骤)
-TA_COMMIT_BUDGET=3 ta-kickoff ...
-
-# 更粗粒度(每步最多 ~10 个 commit)
-TA_COMMIT_BUDGET=10 ta-kickoff ...
-```
-
---
-## 工作流示意图
-
-### Mermaid 流程图
-
-```mermaid
-flowchart TD
- A["ta-kickoff"]
- A --> B["Phase 0: Initialize"]
- B --> C["Phase 1: Detect Commits"]
- C -->|no new commits| D["Done: Already Up-to-Date"]
- C -->|new commits found| E["Phase 2A: Merge Upstream"]
- E -->|conflicts| F["Phase 2B: AI Resolve Conflicts"]
- E -->|no conflicts| G["Phase 2C: Build and Test"]
- F -->|resolved| G
- F -->|max retries| H["Failure: write FAILURE.md"]
- G --> G1["Build: setup.py install"]
- G1 -->|failed| I["AI Fix Code"]
- G1 -->|passed| G2["Test: pytest"]
- G2 -->|failed| I
- G2 -->|passed| J["Commit Fixes"]
- I -->|retry| G1
- I -->|max retries| H
- J --> K["Phase 2D: Finalize"]
- K --> L["Success"]
- L -->|push enabled| M["Push Branch and Create PR"]
- L -->|push disabled| N["Done: work branch kept"]
-```
-
-### ASCII 文本流程图
+## 工作流总览
```
-┌──────────────────────────────────────────────────────────────────────────────┐
-│ TA Main2Main Auto-Sync Flow │
-└──────────────────────────────────────────────────────────────────────────────┘
-
- ┌──────────────┐
- │ START │
- │ (ta-kickoff) │
- └──────┬───────┘
- │
- ▼
- ┌────────────────────────────────┐
- │ Phase 0: initialize │
- │ ────────────────────────── │
- │ • 清空 workspace/ │
- │ • 读取 CLI / 环境变量配置 │
- │ • 记录当前分支和 HEAD │
- │ • 中止残留的 git merge │
- └────────────────┬───────────────┘
- │
- ▼
- ┌────────────────────────────────┐
- │ Phase 1: detect_commits │
- │ ────────────────────────── │
- │ • fetch upstream-triton │
- │ • 找到 merge-base │
- │ • 列出待合并的 upstream commit │
- │ • 统计改动文件数和行数 │
- └────────────────┬───────────────┘
- │
- ┌──────────┴──────────┐
- │ 有新 commit 吗? │
- └──────────┬──────────┘
- Yes │ No
- ┌─────────────────┘ └──────────────┐
- ▼ ▼
- ┌──────────────────────────────┐ ┌──────────────────────────────┐
- │ Phase 2: execute_sync │ │ Already Up-to-Date │
- │ (单节点编排,内部4步) │ │ 无需同步,直接结束 │
- └──────────────┬───────────────┘ └──────────────────────────────┘
- │
- ┌──────────────┴─────────────────────────────────────────────────────┐
- │ │
- ▼ │
-┌─────────────────────────────────────────────┐ │
-│ Step 2A: _do_merge │ │
-│ ─────────────────────────── │ │
-│ • 从 triton-lang/triton-ascend 的 │ │
-│ 最新 main 创建 work branch │ │
-│ • git fetch upstream-triton │ │
-│ • git merge │ │
-│ • 记录冲突文件(如有) │ │
-└──────────────┬──────────────────────────────┘ │
- │ │
- ▼ │
- ┌──────┴──────┐ │
- │ 有冲突吗? │ │
- └──────┬──────┘ │
- Yes │ No │
- ┌──────────┘ └───────────┐ │
- ▼ ▼ (直接跳到 Step 2C) │
-┌─────────────────────────────────────────────┐ │
-│ Step 2B: _do_resolve_conflicts │ │
-│ ────────────────────────────── │ │
-│ • AI (opencode/claude) 读取冲突文件 │ │
-│ • AI 分析并解决冲突 │ │
-│ • 最多重试 3 次 │ │
-│ • 解决后 git add -u && git commit │ │
-│ • 运行 pre-CI 检查(残留标记/语法等) │ │
-└──────────────┬──────────────────────────────┘ │
- │ (冲突解决成功或本来就没有冲突) │
- ▼ │
-┌─────────────────────────────────────────────┐ │
-│ Step 2C: _do_build_and_fix_loop │ │
-│ ─────────────────────────────── │ │
-│ │ │
-│ ┌─────────────────────────────────────┐ │ │
-│ │ ATTEMPT LOOP │ │ │
-│ │ (最多 max_retries=3 轮) │ │ │
-│ │ │ │ │
-│ │ ┌─► build_triton_ascend() │ │ │
-│ │ │ python3 setup.py install │ │ │
-│ │ │ (带 Ascend 编译环境变量) │ │ │
-│ │ │ │ │ │
-│ │ │ Build 失败? ──Yes──► AI fix ──┘ │ │
-│ │ │ │ │ │
-│ │ │ No │ │
-│ │ │ ▼ │ │
-│ │ │ run_tests() │ │
-│ │ │ pytest -n 16 unittest/pytest_ut │ │
-│ │ │ │ │
-│ │ │ Test 全部通过? ──Yes──► 跳出循环 │ │
-│ │ │ │ │ │
-│ │ │ │ 有 NPU OOM? │ │
-│ │ │ │ (重跑全部) │ │
-│ │ │ │ │ │
-│ │ │ No (有失败) │ │
-│ │ │ │ │ │
-│ │ └── AI fix ◄────────────────────────────┘ │
-│ │ (opencode fix 模式) │
-│ └─────────────────────────────────────┘ │
-│ │
-│ 所有测试通过 ▼ │
-│ _commit_fixes() │
-│ git add -u && git commit │
-└──────────────┬──────────────────────────────────────────────┐ │
- │ │ │
- ▼ │ │
-┌─────────────────────────────────────────────┐ │ │
-│ Step 2D: _do_finalize │ │ │
-│ ────────────────────── │ │ │
-│ • 生成 final_summary.md │ │ │
-│ • 生成 cumulative patch (final_target.patch)│ │ │
-│ • 切回原始分支 (git checkout original) │ │ │
-│ • 保留 work branch 供检查 │ │ │
-└─────────────────────────────────────────────┘ │ │
-│ │ │
-└────────────────────────────────────────────────────────────┘ │
- │
- ┌────────────────────────────────────────────────────────┐ │
- │ execute_sync 返回 UpgradeCompleted 或 │ │
- │ UpgradeFailed │ │
- └────────────┬──────────────────────┬─────────────────────┘ │
- │ │
- Completed Failed
- │ │
- ▼ ▼
- ┌──────────────────────┐ ┌──────────────────────┐
- │ push_to_github │ │ handle_failure │
- │ ──────────────── │ │ ────────────── │
- │ (PUSH_TO_GITHUB=true │ │ • 写 FAILURE.md │
- │ 时才触发) │ │ • 打印诊断信息 │
- │ │ │ • 保留 work branch │
- │ • git add -u │ │ 供手动排查 │
- │ • git commit │ │ │
- │ • git push origin │ └──────────────────────┘
- │ • gh pr create │
- └──────────────────────┘
- │
- ▼
- ┌──────────────┐
- │ DONE │
- │ PR URL or │
- │ FAILURE.md │
- └──────────────┘
+Phase 0: Prepare 克隆/配置 repo,设置 remotes,fetch
+Phase 1: Detect 检测待合并的上游 commits,计算 merge-base
+Phase 2: Plan 按行数预算切分为步骤 (steps.json)
+Phase 3: Per-Step 对每个步骤执行 merge→resolve→build→[ir_patch]→test→commit
+Phase 4: Finalize 生成 cumulative patch + summary
+Phase Push: (可选) push + gh pr create
```
----
-
-## 前置条件
-
-### 必需环境
-
-| 依赖 | 说明 |
-|------|------|
-| Python 3.10–3.13 | 运行工作流本身 |
-| git | 所有版本控制操作 |
-| `triton` 仓库 | 上游 Triton 的本地 clone |
-| `triton-ascend` 仓库 | triton-ascend 的本地 clone(Ascend 适配版) |
-
-### AI 后端(二选一)
+详见 [workflow.md](workflow.md) 中的 Mermaid 流程图。
-| 后端 | 安装方式 |
-|------|---------|
-| **opencode**(推荐) | 安装 [`opencode`](https://opencode.ai) CLI,确保在 `$PATH` 中 |
-| **claude** | 安装 `claude` CLI (`npm install -g @anthropic-ai/claude`) |
-
-### 构建环境
-
-- **LLVM**:设置 `LLVM_INSTALL_PREFIX` 指向 LLVM 安装目录
-- **Ascend NPU**(可选):测试需要 Ascend 硬件;可设置 `SKIP_E2E_TEST=true` 跳过
-- **Conda**(推荐):默认使用 `ta-upgrade` 环境
-
-### PR 创建(可选)
+---
-- [GitHub CLI](https://cli.github.com/) (`gh`) 已登录认证
-- 设置 `PUSH_TO_GITHUB=true` 和 `GITHUB_REPO=triton-lang/triton-ascend`
+## 环境变量
+
+所有配置通过 `TAConfig.from_env()` 集中读取。CLI 参数可覆盖。
+
+| 变量 | 说明 | 默认值 |
+|------|------|--------|
+| `TRITON_ASCEND_PATH` | triton-ascend 仓库路径 | (workspace clone) |
+| `TRITON_PATH` | 上游 triton 仓库路径 | -- |
+| `TRITON_TARGET_COMMIT` | 目标 upstream commit | upstream HEAD |
+| `AI_BACKEND` | AI 后端: opencode / claude / auto | auto |
+| `TA_MAX_RETRIES` | AI 修复最大重试次数 | 10 |
+| `IR_MAX_ITERATIONS` | IR 补丁最大迭代次数 | 3 |
+| `SKIP_AI_ANALYSIS` | 跳过 AI 调用 | false |
+| `SKIP_BUILD` | 跳过编译 | false |
+| `SKIP_E2E_TEST` | 跳过测试 | false |
+| `SKIP_LLVM_REBUILD` | 跳过 LLVM 重编译 | false |
+| `SKIP_IR_PATCH` | 跳过 IR 补丁阶段 | false |
+| `PUSH_TO_GITHUB` | 成功后创建 PR | false |
+| `GITHUB_REPO` | PR 目标仓库 owner/name | triton-lang/triton-ascend |
+| `GH_TOKEN` | GitHub token (PR 创建) | -- |
+| `LLVM_PROJECT_PATH` | llvm-project 仓库路径 | ~/llvm-project |
+| `LLVM_INSTALL_PREFIX_SYNC` | LLVM 安装前缀 | ~/llvm-install-sync |
+| `NUM_PROCS` | pytest 并行 worker 数 | 16 |
+| `TA_BASE_BRANCH` | 工作分支基准 | upstream_sync |
+| `TA_PR_BASE_BRANCH` | PR 目标分支 | upstream-sync |
+| `TA_LINE_BUDGET` | 每步骤最大变更行数 | 1000 |
+| `TA_RESUME` | 断点续传模式 | false |
---
## 安装与运行
-### 安装
-
```bash
cd TA_main2main_workflow
pip install -e .
-```
-安装后可使用两个命令:`ta-kickoff`(执行同步)和 `ta-plot`(生成流程图)。
+# 基本用法
+ta-kickoff --triton-ascend-path ./triton-ascend
-### 基本用法
+# 指定目标 commit
+ta-kickoff --target-commit abc123def --max-retries 5 --num-procs 8
-```bash
-# 最简单的用法 — 两个 repo 都在当前目录
-ta-kickoff --triton-ascend-path ./triton-ascend --triton-path ./triton
-
-# 指定要同步到的上游 commit
-ta-kickoff \
- --triton-ascend-path ./triton-ascend \
- --triton-path ./triton \
- --target-commit abc123def456789...
-
-# Dry-run:只走流程不调 AI、不跑测试
-SKIP_AI_ANALYSIS=true SKIP_E2E_TEST=true SKIP_BUILD=true \
-ta-kickoff --triton-ascend-path ./triton-ascend --triton-path ./triton
+# Dry-run
+SKIP_AI_ANALYSIS=true SKIP_BUILD=true SKIP_E2E_TEST=true ta-kickoff
```
---
## 各阶段详解
-### Phase 0 — Initialize(初始化)
-
-这是每次运行的起点,执行以下操作:
-
-1. **清空 workspace**:删除上一轮的 `workspace/` 目录,重新创建
-2. **读取配置**:按优先级 CLI args → 环境变量 → 默认值 读取路径和参数
-3. **记录原始状态**:保存当前 git 分支名和 HEAD commit,流程结束后会切回来
-4. **清理残留状态**:如果上次运行异常中断留下了 `MERGE_HEAD`,自动 `git merge --abort`
-
-### Phase 1 — Detect Commits(检测更新)
-
-确定需要同步的 upstream commit 范围:
-
-1. Fetch `upstream-triton` 远程的最新引用
-2. 计算 `merge-base`(triton-ascend 与 upstream 的共同祖先)
-3. 列出 `merge-base..target_commit` 之间所有 upstream commit
-4. 统计改动范围(文件数、行数、各目录的改动量)
-
-**输出文件**:`workspace/detect.json`
+### Phase 0 -- Prepare (准备环境)
-如果 `merge-base == target_commit`(没有新 commit),流程直接结束,返回 "已是最新"。
+1. 确保 triton-ascend 仓库存在(clone 或使用本地路径)
+2. 配置 origin 和 triton-upstream remotes
+3. Fetch 两个 remote
+4. 基于 base branch 创建 work branch
+5. 解析 ascend HEAD 和 target commit
-当检测到新 commit 且 `TA_PROGRESSIVE_MERGE=true`(默认)时,
-Phase 1 还会自动执行 **步骤规划**(Step Planning):
+**负责模块**: `pipeline/prepare.py`
-1. 逐 commit 统计每个 upstream commit 在关键源码目录中的变更行数
-2. 按行数预算(`TA_LINE_BUDGET`,默认 1000 行)和 commit 数预算切分为多个步骤
-3. 不涉及源码的 commit 会被跳过(如文档、CI 配置修改)
-4. 生成 `workspace/steps.json` 和每个步骤的 patch/commit 列表
+### Phase 1 -- Detect (检测更新)
-### Phase 2 — Execute Steps(步骤执行)
+1. 计算 merge-base
+2. 列出待合并的 upstream commits
+3. 统计改动文件数和行数
+4. 输出 `workspace/detect.json`
-> **渐进式模式**(`TA_PROGRESSIVE_MERGE=true`,默认):对每个规划的步骤执行完整的
-> 合并→解决冲突→编译→测试→修复→提交 流程,前一步成功后才进入下一步。
->
-> **单步模式**(`TA_PROGRESSIVE_MERGE=false`):一次性合并所有 upstream commit,
-> 与渐进式模式使用完全相同的内部方法。
+**负责模块**: `pipeline/detect.py`
-每个步骤(或单步模式下的全局)执行以下子阶段:
+### Phase 2 -- Plan (规划步骤)
-#### Phase 2A — Merge Upstream(合并上游)
+1. 逐 commit 统计源码行变更
+2. 检测含 LLVM hash 变更的 commit(`cmake/llvm-hash.txt`)
+3. 按行数预算分组,LLVM 变更和大 commit 独占步骤
+4. 输出 `workspace/steps.json`
-将当前步骤的 upstream 变更合入 triton-ascend:
+**负责模块**: `pipeline/plan.py`
-1. **创建工作分支**(仅第一步):
- - 确保 triton-ascend 仓库有指向 `https://github.com/triton-lang/triton-ascend.git` 的 remote
- - `git fetch main` 拉取最新的 main
- - `git checkout -B auto/upstream-sync- /main` 基于最新 main 创建分支
+### Phase 3 -- Per-Step Loop (步骤执行)
-2. **合并步骤目标 commit**:
- - Fetch upstream-triton 的最新数据
- - `git merge --no-ff --no-edit `
- - 后续步骤在上一步的 work branch 上继续 merge(累积合并)
+对每个步骤依次执行以下子阶段:
-3. **记录结果**:
- - 如果有冲突,保存每个冲突文件的完整内容到 `workspace/conflicts/`
- - 输出 merge 日志到 `workspace/merge.log`
+#### Step A: Merge (合并)
-**输出文件**:`workspace/merge_result.json`、`workspace/merge.log`、`workspace/conflicts/`
+`git merge --no-ff` 合并步骤的 end_commit。检测冲突文件。
-#### Phase 2B — Resolve Conflicts(AI 解决冲突)
+**负责模块**: `pipeline/merge.py`
-> 仅当合并产生冲突时执行。无冲突则跳过。
+#### Step B: Resolve (AI 解决冲突)
-1. **AI 分析冲突**:将冲突文件快照传递给 AI(opencode/claude),AI 在 `conflict` 模式下分析并解决
-2. **最多重试 3 次**:每次 AI 尝试后检查是否还有残留冲突标记
-3. **提交解决结果**:`git add -u && git commit -s`
-4. **Pre-CI 检查**:扫描是否还有冲突标记、临时文件、语法错误
+仅在 merge 有冲突时执行。AI 分析冲突并解决,最多 `max_retries` 次尝试。
-AI 参考 `reference/` 目录下的知识库文件来理解 triton-ascend 的代码结构、常见错误模式和适配策略。
+**负责模块**: `pipeline/resolve.py`
-**注意**:`git add -u` 只 stage 已跟踪文件的修改,不会把测试过程中产生的临时文件、缓存、日志等误提交。
-
-#### Phase 2C — Build & Fix Loop(编译→测试→修复循环)
-
-这是核心循环,最多执行 `max_retries+1`(默认 4 次 = 首轮 + 3 轮修复):
+#### Step C: Build + Fix (编译修复循环)
```
-Round 0: build → test (首次尝试,不做修复)
-Round 1: AI fix → build → test
-Round 2: AI fix → build → test
-Round 3: AI fix → build → test (最后一次)
+Round 0: build (首次尝试)
+Round 1-N: AI fix → validate_fix → build
```
-#### 编译步骤
+**修复校验 (validate_fix)**:所有 AI 改动必须在 `third_party/ascend/` 下。
+校验失败则 git revert + 写入反馈文件 + 不消耗 attempt 次数。
-```bash
-# 在 triton-ascend 目录下,带 Ascend 编译环境变量
-LLVM_SYSPATH=$LLVM_INSTALL_PREFIX \
-TRITON_BUILD_WITH_CCACHE=true \
-TRITON_BUILD_WITH_CLANG_LLD=true \
-TRITON_BUILD_PROTON=OFF \
-DEBUG=1 \
-TRITON_WHEEL_NAME="triton-ascend" \
-TRITON_APPEND_CMAKE_ARGS="-DTRITON_BUILD_UT=OFF" \
-python3 setup.py install
-```
+**负责模块**: `pipeline/build.py`, `pipeline/fix.py`
-编译成功后会清理 `~/.triton/cache/`。
+**校验流程详见**: [fix-validation-flow.md](fix-validation-flow.md)
-#### 测试步骤
-
-```bash
-pytest -n 16 third_party/ascend/unittest/pytest_ut/
-```
+#### IR Patch Pipeline (LLVM hash 变更时)
-#### AI 修复
+当步骤的 reason 为 `llvm_version` 时,走完整的 IR 补丁管线:
-当编译或测试失败时,AI 以 `fix` 模式运行:
-- 读取错误日志(`build_result.json` 或 `test_result.json`)
-- 分析失败原因并修改 triton-ascend 源码
-- 修改后运行 pre-CI 检查
+**Phase 1 -- 编译适配**:
+1. 编译 clean LLVM(无补丁)
+2. 编译 TA + AI 修复编译错误
+3. 只允许修改 `third_party/ascend/` 下的代码
+4. AscendNPU-IR 错误自动检测,AI 参考专项文档
-#### NPU OOM 特殊处理
+**Phase 2 -- IR 补丁 + 测试循环**:
+1. OP 使用分析 → LLVM 变更分析 → AI 生成 IR 补丁
+2. 应用补丁 + 编译 LLVM(apply/Build 失败时 AI 修复补丁,最多 10 次)
+3. 编译 TA → pytest
+4. 测试失败时 AI 分类:IR 问题则重新生成补丁(最多 3 次),代码问题则 AI 修复
-如果测试日志中出现 NPU OOM(`npu.OutOfMemoryError`、`device memory allocation failed` 等),
-**不修改代码**,而是重跑全部测试用例。OOM 是瞬时的资源分配问题,非代码缺陷。
-详见 `reference/npu-oom-handling.md`。
+**负责模块**: `pipeline/ir_patch.py`
-#### 修复提交
+#### Step D: Test + Fix (测试修复循环)
-当所有测试通过后,自动提交修复:
-
-```bash
-git add -u
-git commit -s -m "fix:
-
-Upstream target:
-Fix attempt: <尝试次数>
-Work branch: <分支名>"
```
-
-#### Phase 2D — Commit Step(提交步骤进度)
-
-每个步骤通过测试后,自动提交该步骤的进度:
-
-```bash
-git add -u
-git commit -s -m "sync: merge upstream commits for step
-...
-Upstream range: ..
-Step: N/M
-Commits in step: "
+Round 0: pytest (首次)
+Round 1-N: [OOM? → 降并发重跑] → AI fix → validate_fix → rebuild → pytest
```
-#### Phase 2E — Finalize(收尾)
-
-1. **生成产物**:
- - `workspace/final_summary.md` — 最终同步摘要(渐进式模式下合并所有步骤的摘要)
- - `workspace/final_target.patch` — 本次同步的**完整累积 diff**(从原始 ascend HEAD 到最新 work branch HEAD)
+- **OOM 处理**:检测到 NPU OOM 时自动降并发(test_procs 减半)重跑,不调用 AI
+- **修复校验**:同编译修复,改动的文件必须符合要求
+- **测试修复允许范围**:优先 `third_party/ascend/`,根因在上游时允许最小化修改
-2. **恢复分支**:`git checkout ` 切回同步前的分支
- - work branch 保留不删,方便手动检查
+**负责模块**: `pipeline/test.py`, `pipeline/fix.py`
-3. **打印汇总表**:所有阶段(含每个步骤)的通过/失败/跳过状态一览
+#### Step E: Commit (提交)
-### Terminal — Push to GitHub / Handle Failure
+提交当前步骤的进度,含 AscendNPU-IR submodule 变更。
-#### 成功 → Push to GitHub
+**负责模块**: `pipeline/commit.py`
-**渐进式模式**(默认,`TA_PROGRESSIVE_MERGE=true`):
+### Phase 4 -- Finalize (收尾)
-每个步骤完成后,当 `PUSH_TO_GITHUB=true` 时:
+1. 生成 `final_summary.md`
+2. 生成累积 patch (`final_target.patch`)
+3. 生成 PR body (`pr_body.md`)
-1. 认证 GitHub CLI
-2. `git push -u origin ` 推送当前 work branch
-3. **第一步**:`gh pr create` 创建 PR(标题含步骤编号,如 `[Step 1/3] sync: upstream triton merge`)
-4. **后续步骤**:直接 push 到同一分支,PR 自动更新
-5. 全部步骤完成后,更新 PR 描述:列出所有已完成步骤的摘要
+**负责模块**: `pipeline/finalize.py`
-**单步模式**(`TA_PROGRESSIVE_MERGE=false`):
+### Phase Push -- PR (可选)
-1. 认证 GitHub CLI
-2. `git add -u && git commit -s`(提交最终变更)
-3. `git push -u origin `
-4. `gh pr create` 创建 PR,base 指向仓库默认分支
+当 `PUSH_TO_GITHUB=true` 时:
+1. `git push -u origin `
+2. `gh pr create` 创建 PR
-#### 失败 → Handle Failure
-
-1. 写 `workspace/FAILURE.md` 包含完整诊断信息
-2. 打印恢复命令(如何切回原始分支、如何删除失败的 work branch)
-3. work branch 保留供手动排查
+**负责模块**: `pipeline/push_pr.py`
---
-## 环境变量完整参考
-
-| 变量 | 用途 | 默认值 | 生效阶段 |
-|------|------|--------|---------|
-| `TRITON_ASCEND_PATH` | triton-ascend 仓库路径 | 当前目录 | Initialize |
-| `TRITON_PATH` | 上游 triton 仓库路径 | 当前目录 | Initialize |
-| `TRITON_TARGET_COMMIT` | 目标 upstream commit | triton HEAD | Detect |
-| `AI_BACKEND` | AI 后端:`opencode` 或 `claude` | 自动检测 | Resolve / Fix |
-| `SKIP_AI_ANALYSIS` | 跳过所有 AI 调用 | `false` | Resolve / Fix |
-| `SKIP_BUILD` | 跳过编译 | `false` | Build |
-| `SKIP_E2E_TEST` | 跳过测试 | `false` | Test |
-| `PUSH_TO_GITHUB` | 成功后创建 PR(渐进式模式下每步都推) | `false` | Push |
-| `GITHUB_REPO` | PR 目标仓库 `owner/name` | `TecJesh/triton-ascend` | Push |
-| `LLVM_INSTALL_PREFIX` | LLVM 安装路径 | — | Build |
-| `CONDA_ENV` | Conda 环境名称 | `ta-upgrade` | Build / Test |
-| `NUM_PROCS` | pytest 并行 worker 数 | `16` | Test |
-| `AUTO_STASH` | 创建 work branch 前自动 stash | `false` | Merge |
-| `TA_PROGRESSIVE_MERGE` | 启用渐进式步骤合并 | `true` | Detect / Plan |
-| `TA_LINE_BUDGET` | 每步骤最大源码变更行数 | `1000` | Plan |
-| `TA_COMMIT_BUDGET` | commit 数预算基数(越小步骤越细) | `5` | Plan |
-
-> **步骤切分策略**:plan_steps 按**行数预算**和 **commit 数预算**两个维度同时切分,
-> 任一维度超限即开始新步骤。commit 数预算 = `max(1, round(TA_COMMIT_BUDGET * sqrt(TA_LINE_BUDGET / 1000)))`。
-> 默认 TA_COMMIT_BUDGET=5、TA_LINE_BUDGET=1000 → 每步最多 5 个 commit。
->
-> 如果觉得步骤太粗(commit 太多合在一步),减小 `TA_COMMIT_BUDGET`(如 3)。
-> 如果觉得步骤太细,增大 `TA_COMMIT_BUDGET`(如 10)。
+## 关键设计
----
+### 不可变状态传递
-## 输出文件结构
-
-```
-workspace/
-├── detect.json # 检测结果:merge-base、target、commit 列表
-├── steps.json # 步骤规划:每个步骤的 commit 范围、行数统计
-├── merge_result.json # 合并结果:分支名、冲突状态
-├── merge.log # git merge 原始输出
-├── build_result.json # 编译结果:各步骤通过/失败
-├── build.log # 编译原始输出
-├── test_result.json # 测试结果:通过/失败计数
-├── test-logs/
-│ ├── pytest.log # pytest 原始输出
-│ └── precommit.log # pre-commit 检查输出
-├── conflicts/ # 冲突文件快照(如有)
-│ └── _.conflict
-├── fixes/
-│ └── fix-/ # 每轮修复的日志
-│ └── opencode.log
-├── steps/ # 渐进式合并的每步产物
-│ └── step-/
-│ ├── commits.txt # 该步骤包含的 commit 列表
-│ ├── upstream.patch # 该步骤对应的上游 diff
-│ ├── changed_files.txt # 该步骤改动的文件列表
-│ ├── step_summary.md # AI 生成的步骤摘要
-│ ├── step_target.patch # 该步骤适配后的 diff
-│ ├── analysis.md # 修复诊断
-│ ├── review.md # 修复自查
-│ └── opencode.log # AI 调用日志
-├── step-0/ # 单步模式(TA_PROGRESSIVE_MERGE=false)产物
-│ ├── step_summary.md
-│ ├── step_target.patch
-│ ├── analysis.md
-│ ├── review.md
-│ └── opencode.log
-├── final_summary.md # 最终同步摘要
-├── final_target.patch # 完整 diff(累积所有步骤,用于 PR)
-├── FAILURE.md # 失败诊断(仅失败时)
-└── sync_meta.json # 同步元数据
-```
-
-> **注意**:渐进式模式下产物写入 `workspace/steps/step-/`;单步模式下
-> 产物写入 `workspace/step-0/`(向后兼容)。
-
----
-
-## 常见场景
+`WorkflowContext` 是 dataclass,通过 `ctx.copy_with(field=value)` 返回新实例。
+管道步骤从不原地修改 context,使得数据流完全显式且可测试。
-### 场景 1:无冲突同步
+### 配置集中管理
-```
-merge → (无冲突) → build → test → all pass → finalize → DONE
-```
+`TAConfig` dataclass 通过 `from_env()` 一次性读取所有环境变量。管道步骤接收
+config 作为第二个参数,不再散落各处的 `os.getenv()` 调用。
-最简单的情况。AI 只用于可能需要的修复,冲突解决直接跳过。
+### 日志系统
-### 场景 2:有冲突但无代码错误
+`TALogger` 提供统一的格式化输出:`header()`, `section()`, `status()`, `key_value()`,
+`step()`, `ai_call()`, `ai_result()`, `table()`, `elapsed()`。
-```
-merge → conflict → AI resolve(1次成功) → build → test → all pass → finalize → DONE
-```
+### Git 自动重试
-AI 解决了冲突,代码编译和测试都直接通过。
+`run_git()` 对网络操作(fetch, clone, push)自动重试 3 次,本地操作直接抛出异常。
-### 场景 3:需要多轮修复
+### 修复评价三层防护
-```
-merge → conflict → AI resolve(2次) → build pass → test fail
- → AI fix(Round 1) → build pass → test fail
- → AI fix(Round 2) → build pass → test all pass
- → commit fixes → finalize → DONE
-```
+1. **AI 自检** (prompt.md):提交前自查文件路径和修复根因
+2. **代码硬校验** (validate_fix):检查修改文件路径,不通过则 revert + 反馈
+3. **实际验证**:编译/测试结果
-### 场景 4:修复耗尽,同步失败
+详见 [fix-validation-flow.md](fix-validation-flow.md)
-```
-merge → AI resolve → build pass → test fail
- → AI fix(R1) → test fail
- → AI fix(R2) → test fail
- → AI fix(R3) → test fail
- → max_retries exhausted → UpgradeFailed → FAILURE.md
-```
+---
-work branch 保留,可以手动 `git checkout auto/upstream-sync-...` 继续排查。
+## 常见场景
-### 场景 5:渐进式多步骤合并(默认模式)
+### 无冲突同步
+merge → (无冲突) → build → test → pass → commit → DONE
-上游有 45 个 commit(其中 30 个涉及源码,共约 2800 行变更),
-自动切分为 3 个步骤:
+### 有冲突需修复
+merge → conflict → AI resolve → build fail → AI fix → build pass → test fail → AI fix → test pass → commit → DONE
-```
-detect(45 commits, 2800 lines) → plan(3 steps)
- → Step 1(12 commits, 950 lines): merge → (clean) → build → test → pass → commit
- → Step 2(10 commits, 980 lines): merge → conflict → AI resolve → build → test → pass → commit
- → Step 3(8 commits, 870 lines): merge → (clean) → build → test fail
- → AI fix(R1) → build → test pass → commit
- → finalize: 生成累积 patch(2800+ lines) + summary → DONE
-```
+### LLVM 版本变更
+merge → IR Patch Phase 1: build LLVM → fix compile errors → Phase 2: OP 分析 → 生成补丁 → 编译 LLVM → build TA → pytest → commit → DONE
-每个步骤只处理约 1000 行以内的变更,冲突范围可控、修复定位精确。
-即使某一步失败,前几步的进度已提交,不会全部丢失。
+### 修复耗尽
+build → AI fix (R1 rejected) → AI fix (R2) → build fail → ... → max_retries exhausted → UpgradeFailed
-### 场景 6:Dry-run 调试
-
-```bash
-SKIP_AI_ANALYSIS=true SKIP_BUILD=true SKIP_E2E_TEST=true \
-ta-kickoff --triton-ascend-path ... --triton-path ...
-```
-
-跳过 AI 和编译测试,只走 git merge 流程。适合验证工作流本身是否正常。
+work branch 保留,可手动排查。
---
-## 故障排查
-
-### work branch 创建失败
+## 输出文件
```
-[merge] ERROR: Working tree has uncommitted changes
+workspace/
+├── detect.json
+├── steps.json
+├── build_result.json / build.log
+├── test_result.json
+├── test-logs/ (pytest JUnit XML + 日志)
+├── llvm_build.log
+├── fixes/ (每轮 AI 修复日志)
+├── steps/step-N/ (每步产物)
+├── ir-analysis/ (IR 分析报告)
+├── final_summary.md
+├── final_target.patch
+├── pr_body.md
+└── FAILURE.md (仅失败时)
```
-**原因**:triton-ascend 仓库有未提交的修改。
+---
-**解决**:
-- 手动 `git stash` 暂存修改
-- 或设置 `AUTO_STASH=true` 让工作流自动 stash
+## 故障排查
### AI 后端不可用
+安装 opencode CLI 或设置 `AI_BACKEND=claude`,或 `SKIP_AI_ANALYSIS=true` 手动处理。
-```
-AI backend not available
-```
-
-**原因**:`opencode` 或 `claude` 不在 `$PATH` 中。
-
-**解决**:
-- 安装 opencode:参考 https://opencode.ai
-- 或设置 `AI_BACKEND=claude` 使用 claude CLI
-- 或设置 `SKIP_AI_ANALYSIS=true` 跳过 AI(需要手动解决冲突和修复)
-
-### 编译失败
-
-检查 `workspace/build.log` 中的错误。常见原因:
-- `LLVM_INSTALL_PREFIX` 未设置或指向错误路径
-- Conda 环境未激活或缺少依赖
+### LLVM 编译失败
+检查 `llvm_build.log`。确认 `LLVM_PROJECT_PATH` 和 `LLVM_INSTALL_PREFIX_SYNC` 正确。
### 测试持续 OOM
-
-1. 查看 `workspace/test-logs/pytest.log` 确认是 NPU OOM
-2. 减少并行度:`NUM_PROCS=8 ta-kickoff ...`
-3. OOM 是瞬时资源问题,多次重跑通常会消失
-4. 参考 `reference/npu-oom-handling.md`
+- 减少并行度: `NUM_PROCS=8 ta-kickoff`
+- OOM 是瞬时资源问题,工作流会自动降并发重跑
+- 参考 `reference/npu-oom-handling.md`
### PR 创建失败
-
-- 确认 `gh auth status` 显示已登录
-- 确认 `GITHUB_REPO` 格式正确(`owner/name`,不含 `https://`)
-
----
-
-## 架构说明
-
-### 为什么用单节点编排而不是 CrewAI 信号链?
-
-CrewAI 的 `@listen → @listen` 信号链在某些版本中不能正确传递返回值,
-导致下游节点收不到信号。因此 `execute_sync` 设计为**单个 `@router` 节点**,
-内部以普通 Python 方法调用的方式串联所有子步骤。
-
-在渐进式模式下,内部是一个 `while` 循环,对每个步骤执行完整的
-合并→解决→编译→测试→修复→提交 流程:
-
-```python
-@router(detect_commits)
-def execute_sync(self):
- while self.state.current_step < self.state.total_steps:
- step = self.state.steps[self.state.current_step]
- self._do_step_merge(step) # 合并该步骤的 end_commit
- if self.state.merge_has_conflicts:
- if not self._do_resolve_conflicts(): return UpgradeFailed
- if not self._do_build_and_fix_loop(): return UpgradeFailed
- self._do_commit_step(step) # 提交该步骤进度
- self.state.current_step += 1
- self._do_finalize() # 生成累积 patch 和 summary
- return UpgradeCompleted
-```
-
-这样保证了流转逻辑 100% 可控,不受 CrewAI 版本的信号路由行为影响。
-同时也实现了渐进式提交——即使后续步骤失败,已完成的步骤进度也不会丢失。
-
-### 为什么用 `git add -u` 而不是 `git add -A`?
-
-`git add -A` 会 stage 所有文件(包括 untracked 的新文件),可能把编译产物、
-测试日志、缓存文件等临时文件误提交到 git。`git add -u` 只 stage 已跟踪文件的修改,
-避免了这个风险。
-
-### 为什么不更新 version.txt?
-
-历史版本会在 `_do_finalize` 中写入 `version.txt`,但这个文件对工作流无实际用途,
-反而可能被误提交。已移除该逻辑。
+- 确认 `gh auth status` 已登录
+- 确认 `GITHUB_REPO` 格式正确
diff --git a/docs/workflow.md b/docs/workflow.md
index 44f0316..22d4b71 100644
--- a/docs/workflow.md
+++ b/docs/workflow.md
@@ -1,23 +1,164 @@
+# 单步模式工作流
+
```mermaid
flowchart TD
- A["ta-kickoff"]
- A --> B["Phase 0: Initialize"]
- B --> C["Phase 1: Detect Commits"]
- C -->|no new commits| D["Done: Already Up-to-Date"]
- C -->|new commits found| E["Phase 2A: Merge Upstream"]
- E -->|conflicts| F["Phase 2B: AI Resolve Conflicts"]
- E -->|no conflicts| G["Phase 2C: Build and Test"]
- F -->|resolved| G
- F -->|max retries| H["Failure: write FAILURE.md"]
- G --> G1["Build: setup.py install"]
- G1 -->|failed| I["AI Fix Code"]
- G1 -->|passed| G2["Test: pytest"]
- G2 -->|failed| I
- G2 -->|passed| J["Commit Fixes"]
- I -->|retry| G1
- I -->|max retries| H
- J --> K["Phase 2D: Finalize"]
- K --> L["Success"]
- L -->|push enabled| M["Push Branch and Create PR"]
- L -->|push disabled| N["Done: work branch kept"]
-```
\ No newline at end of file
+ A["ta-kickoff"] --> B["Phase 0: Prepare
克隆/配置 repo、remotes"]
+ B --> C["Phase 1: Detect
检测待合并的上游 commits"]
+ C -->|无新 commit| D["Done: Already Up-to-Date"]
+ C -->|有新 commit| E["Phase 2: Plan
按行数预算切分步骤"]
+ E --> F["Phase 3: Per-Step Loop"]
+
+ subgraph STEP["每个步骤 (while current_step < total_steps)"]
+ F1["Step A: Merge
git merge upstream commits"]
+ F1 --> F2{有冲突?}
+ F2 -->|是| F3["Step B: Resolve
AI 解决冲突 (max_retries)"]
+ F3 -->|未解决| FAIL["UpgradeFailed"]
+ F3 -->|已解决| F4
+ F2 -->|否| F4{LLVM 版本变更?}
+
+ F4 -->|是| F5["Step C: IR Patch Pipeline
Phase 1: 编译适配
Phase 2: IR 补丁生成+测试"]
+ F4 -->|否| F6["Step C: Build + Fix Loop
编译 - AI 修复 - 重编译"]
+
+ F5 --> F7{IR 补丁通过?}
+ F7 -->|否| FAIL
+ F7 -->|是| F8
+
+ F6 --> F8{编译通过?}
+ F8 -->|否| FAIL
+ F8 -->|是| F9["Step D: Test + Fix Loop
pytest - OOM 重跑 - AI 修复"]
+
+ F9 --> F10{测试通过?}
+ F10 -->|否| FAIL
+ F10 -->|是| F11["Step E: Commit
提交步骤进度"]
+
+ F11 --> F12["current_step += 1"]
+ F12 -->|还有步骤| F1
+ F12 -->|全部完成| G
+ end
+
+ STEP --> G["Phase 4: Finalize
生成 summary + cumulative patch"]
+
+ G -->|PUSH_TO_GITHUB=true| H["Push + Create PR"]
+ G -->|PUSH_TO_GITHUB=false| I["Done: work branch 保留"]
+
+ style FAIL fill:#d73,stroke:#333
+ style D fill:#4a9,stroke:#333
+ style I fill:#4a9,stroke:#333
+ style H fill:#4a9,stroke:#333
+```
+
+## IR Patch Pipeline (LLVM hash 变更时)
+
+```mermaid
+flowchart TD
+ subgraph IR["IR Patch Pipeline"]
+ P1["Phase 1: 编译适配"]
+ P1A["Build clean LLVM"] --> P1B["Build TA"]
+ P1B --> P1C{编译通过?}
+ P1C -->|否| P1D["AI 修复 (只改 third_party/ascend/)
AscendNPU-IR 专项文档"]
+ P1D -->|retries 耗尽| IR_FAIL["IR Pipeline Failed"]
+ P1D --> P1B
+ P1C -->|是| P2["Phase 2: IR 补丁循环"]
+
+ P2A["OP 分析 + LLVM 变更分析"] --> P2B["AI 生成 IR 补丁"]
+ P2B --> P2C["应用补丁 + 编译 LLVM
(失败时 AI 修复, 最多 10 次)"]
+ P2C -->|成功| P2D["Build TA + pytest"]
+ P2C -->|10 次耗尽| IR_FAIL
+ P2D -->|通过| IR_PASS["Pipeline Passed"]
+ P2D -->|失败| P2F["诊断: IR vs 代码"]
+ P2F -->|IR 问题| P2G["重新生成补丁 (最多 3 次)"]
+ P2G --> P2C
+ P2F -->|代码问题| P2H["AI 修复 → Rebuild → 重测"]
+ P2H --> P2D
+ end
+
+ style IR_FAIL fill:#d73,stroke:#333
+ style IR_PASS fill:#4a9,stroke:#333
+```
+
+## 修复代码评价机制
+
+```
+AI 修复
+ ↓
+Layer 1: AI 自检 (prompt.md step 6)
+ └─ 检查文件路径 → 不在允许目录则自行回退
+ ↓
+Layer 2: 代码校验 (validate_fix)
+ └─ 硬检查 → 不通过则 git revert + 反馈 + 不消耗 attempt
+ ↓
+Layer 3: 实际验证
+ └─ 编译/测试结果 → 失败则继续 fix loop
+```
+
+详见 [fix-validation-flow.md](fix-validation-flow.md)
+
+## 环境变量速查
+
+### 仓库 & 分支
+| 变量 | 默认值 | 说明 |
+|------|--------|------|
+| `TRITON_ASCEND_PATH` | (当前目录) | triton-ascend 本地路径 |
+| `TRITON_PATH` | (空) | 上游 triton 本地路径 |
+| `TRITON_TARGET_COMMIT` | (upstream HEAD) | 要合并的目标 commit |
+| `TA_BASE_BRANCH` | `upstream_sync` | 基线分支名 |
+| `TA_WORK_BRANCH_BASE` | `upstream-ascend` | 工作分支 base remote |
+
+### 构建 & 测试
+| 变量 | 默认值 | 说明 |
+|------|--------|------|
+| `BUILD_PROCS` / `MAX_JOBS` | 32 | 并行编译数(`MAX_JOBS` 为旧名,向后兼容) |
+| `TEST_PROCS` / `NUM_PROCS` | 16 | 并行 pytest 数(`NUM_PROCS` 为旧名) |
+| `LLVM_PROJECT_PATH` | `~/llvm-project` | LLVM 源码路径 |
+| `LLVM_INSTALL_PREFIX_SYNC` | `~/llvm-install-sync` | LLVM 安装路径 |
+| `LLVM_INSTALL_PREFIX` | (空) | LLVM 安装前缀(优先级高于上面) |
+| `SKIP_BUILD` | `false` | 跳过编译 |
+| `SKIP_E2E_TEST` | `false` | 跳过测试 |
+| `SKIP_BASELINE_LLVM` | `false` | 跳过基线 LLVM 编译(已有 LLVM 时使用) |
+| `SKIP_LLVM_REBUILD` | `false` | 跳过 LLVM 版本变更时的重编译 |
+
+### 测试目录(可插拔)
+| 变量 | 默认值 | 说明 |
+|------|--------|------|
+| `TA_TEST_DIR` | `third_party/ascend/unittest/pytest_ut` | 主测试目录 |
+| `TA_EXTRA_TEST_DIRS` | (空) | 额外 pytest 目录,逗号/空格分隔 |
+| `TA_TEST_COMMAND` | (空) | 自定义测试命令(在 pytest ut 之后额外执行) |
+
+### AI & 重试
+| 变量 | 默认值 | 说明 |
+|------|--------|------|
+| `AI_BACKEND` | `auto` | AI 后端(`opencode` / `claude`) |
+| `SKIP_AI_ANALYSIS` | `false` | 跳过 AI 调用 |
+| `TA_MAX_RETRIES` | 10 | AI 修复最大重试次数 |
+| `TA_LINE_BUDGET` | 1000 | 每步最大源码行数 |
+
+### PR
+| 变量 | 默认值 | 说明 |
+|------|--------|------|
+| `PUSH_TO_GITHUB` | `false` | 自动创建 PR |
+| `GITHUB_REPO` | `triton-lang/triton-ascend` | PR 目标仓库 |
+
+### 其他
+| 变量 | 默认值 | 说明 |
+|------|--------|------|
+| `TA_SINGLE_STEP_MODE` | `true` | 单步模式 |
+| `TA_RESUME` | `false` | 从缓存恢复(跳过已完成的步骤) |
+| `PYTHON` | `python3` | Python 解释器 |
+| `CONDA_ENV` | `ta-upgrade` | Conda 环境名 |
+| `TA_MAIN2MAIN_WORKSPACE` | `./workspace` | 工作区目录 |
+
+## CLI 参数
+
+```
+ta-kickoff [选项]
+
+ --triton-ascend-path PATH triton-ascend 仓库路径
+ --triton-path PATH 上游 triton 仓库路径
+ --target-commit SHA 目标 commit
+ --llvm-prefix PATH LLVM 安装前缀
+ --conda-env NAME Conda 环境名
+ --build-procs N 并行编译数
+ --test-procs N 并行测试数
+ --extra-test-dirs DIRS 额外测试目录(逗号分隔)
+ --test-command CMD 自定义测试命令
+```
diff --git a/pyproject.toml b/pyproject.toml
index 302cb82..3b1b7ea 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -1,20 +1,14 @@
[project]
name = "TA_main2main_workflow"
-version = "0.1.0"
-description = "TA_main2main_workflow — Triton-Ascend upstream sync using CrewAI"
+version = "0.2.0"
+description = "TA_main2main_workflow — Triton-Ascend upstream sync (modular pipeline)"
authors = [{ name = "Your Name", email = "you@example.com" }]
requires-python = ">=3.10,<3.14"
-dependencies = [
- "crewai[tools]==1.14.5"
-]
+dependencies = []
[project.scripts]
ta-kickoff = "TA_main2main_workflow.main:kickoff"
-ta-plot = "TA_main2main_workflow.main:plot"
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
-
-[tool.crewai]
-type = "flow"
diff --git a/src/TA_main2main_workflow/agent/prompt.md b/src/TA_main2main_workflow/agent/prompt.md
index cd45445..babb55b 100644
--- a/src/TA_main2main_workflow/agent/prompt.md
+++ b/src/TA_main2main_workflow/agent/prompt.md
@@ -33,30 +33,48 @@ The active mode is: {mode}
── report mode ────────────────────────────────────────────────────
- Trigger: {mode} is "report" (sync complete, generate summary).
-
- ALL data is in {error_logs} (JSON context file). Read it first.
-
- Generate a comprehensive report to {step_dir}/step_summary.md:
-
- ## 1. Executive Summary (总体概况)
- - Upstream commits synced, steps, conflicts, build/test fixes, AI rounds
-
- ## 2. Per-Step Analysis (逐步分析)
- - Commits merged, modules affected, conflicts and resolutions
- - Build errors: root causes and fixes (specific files and error messages)
- - Test failures: root causes and fixes (specific cases and fixes)
-
- ## 3. Fix Pattern Analysis (修复模式总结)
- - Cross-step patterns, API changes, recurring issues
- - Fixes that required multiple attempts
-
- ## 4. Recommendations (建议)
- - Preventative measures, fragile areas
-
- Rules: DO NOT modify source code. Write in Chinese (中文).
- Be specific with file paths, error messages, commit SHAs.
- 用中文写同步工作流总结报告
+ Trigger: {mode} is "report" (sync complete, generate PR description).
+
+ ALL context data is in {error_logs} (JSON context file). Read it first.
+
+ Generate the PR description to {step_dir}/step_summary.md with these
+ sections (write in English — this goes to a GitHub PR):
+
+ ## Summary
+ - Upstream Triton commits synced: {upstream_commits_count} commits
+ - Steps: {total_steps} step(s)
+ - Conflicts resolved: {conflict_files_resolved} file(s)
+ - AI build fixes: {build_fix_count} round(s)
+ - AI test fixes: {test_fix_count} round(s)
+ - Status: {final_status}
+
+ ## Background
+ - Source: triton-lang/triton (upstream)
+ - Target: this triton-ascend fork
+ - Why this sync is needed (e.g., keeping Ascend backend aligned with
+ upstream API changes, LLVM version updates, new features)
+
+ ## Changes
+ - Per-step breakdown: commits merged, source lines changed, modules affected
+ - Any LLVM version changes and IR compatibility patches applied
+ - List key files modified by AI fixes (from commit history or fix_errors)
+
+ ## Impact
+ - Which Ascend backend modules are affected (python/triton_ascend/,
+ third_party/ascend/, lib/Target/Ascend/)
+ - Any API/ABI changes that downstream consumers need to know about
+ - Test results: passed/failed counts per suite
+
+ ## Additional Notes
+ - Any known limitations or follow-up work needed
+ - Recommendations for reviewers
+
+ Rules:
+ - DO NOT modify source code
+ - Write in English (this is a GitHub PR description)
+ - Be specific: cite commit SHAs, file paths, error messages
+ - Read all context files in {error_logs} before writing
+ - Keep it concise but thorough — reviewers depend on this
── conflict mode ──────────────────────────────────────────────────
@@ -101,6 +119,31 @@ The active mode is: {mode}
Trigger: {mode} is "fix" (build or tests failed).
+ ═══ AscendNPU-IR compile errors (ascend_npu_ir_fix=true) ═══════════
+
+ When `ascend_npu_ir_fix` is "true", the build failure originates from
+ AscendNPU-IR (bishengir) code under `third_party/ascend/AscendNPU-IR/`.
+ LLVM version upgrades commonly break this code. You MUST read and
+ apply the patterns from BOTH of these references:
+
+ 1. {ascend_npu_ir_compat_ref}
+ — Complete catalog of all AscendNPU-IR LLVM 20→21→22 adaptations
+ (CMake compat macros, TableGen API changes, C++ API migrations,
+ dialect registration, pass infrastructure, build system fixes)
+
+ 2. {reference_dir}/02-llvm-version-adaptation-and-compile-fixes.md
+ — LLVM/MLIR API change table, compat macros, patch mechanism
+
+ For each AscendNPU-IR compile error:
+ - Match the error against the catalog in (1) to find the exact fix
+ pattern (e.g. getDirectSuperClasses() API change → __LLVM_MAJOR_VERSION_22_COMPATIBLE__)
+ - Apply the fix using the version-compat macros when available
+ - DO NOT modify third_party LLVM source directly — use compat macros
+ - If the error is NOT in the catalog, apply the general LLVM API
+ adaptation patterns from (2)
+
+ ═══════════════════════════════════════════════════════════════════════
+
Workflow:
1. Read structured error output from {error_logs}
2. Classify each failure:
@@ -125,19 +168,42 @@ The active mode is: {mode}
NPUIR updates) →
{reference_dir}/04-ir-compatibility-and-backend-adaptation.md
4. Apply minimal fixes:
- - Update imports when upstream moves modules
- - Update function signatures when upstream changes APIs
- - Update CMakeLists.txt when build configuration changes
- - Fix pytest assertions when expected behavior changes
- 5. Do NOT modify upstream triton code in python/triton/ unless it contains
- Ascend-specific changes (marked with triton_ascend imports or ascend checks)
- 6. Write fix summary to {step_dir}/step_summary.md
- 7. Write a ONE-LINE commit message to {step_dir}/commit_message.txt
- - Format: ": "
- - Types: fix, build, test, cmake, compat
- - Example: "fix: update AscendDotOp::build() signature for LLVM 22"
- - Example: "test: fix pytest assertion for renamed attribute getLhs→getA"
- - Keep under 72 characters, be specific about WHAT was fixed
+ - For BUILD / COMPILE errors: ONLY modify code under
+ {ascend_path}/third_party/ascend/. All other paths are read-only.
+ If an upstream API change broke the build, adapt the Ascend backend
+ code that depends on it.
+ - For TEST / PYTEST failures: FIRST try to fix in Ascend-specific
+ code under:
+ - {ascend_path}/third_party/ascend/
+ - {ascend_path}/python/triton/extension/
+ - {ascend_path}/python/triton/runtime/libentry.py
+ Most test failures can be resolved by adapting the Ascend backend
+ without touching upstream code. If — and only if — root cause
+ analysis shows the issue is inherently in upstream code with no
+ Ascend-side workaround, then apply a minimal targeted fix at the
+ specific point in the upstream file.
+ 5. Update imports and signatures when upstream changes APIs — adapt
+ the Ascend call sites, not the upstream declarations.
+ 6. SELF-REVIEW before returning:
+ - List every file you modified.
+ - For build fixes, verify it is under {ascend_path}/third_party/ascend/
+ - For test fixes, verify it is under one of:
+ {ascend_path}/third_party/ascend/
+ {ascend_path}/python/triton/extension/
+ {ascend_path}/python/triton/runtime/libentry.py
+ - If ANY modified file is outside these paths, REVERT that change
+ BEFORE returning — the fix will be rejected by the workflow.
+ - Confirm the fix directly addresses the root cause, not just
+ silences the error.
+ 7. Write fix summary to {step_dir}/step_summary.md
+ 8. Write a ONE-LINE commit subject to {step_dir}/commit_message.txt
+ - Describe WHAT was fixed, be specific (file/module and change)
+ - Keep under 72 characters
+ - Do NOT add a type prefix like "fix:" or "build:" — the workflow
+ will wrap it as [Sync](fix) automatically
+ - Good Example: "Update AscendDotOp::build() signature for LLVM 22"
+ - Good Example: "Fix pytest assertion for renamed attribute getLhs to getA"
+ - Bad Example: "fix: update AscendDotOp::build() signature" (redundant fix:)
- This line will be used as the git commit subject
Common failure patterns in Triton-Ascend:
@@ -147,6 +213,23 @@ The active mode is: {mode}
- third_party/nvidia/ changes → Ascend third_party/ascend/ may need matching updates
- CMakeLists.txt changes → Ascend CMake configuration may need adjusting
+ ── TEST-FAILURE-ONLY: LLVM/MLIR Op name swap in generated IR ──────────
+
+ ⚠️ APPLY ONLY WHEN FIXING TEST FAILURES (pytest / unit-test errors).
+ Do NOT apply this during compile-error fixing — build errors caused
+ by missing ToBufferOp/ToMemrefOp should be fixed by code adaptation
+ (see {reference_dir}/02-llvm-version-adaptation-and-compile-fixes.md).
+
+ bufferization::ToMemrefOp ↔ bufferization::ToBufferOp:
+ These two names have swapped across LLVM versions. If test logs
+ show the compiler cannot recognize `ToBufferOp` in generated IR,
+ the target LLVM uses `ToMemrefOp`. Fix: replace ALL occurrences
+ of the unrecognized Op name with the recognized one in the Ascend
+ backend (third_party/ascend/ and lib/Target/Ascend/).
+
+ grep -rn "ToBufferOp\|ToMemrefOp" {ascend_path}/third_party/ascend/ \
+ {ascend_path}/lib/Target/Ascend/ --include="*.cpp" --include="*.h"
+
━━━ REPOSITORIES ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
triton-ascend: {ascend_path}
@@ -287,107 +370,381 @@ The active mode is: {mode}
Target LLVM hash: {target_llvm_hash}
llvm-project repo: {llvm_project_path}
- ═══ HOW TO COMPARE — use git in the llvm-project repo ═══════════════════
+ ═══ CHANGE TYPE TAXONOMY — every OP MUST be checked for ALL 7 types ═══
+
+ For each OP, check these 7 dimensions. Mark `needs_patch: true` when
+ the change could cause IR generated by the target LLVM to be unparseable
+ by the old AscendNPU-IR / BishengIR compiler.
+
+ ┌─────────────────────────────────────────────────────────────────────┐
+ │ 1. OP_NAME_CHANGED — OP was renamed upstream │
+ │ Detect: grep "def " at target returns different name │
+ │ Impact: Ascend backend references old name → IR parse error │
+ │ needs_patch: true (add backward-compatible alias or op mapping) │
+ ├─────────────────────────────────────────────────────────────────────┤
+ │ 2. ASSEMBLY_FORMAT_CHANGED — assemblyFormat string differs │
+ │ Detect: diff the `let assemblyFormat = "...";` line │
+ │ cmd: git diff baseline..target -- │
+ │ Impact: new LLVM emits IR in format old parser cannot handle │
+ │ needs_patch: true │
+ ├─────────────────────────────────────────────────────────────────────┤
+ │ 3. ASSEMBLY_FORMAT_ADDED — OP gained assemblyFormat (previously │
+ │ used custom printer/parser or had no format at all) │
+ │ Detect: baseline lacks `let assemblyFormat`, target has it │
+ │ Impact: IR output switches from custom format to declarative; │
+ │ old parser may not understand the new format │
+ │ needs_patch: true (add backward-compatible custom printer) │
+ ├─────────────────────────────────────────────────────────────────────┤
+ │ 4. ATTRIBUTES_CHANGED — Op arguments/attributes added, removed, │
+ │ renamed, type-changed, or default-value-changed │
+ │ Detect: diff `let arguments = (ins ...);` block │
+ │ Sub-types: renamed, added, removed, type_changed, default_changed│
+ │ Impact: Ascend code references old attribute → compile error │
+ │ needs_patch: true if attribute is used in Ascend backend │
+ ├─────────────────────────────────────────────────────────────────────┤
+ │ 5. CUSTOM_PRINTER_PARSER_CHANGED — print()/parse() implementation │
+ │ differs between baseline and target │
+ │ Detect: diff the .cpp file containing print/parse methods │
+ │ cmd: git diff baseline..target -- mlir/lib/Dialect// │
+ │ Impact: IR text output/input format changes │
+ │ needs_patch: true │
+ ├─────────────────────────────────────────────────────────────────────┤
+ │ 6. CREATE_BUILDER_CHANGED — create()/build() method signature │
+ │ changed (params added/removed/reordered/retyped) │
+ │ Detect: diff `let builders = [...]` or build() methods in .cpp │
+ │ Impact: Ascend calls old signature → compile error │
+ │ needs_patch: true (add backward-compatible overload) │
+ ├─────────────────────────────────────────────────────────────────────┤
+ │ 7. TRAITS_CHANGED — Op traits added or removed │
+ │ Detect: diff `let traits = [...]` or template Traits<...> │
+ │ Impact: removed trait may break Ascend pass that depends on it │
+ │ needs_patch: true only if Ascend backend references the trait │
+ └─────────────────────────────────────────────────────────────────────┘
+
+ ═══ PER-OP COMPARISON PROCEDURE ═══════════════════════════════════════
+
+ For EVERY OP in ops_report.json, execute this procedure:
+
+ A. LOCATE the .td file:
+ grep -r "def " {llvm_project_path}/mlir/ --include="*.td"
+
+ B. GET both versions of the definition:
+ git -C {llvm_project_path} show {baseline_llvm_hash}:.td
+ git -C {llvm_project_path} show {target_llvm_hash}:.td
+
+ C. DIFF the two versions:
+ git -C {llvm_project_path} diff {baseline_llvm_hash}..{target_llvm_hash} -- .td
+
+ D. CROSS-REFERENCE with Ascend backend usage:
+ grep -r "" {ascend_path}/third_party/ascend/ --include="*.cpp" --include="*.h" -l
+ For each usage site, check whether the detected change breaks that code.
+
+ E. CLASSIFY every diff against the 7-type taxonomy above.
+ One OP can have MULTIPLE change types — record each in the
+ change_types array. An OP with ANY change automatically
+ gets needs_patch: true unless proven harmless.
+
+ ═══ OUTPUT JSON SCHEMA ════════════════════════════════════════════════
+
+ Output to `{step_dir}/changes_report.json`:
+
+ {{
+ "source_llvm_hash": "{baseline_llvm_hash}",
+ "target_llvm_hash": "{target_llvm_hash}",
+ "changes": [
+ {{
+ "op_name": "arith::CmpFOp",
+ "td_file": "mlir/include/mlir/Dialect/Arith/IR/ArithOps.td",
+ "cpp_file": "mlir/lib/Dialect/Arith/IR/ArithOps.cpp",
+ "change_types": ["attributes_changed", "create_builder_changed"],
+ "details": {{
+ "attributes_changed": {{
+ "added": ["fastmath: FastMathFlagsAttr (optional)"],
+ "removed": [],
+ "renamed": [],
+ "type_changed": [],
+ "default_changed": []
+ }},
+ "create_builder_changed": {{
+ "old_signature": "create(builder, location, predicate, lhs, rhs)",
+ "new_signature": "create(builder, location, predicate, lhs, rhs, fastmath)"
+ }}
+ }},
+ "ascend_usage_files": [
+ "third_party/ascend/lib/Conversion/ArithToHFusion/ArithToHFusion.cpp"
+ ],
+ "needs_patch": true,
+ "reason": "create() gained fastmath param; Ascend calls old 5-arg signature"
+ }},
+ {{
+ "op_name": "scf::ForOp",
+ "td_file": "mlir/include/mlir/Dialect/SCF/IR/SCFOps.td",
+ "cpp_file": null,
+ "change_types": ["assembly_format_added"],
+ "details": {{
+ "assembly_format_added": {{
+ "baseline": "hasCustomAssemblyFormat = 1 (custom printer/parser)",
+ "target": "let assemblyFormat = \\"...\\" (declarative format)"
+ }}
+ }},
+ "ascend_usage_files": [],
+ "needs_patch": true,
+ "reason": "new declarative format may emit IR old BishengIR cannot parse"
+ }},
+ {{
+ "op_name": "arith::AddIOp",
+ "td_file": "mlir/include/mlir/Dialect/Arith/IR/ArithOps.td",
+ "cpp_file": null,
+ "change_types": ["assembly_format_changed"],
+ "details": {{
+ "assembly_format_changed": {{
+ "old_format": "$attr `,` $lhs `,` $rhs attr-dict `:` type($result)",
+ "new_format": "$lhs `,` $rhs attr-dict `:` type($result)"
+ }}
+ }},
+ "ascend_usage_files": [],
+ "needs_patch": true,
+ "reason": "old format includes $attr prefix; BishengIR expects it"
+ }},
+ {{
+ "op_name": "linalg::MatmulOp",
+ "td_file": "mlir/include/mlir/Dialect/Linalg/IR/LinalgStructuredOps.td",
+ "cpp_file": null,
+ "change_types": ["op_name_changed"],
+ "details": {{
+ "op_name_changed": {{
+ "old_name": "linalg::MatmulOp",
+ "new_name": "linalg::MatmulTransposeOp"
+ }}
+ }},
+ "ascend_usage_files": [],
+ "needs_patch": true,
+ "reason": "Ascend references old MatmulOp name"
+ }},
+ {{
+ "op_name": "arith::ConstantOp",
+ "td_file": "mlir/include/mlir/Dialect/Arith/IR/ArithOps.td",
+ "cpp_file": null,
+ "change_types": [],
+ "details": {{}},
+ "ascend_usage_files": ["third_party/ascend/lib/Conversion/SomePass.cpp"],
+ "needs_patch": false,
+ "reason": "OP definition is identical across baseline and target"
+ }}
+ ],
+ "summary": {{
+ "total_ops_analyzed": 42,
+ "ops_needing_patch": 5,
+ "ops_unchanged": 37,
+ "by_change_type": {{
+ "op_name_changed": 1,
+ "assembly_format_changed": 1,
+ "assembly_format_added": 1,
+ "attributes_changed": 1,
+ "custom_printer_parser_changed": 0,
+ "create_builder_changed": 1,
+ "traits_changed": 0
+ }}
+ }}
+ }}
+
+ ═══ RULES ═══════════════════════════════════════════════════════════════
+
+ - Check ALL 7 change types for EVERY OP — do not stop at the first hit.
+ - Use `git show` / `git diff` in llvm-project — do NOT read the working
+ tree directly (it may be at an arbitrary commit).
+ - `needs_patch: false` ONLY when the OP definition is IDENTICAL across
+ both LLVM versions for all 7 dimensions.
+ - The "details" field MUST contain specific old-vs-new values for each
+ detected change_type — file paths, old/new signatures, diffs.
+ - Cross-reference with Ascend backend usage (grep in
+ {ascend_path}/third_party/ascend/) — if Ascend never references the
+ changed attribute/API, note it but still include the OP.
+ - If an OP in ops_report.json no longer exists at the target hash,
+ record it as op_name_changed with the old name and empty new_name.
+ - An OP that was checked and found unchanged across all 7 types still
+ goes in the output with change_types: [], needs_patch: false.
+
+ Reference:
+ {reference_dir}/02-llvm-version-adaptation-and-compile-fixes.md
+ {reference_dir}/04-ir-compatibility-and-backend-adaptation.md
+ {reference_dir}/05-ir-patch-generation-guide.md
+
+ Rules: DO NOT modify source code. Output ONLY the structured JSON report.
- For EVERY OP in ops_report.json, you MUST compare its .td definition at
- the two LLVM versions using git. Do NOT guess or skip any OP.
+── ir_generate_patch mode ──────────────────────────────────────────────
- Step A — verify both commits exist:
- cd {llvm_project_path}
- git cat-file -t {baseline_llvm_hash}
- git cat-file -t {target_llvm_hash}
+ Trigger: {mode} is "ir_generate_patch" (generate TA-side LLVM OP patches).
- Step B — find the .td file for each OP:
- Search for the OP's TableGen definition in mlir/include/:
- grep -r "def " mlir/include/ --include="*.td"
+ ═══ PATCH FIX RETRY (patch_error_type present) ═════════════════════════
- Step C — compare the definition at both versions:
- git show {baseline_llvm_hash}:mlir/include/.../.td
- git show {target_llvm_hash}:mlir/include/.../.td
- Then diff the two definitions.
+ When `patch_error_type` is "apply" or "build", a previously generated
+ patch FAILED. You are fixing it — NOT starting from scratch.
- Step D — also check for name changes (OP renamed):
- git diff {baseline_llvm_hash}..{target_llvm_hash} -- mlir/include/ | grep "^[-+].*def "
- This shows which OP definitions were added/removed between the two versions.
+ patch_error_type: {patch_error_type}
+ patch_error_msg: {patch_error_msg}
- Step E — for each OP, cross-reference with the Ascend backend usage:
- Check how the OP is used in {ascend_path}/third_party/ascend/lib/
- and {ascend_path}/lib/Target/Ascend/ — does the OP use create(),
- match(), or transformation patterns that depend on the old definition?
+ Fix strategy:
+ - **apply failure**: the patch does not apply cleanly to the target
+ LLVM commit. Re-read the target files via git show, check line
+ numbers and context, and regenerate the patch to match exactly.
+ - **build failure**: the patch applied but LLVM compilation failed.
+ Read the build error carefully — it tells you exactly which file
+ and line has the problem. Common causes:
+ * Wrong API for the target LLVM version (check git show output)
+ * Missing/extra parameters in create()/build() calls
+ * Type mismatches in attribute definitions
+ * Missing includes or forward declarations
+ Fix the relevant section of the patch while keeping all OTHER
+ sections intact — do NOT drop OPs that were correctly patched.
- ═══════════════════════════════════════════════════════════════════════════
- Workflow:
- 1. Read `{step_dir}/ops_report.json` for the list of OPs to check.
- 2. For each OP, examine its TableGen (.td) definition in the llvm-project
- at BOTH the baseline and target LLVM versions using the git commands above.
- The llvm-project repo is at: {llvm_project_path}
- Baseline (source): {baseline_llvm_hash}
- Target (current): {target_llvm_hash}
- 3. Record deltas per OP:
- - Name change (old_name → new_name)
- - assemblyFormat change (does the old format still parse?)
- - create() / builder parameter signature change
- - Attributes / getters renamed (e.g., getLhs → getA)
- - Traits added/removed
- - Custom printer/parser output format change
- 4. Output to `{step_dir}/changes_report.json`:
- {{
- "source_llvm_hash": "abc123",
- "target_llvm_hash": "def456",
- "changes": [
- {{
- "op_name": "arith::AddIOp",
- "change_type": "assemblyFormat_changed",
- "old_format": "...",
- "new_format": "...",
- "needs_patch": true,
- "reason": "new LLVM generates IR in format old NPU-IR cannot parse"
- }}
- ],
- "summary": {{
- "total_ops_analyzed": 42,
- "ops_needing_patch": 5,
- "renamed_ops": 1,
- "signature_changes": 3
- }}
- }}
+ ═══ PATCH SUPPLEMENT (adjust_mode=supplement) ═══════════════════════════
- Reference:
- {reference_dir}/02-llvm-version-adaptation-and-compile-fixes.md
- {reference_dir}/04-ir-compatibility-and-backend-adaptation.md
- {reference_dir}/05-ir-patch-generation-guide.md
+ When `adjust_mode` is "supplement", the existing patch was ALREADY
+ applied and LLVM was built successfully, but TESTS ARE FAILING with
+ IR compatibility errors. The patch is incomplete — it is missing
+ OP changes that the new LLVM version introduced.
- Rules: DO NOT modify source code. Output ONLY the structured JSON report.
+ ⚠️ This is NOT a full generation. You are SUPPLEMENTING an existing
+ working patch with additional OP changes. DO NOT start from scratch.
-── ir_generate_patch mode ──────────────────────────────────────────────
+ supplement_iteration: {supplement_iteration}
- Trigger: {mode} is "ir_generate_patch" (generate TA-side LLVM OP patches).
+ Workflow:
+ 1. READ the IR diagnosis ({previous_step_summary_path}) to
+ understand which OPs failed and the specific error symptoms.
+
+ 2. READ the FOCUSED CHANGE ANALYSIS ({focused_changes_path}).
+ This file was automatically generated by git-diffing the .td
+ definitions of each affected OP between baseline LLVM
+ ({baseline_llvm_hash}) and target LLVM ({target_llvm_hash}):
+ - Each affected OP has its .td file path and the git diff
+ - Use these diffs to understand exactly what changed upstream
+ - The diff covers all 7 change types: OP name, assemblyFormat,
+ attributes, custom printer/parser, create builder, traits
+
+ 3. READ the EXISTING PATCH ({ascend_patch_file}).
+ The patch content snippet shows the first 5000 bytes:
+ {patch_content_snippet}
+ Read the full file if the snippet is truncated.
+
+ 4. For EACH affected OP in the focused analysis, determine:
+ - Already in the existing patch? → the fix may be incorrect for
+ the current LLVM version → UPDATE it based on the focused diff
+ - NOT in the existing patch? → ADD a new section following the
+ exact same patterns (see KNOWN IR PATCH PATTERNS below and
+ the patch generation guide)
+
+ 5. SUPPLEMENT the patch — modify {ascend_patch_file} in-place:
+ - KEEP all existing entries that are correct — do NOT drop them
+ - ADD new entries for OPs that are missing
+ - FIX existing entries that are incorrect (wrong API, params, etc.)
+ - Use the focused diff to write precise, minimal changes
+ - Follow the exact git format-patch style of the existing patch
+
+ 6. SELF-CHECK before returning:
+ - Every affected OP in the focused analysis is addressed
+ - The patch still applies cleanly at {target_llvm_hash}
+ - No correct existing entries were removed
+
+ ═══════════════════════════════════════════════════════════════════════
Core strategy: patch TA-side LLVM so it generates IR compatible with the
UNMODIFIED AscendNPU-IR. NPU-IR is NOT touched — we cannot patch or
recompile it from the TA side.
- Workflow:
+ ⚠️ The patch MUST target LLVM at commit `{target_llvm_hash}`.
+ The baseline LLVM is `{baseline_llvm_hash}` — changes_report.json
+ describes what changed between these two LLVM versions.
+ The patch will be applied to `{llvm_project_path}` checked out at
+ `{target_llvm_hash}`, so all code modifications must be compatible
+ with the target LLVM's API.
+
+ Workflow (full generation — adjust_mode is NOT "supplement"):
1. Read `{step_dir}/changes_report.json` for ALL OPs needing patches.
- 2. Read the patch template:
+ 2. Read the patch generation guide:
+ `{reference_dir}/05-ir-patch-generation-guide.md`
+ — core strategy, patch patterns, format requirements, validation steps.
+ 3. Read the patch template for concrete format examples:
`{reference_dir}/ir_compatibility_patch_example.patch`
This demonstrates the direct OP patching approach (NOT BC/bytecode).
- 3. Generate a SINGLE complete `.patch` file that covers ALL OPs flagged
+ 4. For each OP, view the TARGET version of its .td/.cpp file using:
+ git -C {llvm_project_path} show {target_llvm_hash}:mlir/include/.../.td
+ git -C {llvm_project_path} show {target_llvm_hash}:mlir/lib/.../.cpp
+ Do NOT read the working tree directly — the checked-out commit may
+ differ from `{target_llvm_hash}`.
+ 5. Generate a SINGLE complete `.patch` file that covers ALL OPs flagged
with `needs_patch: true` in one unified patch. For each OP:
- - Locate its .td / .cpp file in `{llvm_project_path}/mlir/`
- - Apply the appropriate strategy by change type:
+ - Apply the appropriate strategy by change type (per the guide):
— OP renamed: add a backward-compatible alias (old name → new name)
— assemblyFormat changed: modify to also accept/emit old format
— create() params changed: add overload/defaults for old signature
— Pass option renamed: add old option name as alias
- 4. Write the single patch directly to `{ascend_patch_file}` (modify
+ — attributes changed: add backward-compat getter/wrapper
+ — custom printer/parser changed: preserve old output format
+ 6. Write the single patch directly to `{ascend_patch_file}` (modify
the existing file in-place):
- Follow `git format-patch` style with proper headers
- - Apply cleanly to `{llvm_project_path}` as one atomic change
+ - Apply cleanly to `{llvm_project_path}` at `{target_llvm_hash}` as
+ one atomic change
- Cover every OP in changes_report — do NOT leave any out
- Completeness requirement: the generated patch MUST be as complete as
+ ═══ KNOWN IR PATCH PATTERNS ═══════════════════════════════════════════
+
+ Apply these specific fixes when the corresponding OP appears in
+ changes_report.json with needs_patch: true.
+
+ ── empty-properties / assume-op rejection ─────────────────────────────
+
+ Symptom: the old NPU-IR compiler rejects IR containing an OP followed
+ by an empty inline property dict (e.g. the new LLVM prints the OP
+ with a trailing empty dict while the old parser only expects the OP
+ name without properties). This commonly affects ops whose TableGen
+ definition gained `useCustomPropertiesEncoding` or whose printer now
+ emits an inline property dict.
+
+ Fix: in the affected OP's custom printer (print() method in the
+ corresponding .cpp file under mlir/lib/Dialect/), replace the
+ printGenericOp block that emits inline attributes:
+
+ // BEFORE (new LLVM — emits empty dict that old parser rejects):
+ void ::print(OpAsmPrinter &p) {{
+ p.printGenericOp(*this); // ← emits attributes inline
+ }}
+
+ // AFTER (backward-compatible — filters empty properties):
+ void ::print(OpAsmPrinter &p) {{
+ SmallVector filtered;
+ for (NamedAttribute attr : (*this)->getAttrs()) {{
+ if (auto prop = dyn_cast<::mlir::Properties>(attr.getValue())) {{
+ if (prop.isEmpty()) continue; // skip empty properties
+ }}
+ filtered.push_back(attr);
+ }}
+ p << " ";
+ p.printAttribute(DictionaryAttr::get(getContext(), filtered));
+ }}
+
+ If the OP uses `printGenericOp` directly (no custom printer), you
+ must ADD a custom printer that filters out empty properties, AND
+ set `let hasCustomAssemblyFormat = 1;` / remove `let assemblyFormat`
+ in the .td file.
+
+ ── assume-op specific fix ─────────────────────────────────────────────
+
+ This fix applies ONLY to `llvm.assume` / `LLVM::AssumeOp`. Do NOT
+ apply it to any other OP — other ops have their own handling.
+
+ If `LLVM::AssumeOp`'s custom printer has these three lines:
+ os << " <";
+ Impl::printAttribute(prop);
+ os << '>';
+ Simply comment them out (no replacement needed). This suppresses
+ the inline attribute printing that produces the empty dict the old
+ parser cannot handle.
+
+ ─────────────────────────────────────────────────────────────────────
possible. Missing even one OP will cause the outer loop to retry
(costly: LLVM rebuild takes ~2 hours). Review changes_report
thoroughly before writing the patch — every `needs_patch: true` OP
diff --git a/src/TA_main2main_workflow/flow.py b/src/TA_main2main_workflow/flow.py
index f74adbe..ed9c83a 100644
--- a/src/TA_main2main_workflow/flow.py
+++ b/src/TA_main2main_workflow/flow.py
@@ -1,715 +1,170 @@
-"""CrewAI Flow — Triton-Ascend main2main upstream sync (merge-based).
+"""TA Main2Main Workflow — Triton-Ascend upstream sync orchestrator.
-Node order:
- initialize → detect_commits → execute_sync → push_to_github / handle_failure
+Assembles pipeline steps for single-step mode::
-The flow uses a single orchestration node (execute_sync) that internally
-runs merge → AI resolve conflicts → build → test → AI fix in a loop.
-This avoids relying on CrewAI @listen → @listen signal chaining which
-fails to propagate return values in some CrewAI versions.
-
-ALL progress is printed to the local console — no CrewAI web UI needed.
-AI (opencode or claude) is invoked via subprocess for conflict resolution
-and test fixing.
+ prepare → detect → plan → [build_baseline_llvm] → for each step:
+ merge → [resolve] →
+ if LLVM hash changed: per_step_ir_patch (apply existing → build LLVM →
+ build TA → test → supplement IR → loop)
+ else: build_and_fix_loop → test_and_fix_loop
+ → commit
+ → finalize → [push_pr]
"""
-import json
-import os
-import shutil
-import subprocess
-import time
-from pathlib import Path
-from typing import Literal
-
-from pydantic import BaseModel
-
-from crewai.flow import Flow, listen, start, router
+from __future__ import annotations
-from TA_main2main_workflow.agent.opencode_adapter import AIResult, run_opencode_adapter
-from TA_main2main_workflow.scripts.build_test import build_triton_ascend, run_tests
-from TA_main2main_workflow.scripts.detect_commits import detect
-from TA_main2main_workflow.scripts.merge_upstream import run_merge, run_merge_incremental
-from TA_main2main_workflow.scripts.plan_steps import run_plan
-from TA_main2main_workflow.scripts.pre_ci_check import run_pre_ci_check, cleanup_temp_files
-from TA_main2main_workflow.scripts.push_to_github import (
- push_and_create_pr,
-)
+from pathlib import Path
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils.git import run_git
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.tracker import timed, total_elapsed
from TA_main2main_workflow.utils import (
- BUILD_LOG_FILE, BUILD_RESULT_FILE, CONFLICT_LOG_DIR,
- EACH_STEP_SUMMARY_FILE, EACH_STEP_TARGET_PATCH_FILE,
- FINAL_SUMMARY_FILE, FINAL_TARGET_PATCH_FILE, FIX_LOG_DIR,
- HasNewCommits, HasNoNewCommits,
- STEPS_DIR, STEPS_FILE, LINE_BUDGET,
- TEST_RESULT_FILE, UpgradeCompleted, UpgradeFailed,
- WORKSPACE_DIR, has_merge_conflicts, run_git, get_conflict_files,
- commit_submodule, push_submodule, submodule_has_changes,
- IR_ANALYSIS_DIR, IR_OPS_REPORT_FILE,
- IR_CHANGES_REPORT_FILE, IR_DIAGNOSIS_FILE, IR_MAX_ITERATIONS,
- ENV_SINGLE_STEP_MODE, ENV_BASE_BRANCH, get_base_branch_ref, LLVM_CHANGE_ANALYSIS_DIR,
- print_header, print_section, print_step, print_status, print_info,
- print_warn, print_error, print_key_value,
- print_flow_progress, print_conflict_list, print_summary_table,
- print_ai_call_info, print_ai_result, print_elapsed_total,
- start_timer, stop_timer,
+ UpgradeCompleted,
+ UpgradeFailed,
+ HasNewCommits,
+ HasNoNewCommits,
+ WORKSPACE_DIR,
)
-
-_REFERENCE_DIR = str(Path(__file__).parent / "reference")
-
-# Baseline LLVM version that Ascend backend OP usage is built against.
-# IR compatibility patches bridge from this version to the target LLVM.
-_ASCEND_BASELINE_LLVM_HASH = "b5cc222d7429fe6f18c787f633d5262fac2e676f"
-
-
-def _llvm_project_path() -> Path:
- """Return the resolved llvm-project path (expands ~ and $HOME)."""
- return Path(os.path.expanduser(
- os.getenv("LLVM_PROJECT_PATH", "~/llvm-project")))
-
-
-def _llvm_install_prefix() -> Path:
- """Return the resolved LLVM install prefix (expands ~ and $HOME)."""
- return Path(os.path.expanduser(
- os.getenv("LLVM_INSTALL_PREFIX_SYNC", "~/llvm-install-sync")))
-
-
-class TA_Main2MainState(BaseModel):
- triton_ascend_path: str = ""
- triton_path: str = ""
- target_commit: str = ""
- test_log_dir: str = ""
-
- merge_base: str = ""
- ascend_head: str = ""
- work_branch: str = ""
- original_branch: str = ""
-
- upstream_commits_count: int = 0
- merge_has_conflicts: bool = False
- conflict_files: list = []
-
- build_passed: bool = False
- test_passed: bool = False
-
- retry_count: int = 0
- max_retries: int = 10
- fix_errors: list = []
-
- # ── Per-step tracking for sync report ──
- build_fix_count: int = 0 # AI fix attempts for build failures
- test_fix_count: int = 0 # AI fix attempts for test failures
- conflict_files_resolved: int = 0 # Total merge conflicts resolved
- step_details: list = [] # Per-step breakdown for report
- fix_attempts: list = [] # Detailed fix attempt records
-
- final_status: str = ""
- pr_url: str = ""
-
- llvm_prefix: str = ""
- conda_env: str = ""
- test_dir: str = "third_party/ascend/unittest/pytest_ut"
- num_procs: int = 16
-
- # ── Progressive step-by-step merge ──
- steps: list = []
- total_steps: int = 0
- current_step: int = 0
- step_start_ascend_head: str = "" # ascend HEAD before current step
- progressive_merge: bool = True
- step_pr_descriptions: list = [] # accumulated step descriptions for PR body
-
- # ── IR Patch Loop State ──
- ir_analysis_done: bool = False
- ir_ops_report: dict = {}
- ir_changes_report: dict = {}
- ir_patches: list = []
- ir_patch_iteration: int = 0
- ir_max_iterations: int = 3
- ir_issues_found: int = 0
- ir_fix_count: int = 0
- llvm_hash_changed: bool = False
-
- # ── Pytest State ──
- pytest_passed: bool = False
- test_failures_by_python: dict = {}
- ir_loop_details: list = []
-
- summary_rows: list = []
-
-
-class TA_Main2MainFlow(Flow[TA_Main2MainState]):
-
- def __init__(self, **kwargs):
- super().__init__(**kwargs)
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Workspace info helper — prints paths, branches, and git status
- # ═══════════════════════════════════════════════════════════════════════════
-
- def _print_workspace_info(self, label: str = "") -> None:
- """Print all relevant repo paths, current branches, and git status.
-
- Called at key workflow steps to provide full visibility into the
- workspace state — which repos are in play, what branches they're on,
- and whether there are uncommitted changes.
- """
- header = f"Workspace Info{f' — {label}' if label else ''}"
- print_section(header)
-
- # ── Resolve paths ──
- llvm_proj = _llvm_project_path()
- llvm_install = _llvm_install_prefix()
- ascend_str = self.state.triton_ascend_path
- triton_str = self.state.triton_path
-
- # ── Print all relevant paths ──
- print_key_value("LLVM_PROJECT_PATH", str(llvm_proj))
- print_key_value("LLVM_INSTALL_PREFIX_SYNC", str(llvm_install))
- if self.state.llvm_prefix:
- print_key_value("LLVM_INSTALL_PREFIX", self.state.llvm_prefix)
- if ascend_str:
- print_key_value("TRITON_ASCEND_PATH", ascend_str)
- if triton_str:
- print_key_value("TRITON_PATH", triton_str)
-
- # ── Print git branch + status for each repo ──
- repos: list[tuple[str, Path]] = []
- if ascend_str:
- ap = Path(ascend_str)
- if ap.exists():
- repos.append(("triton-ascend", ap))
- if triton_str:
- tp = Path(triton_str)
- if tp.exists():
- # Skip triton if it's the same directory as triton-ascend
- if not ascend_str or tp != Path(ascend_str):
- repos.append(("triton", tp))
- if llvm_proj.exists():
- repos.append(("llvm-project", llvm_proj))
-
- for repo_label, repo_path in repos:
- try:
- branch = run_git(repo_path, "branch", "--show-current").strip()
- print_key_value(f"{repo_label} branch", branch)
- status = run_git(repo_path, "status", "--porcelain").strip()
- if status:
- lines = status.splitlines()
- print_info(
- f"{repo_label} uncommitted changes ({len(lines)} files):"
- )
- for line in lines[:10]:
- print(f" {line}")
- if len(lines) > 10:
- print(f" ... and {len(lines) - 10} more")
- else:
- print_info(f"{repo_label} status: clean")
- except Exception as e:
- print_warn(f"Could not get git info for {repo_label}: {e}")
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Mode dispatch — supports full (CrewAI), merge-only, and fix-only modes
- # ═══════════════════════════════════════════════════════════════════════════
-
- def kickoff(self, inputs: dict | None = None):
- """Override CrewAI Flow.kickoff() to support TA_MODE dispatch.
-
- TA_MODE values:
- full — Original CrewAI flow: merge → resolve → build → test → fix → PR
- merge — Merge + AI resolve only, push work branch, skip build/test.
- Used on ubuntu-latest CI to prepare the work branch before
- NPU testing.
- fix — AI fix on an existing work branch. Reads error logs from
- TA_ERROR_LOGS_PATH, runs AI fix, commits & pushes.
- """
- mode = os.getenv("TA_MODE", "full")
- if os.getenv(ENV_SINGLE_STEP_MODE, "false").lower() == "true":
- return self._run_single_step_mode(inputs)
- elif mode == "merge":
- return self._run_merge_mode(inputs)
- elif mode == "fix":
- return self._run_fix_mode(inputs)
- else:
- return super().kickoff(inputs=inputs)
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Mode: merge — AI merge + resolve ONE step, then push work branch
- # ═══════════════════════════════════════════════════════════════════════════
-
- def _run_merge_mode(self, inputs: dict | None) -> str:
- """Merge + AI-resolve for ONE progressive step. Push work branch, no build/test.
-
- Used in CI (ubuntu-latest) as the merge phase of the per-step pipeline.
- Each call merges exactly one step's batch of upstream commits. The CI
- workflow orchestrates the per-step loop:
-
- For each step N:
- → ta-kickoff --mode=merge (merges step N, resolves conflicts, pushes)
- → NPU build+test
- → AI fix retries (if needed)
- → advance to step N+1
-
- Env vars:
- TA_CURRENT_STEP — which step index to merge (0-based, default 0).
- Step 0 does full init + detect + plan first.
- Step N>0 resumes from an existing work branch.
- """
- current_step = int(os.getenv("TA_CURRENT_STEP", "0"))
-
- # ── Apply inputs to state ──
- if inputs:
- for key, value in inputs.items():
- if hasattr(self.state, key):
- setattr(self.state, key, value)
-
- self._print_workspace_info(f"Merge Mode — step {current_step}")
-
- # ── Force skip build/test in merge mode ──
- os.environ["SKIP_BUILD"] = "true"
- os.environ["SKIP_E2E_TEST"] = "true"
-
- ascend_path = Path(self.state.triton_ascend_path)
-
- if current_step == 0:
- # ── First step: full init + detect + plan ──
- self.initialize()
- result = self.detect_commits()
- if result == HasNoNewCommits:
- print_info("No new commits — nothing to merge")
- metadata_dir = WORKSPACE_DIR / "merge-metadata"
- metadata_dir.mkdir(parents=True, exist_ok=True)
- (metadata_dir / "no_changes.txt").write_text("true", encoding="utf-8")
- self.state.summary_rows.append(
- ("MERGE PHASE", "SKIP", "No new upstream commits")
- )
- print_summary_table(self.state.summary_rows)
- self.state.final_status = UpgradeCompleted
- return UpgradeCompleted
-
- # Store the plan for subsequent steps
- self._write_step_plan()
- else:
- # ── Resume: checkout existing work branch ──
- work_branch = os.getenv("TA_WORK_BRANCH", self.state.work_branch)
- if not work_branch:
- print_error("TA_WORK_BRANCH is required for TA_CURRENT_STEP > 0")
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- # Restore state from work branch metadata
- self.state.work_branch = work_branch
- self.state.triton_ascend_path = (
- self.state.triton_ascend_path
- or os.getenv("TRITON_ASCEND_PATH")
- or str(Path.cwd())
- )
- self.state.target_commit = (
- self.state.target_commit or os.getenv("TRITON_TARGET_COMMIT", "")
- )
-
- # Read step plan saved from step 0
- plan_file = WORKSPACE_DIR / "merge-metadata" / "step_plan.json"
- if not plan_file.exists():
- print_warn("Step plan file not found — re-detecting commits")
- # Lightweight re-init without full initialize
- self.state.triton_ascend_path = self.state.triton_ascend_path or str(Path.cwd())
- self.state.triton_path = os.path.expanduser(
- os.getenv("TRITON_PATH", self.state.triton_ascend_path))
- ascend_path = Path(self.state.triton_ascend_path)
- # Fetch and checkout work branch
- try:
- run_git(ascend_path, "fetch", "origin", work_branch)
- except Exception:
- pass
- run_git(ascend_path, "checkout", work_branch)
- result = self.detect_commits()
- if result == HasNoNewCommits:
- self.state.final_status = UpgradeCompleted
- return UpgradeCompleted
- self._write_step_plan()
- else:
- import json
- plan_data = json.loads(plan_file.read_text(encoding="utf-8"))
- self.state.total_steps = plan_data["total_steps"]
- self.state.steps = plan_data["steps"]
- self.state.upstream_commits_count = plan_data.get("upstream_commits_count", 0)
-
- # Minimal init for resume
- self.state.triton_path = os.path.expanduser(
- os.getenv("TRITON_PATH", str(ascend_path)))
- # Fetch and checkout work branch
- try:
- run_git(ascend_path, "fetch", "origin", work_branch)
- except Exception:
- pass
- run_git(ascend_path, "checkout", work_branch)
-
- # ── Validate step index ──
- if current_step >= self.state.total_steps:
- print_info(f"current_step={current_step} >= total_steps={self.state.total_steps} — "
- f"all steps already merged")
- metadata_dir = WORKSPACE_DIR / "merge-metadata"
- metadata_dir.mkdir(parents=True, exist_ok=True)
- (metadata_dir / "all_steps_done.txt").write_text("true", encoding="utf-8")
- self._write_merge_metadata()
- self.state.final_status = UpgradeCompleted
+from TA_main2main_workflow.pipeline.prepare import prepare
+from TA_main2main_workflow.pipeline.detect import run_detect
+from TA_main2main_workflow.pipeline.plan import run_plan, llvm_hash_changed_after_merge
+from TA_main2main_workflow.pipeline.merge import merge_upstream_commit
+from TA_main2main_workflow.pipeline.resolve import resolve_conflicts
+from TA_main2main_workflow.pipeline.build import build_and_fix_loop
+from TA_main2main_workflow.pipeline.test import test_and_fix_loop
+from TA_main2main_workflow.pipeline.commit import commit_step
+from TA_main2main_workflow.pipeline.finalize import finalize
+from TA_main2main_workflow.pipeline.ir_patch import build_baseline_llvm, per_step_ir_patch
+
+log = get_logger(__name__)
+
+
+class TA_Main2MainFlow:
+ """Orchestrator — builds context, runs pipeline steps, handles PR.
+
+ Single-step mode is the only supported mode. Each step runs the full
+ pipeline: merge → resolve conflicts → build → fix → test → fix → commit.
+ """
+
+ def __init__(self, config: TAConfig | None = None) -> None:
+ self.config = config or TAConfig.from_env()
+
+ def run(self) -> str:
+ """Execute the full sync pipeline. Returns UpgradeCompleted or UpgradeFailed."""
+ log.header("Triton-Ascend Upstream Sync (Single-Step Mode)")
+ log.key_value("AI Backend", self.config.ai_backend)
+ log.key_value("Max Retries", str(self.config.max_retries))
+ log.key_value("Line Budget", str(self.config.line_budget))
+
+ # ── Phase 0: Prepare workspace ──────────────────────────────
+ with timed("prepare"):
+ ctx = prepare(WorkflowContext(), self.config)
+
+ # ── Phase 1: Detect ─────────────────────────────────────────
+ log.header("Phase 1: Detect Upstream Commits")
+ with timed("detect"):
+ ctx = run_detect(ctx, self.config)
+ if not ctx.has_new_commits:
+ log.status(True, "Already up to date")
+ ctx = ctx.copy_with(final_status=UpgradeCompleted)
+ ctx.summary_rows.append(("Detect", "SKIP", "No new upstream commits"))
return UpgradeCompleted
-
- step = self.state.steps[current_step]
- step_id = step["id"]
- is_last_step = (current_step == self.state.total_steps - 1)
- self.state.current_step = current_step
- self.state.retry_count = 0
-
- print_header(
- f"Step {current_step + 1}/{self.state.total_steps}: {step_id}"
- )
- print_key_value("commits in step", str(step["commit_count"]))
- print_key_value("end commit", step["end_commit"][:12])
- print_key_value("is last step", str(is_last_step))
-
- # ── Work-branch guard ──
- if current_step > 0 and self.state.work_branch:
- current_branch = run_git(ascend_path, "branch", "--show-current").strip()
- if current_branch != self.state.work_branch:
- print_warn(
- f"Expected work branch '{self.state.work_branch}' "
- f"but on '{current_branch}' — switching"
- )
- run_git(ascend_path, "checkout", self.state.work_branch)
-
- self.state.step_start_ascend_head = run_git(
- ascend_path, "rev-parse", "HEAD"
- ).strip()
-
- # ── Step A: git merge this step's commits ──
- merge_result = self._do_step_merge(step)
- if merge_result == UpgradeFailed:
- self.state.final_status = UpgradeFailed
- self._write_merge_metadata()
+ log.status(True, f"Found {ctx.upstream_commits_count} upstream commits")
+
+ # ── Phase 2: Plan ───────────────────────────────────────────
+ log.header("Phase 2: Plan Steps")
+ with timed("plan"):
+ ctx = run_plan(ctx, self.config)
+ log.status(True, f"Planned {ctx.total_steps} step(s)")
+
+ # ── Phase 2.5: Build baseline LLVM (pre-merge) ──────────────
+ log.header("Build Baseline LLVM")
+ with timed("baseline-llvm"):
+ ctx = build_baseline_llvm(ctx, self.config)
+ if not ctx.build_passed:
+ log.error("Baseline LLVM build failed — cannot proceed")
return UpgradeFailed
+ log.status(True, "Baseline LLVM ready")
- # ── Step B: AI resolve conflicts ──
- if self.state.merge_has_conflicts:
- if not self._do_resolve_conflicts():
- self.state.final_status = UpgradeFailed
- self._write_merge_metadata()
- return UpgradeFailed
-
- # ── Step C: Skip build/test (NPU CI runs these) ──
- print_header("Build & Test — Merge Mode")
- print_info(f"Merge mode: deferring build/test for step {step_id} to NPU CI")
- self.state.build_passed = True
- self.state.test_passed = True
- self.state.summary_rows.append(("Build", "DEFER", "Runs on NPU CI"))
- self.state.summary_rows.append(("Tests", "DEFER", "Runs on NPU CI"))
-
- # ── Step D: Commit step merge progress ──
- self._do_commit_step(step)
+ # ── Phase 3: Per-step loop ──────────────────────────────────
+ log.header("Single-Step Mode — Per-Step Full Pipeline")
+ log.key_value("Total steps", str(ctx.total_steps))
- # Record step description
- desc = (
- f"✅ **{step_id}**: {step['commit_count']} commits, "
- f"end_commit=`{step['end_commit'][:12]}`, "
- f"source lines={step.get('source_changed_lines', '?')}"
- )
- self.state.step_pr_descriptions.append(desc)
- print_status(True, f"Step {step_id} merge committed")
-
- # ── Push work branch ──
- self._push_work_branch_to_remote()
-
- # ── Write metadata for CI orchestration ──
- self._write_merge_metadata()
-
- # Print summary
- print_header(f"Merge Phase Complete — Step {step_id}")
- print_key_value("Work branch", self.state.work_branch)
- print_key_value("Current step", f"{current_step + 1}/{self.state.total_steps}")
- print_key_value("Target commit", self.state.target_commit[:12])
- print_key_value("Is last step", str(is_last_step))
- print_info(f"Pushed to origin/{self.state.work_branch}")
- if is_last_step:
- print_info("This is the last step — PR will be created if tests pass")
- else:
- next_step_id = self.state.steps[current_step + 1]["id"]
- print_info(f"Next: NPU tests on this step, then merge step {next_step_id}")
-
- self.state.summary_rows.append(
- ("MERGE PHASE", "PASS", f"Step {step_id}, branch: {self.state.work_branch}")
- )
- print_summary_table(self.state.summary_rows)
-
- self.state.final_status = UpgradeCompleted
- return UpgradeCompleted
-
- def _write_step_plan(self) -> None:
- """Persist the step plan so subsequent merge-mode calls can resume."""
- import json
- metadata_dir = WORKSPACE_DIR / "merge-metadata"
- metadata_dir.mkdir(parents=True, exist_ok=True)
- plan_data = {
- "total_steps": self.state.total_steps,
- "steps": self.state.steps,
- "upstream_commits_count": self.state.upstream_commits_count,
- }
- (metadata_dir / "step_plan.json").write_text(
- json.dumps(plan_data, indent=2, ensure_ascii=False), encoding="utf-8"
- )
- print_info(f"Step plan saved: {self.state.total_steps} step(s)")
-
- def _push_work_branch_to_remote(self) -> None:
- """Push the work branch to origin so NPU CI can access it."""
- ascend_path = Path(self.state.triton_ascend_path)
-
- # Check we're on the work branch
- current = run_git(ascend_path, "branch", "--show-current").strip()
- if current != self.state.work_branch:
- run_git(ascend_path, "checkout", self.state.work_branch)
-
- # ── Configure git auth (same logic as push_to_github._ensure_gh_auth) ──
- self._setup_git_auth_for_push(ascend_path)
-
- # ── Push AscendNPU-IR submodule first ──
- self._push_submodule_if_needed()
-
- print_header("Push Work Branch")
- try:
- run_git(ascend_path, "push", "-u", "origin", self.state.work_branch)
- print_status(True, f"Pushed {self.state.work_branch} to origin")
- self.state.summary_rows.append(
- ("Push branch", "PASS", self.state.work_branch)
- )
- except Exception as e:
- print_error(f"Failed to push work branch: {e}")
- # Try with force if normal push fails (e.g., branch exists from prior run)
- try:
- print_warn("Retrying with --force...")
- run_git(
- ascend_path, "push", "-u", "--force",
- "origin", self.state.work_branch,
- )
- print_status(True, f"Force-pushed {self.state.work_branch}")
- except Exception:
- print_error("Force push also failed")
- raise
-
- def _setup_git_auth_for_push(self, repo: Path) -> None:
- """Configure git authentication for pushing to GitHub.
-
- 1. Login gh CLI explicitly against github.com (needed when git
- remotes point to a proxy host that gh doesn't recognize).
- 2. Run 'gh auth setup-git' to configure the git credential helper.
- 3. Rewrite the origin URL to embed the token so git push works
- even through url.insteadOf proxy rewriting.
- """
- gh_token = os.getenv("GH_TOKEN", "")
- if gh_token:
- print_info("GH_TOKEN set — configuring git credential helper")
-
- # Explicit gh login against github.com — essential when the
- # git remote points to a proxy host (gh needs to know about
- # github.com independently of git remotes).
- result = subprocess.run(
- ["gh", "auth", "login", "--with-token", "--hostname", "github.com"],
- input=gh_token + "\n", text=True, capture_output=True,
- )
- if result.returncode == 0:
- print_info("gh auth login --with-token: success")
- else:
- print_warn(f"gh auth login stderr: {result.stderr.strip()}")
-
- result = subprocess.run(
- ["gh", "auth", "setup-git", "--hostname", "github.com"],
- capture_output=True, text=True,
- )
- if result.returncode == 0:
- print_info("gh auth setup-git: success")
- else:
- print_warn(f"gh auth setup-git skipped "
- f"(exit {result.returncode}): {result.stderr.strip()}")
- # Rewrite origin URL to embed token (for git push through proxy)
- try:
- origin_url = run_git(repo, "remote", "get-url", "origin").strip()
- if origin_url.startswith("https://"):
- clean_url = origin_url.replace("https://", "", 1)
- if "@" in clean_url:
- clean_url = clean_url.split("@", 1)[1]
- new_url = f"https://x-access-token:{gh_token}@{clean_url}"
- run_git(repo, "remote", "set-url", "origin", new_url)
- safe = f"https://x-access-token:***@{clean_url}"
- print_info(f"origin URL rewritten with token: {safe}")
- except Exception as exc:
- print_warn(f"Could not rewrite origin URL: {exc}")
- else:
- # Verify gh CLI is authenticated (interactive or env-based)
- try:
- subprocess.run(
- ["gh", "auth", "status"],
- check=True, capture_output=True, text=True,
- )
- subprocess.run(
- ["gh", "auth", "setup-git"],
- check=True, capture_output=True, text=True,
- )
- print_info("Git credential helper configured via gh")
- except subprocess.CalledProcessError as e:
- print_error(
- f"gh not authenticated and GH_TOKEN not set: {e.stderr.strip()}"
- )
- raise RuntimeError(
- "Cannot push to GitHub: no GH_TOKEN and gh CLI not authenticated. "
- "Run 'gh auth login' locally or set GH_TOKEN in CI."
- )
-
- def _push_submodule_if_needed(self) -> None:
- """Push AscendNPU-IR submodule changes to its remote.
-
- Raises RuntimeError on failure so the error is surfaced to
- GitHub Actions and the workflow exits with code 1.
-
- Uses the same branch name as the parent repo's work branch so the
- two repos stay in sync. Pushes with --force-with-lease to avoid
- clobbering existing remote state.
- """
- ascend_path = Path(self.state.triton_ascend_path)
- if not push_submodule(ascend_path, self.state.work_branch):
- raise RuntimeError(
- f"Failed to push AscendNPU-IR submodule branch "
- f"'{self.state.work_branch}'")
- self.state.summary_rows.append(
- ("Push AscendNPU-IR", "PASS", self.state.work_branch)
- )
-
- def _write_merge_metadata(self) -> None:
- """Write work branch, target commit, and step progress for CI orchestration."""
- metadata_dir = WORKSPACE_DIR / "merge-metadata"
- metadata_dir.mkdir(parents=True, exist_ok=True)
-
- (metadata_dir / "work_branch.txt").write_text(
- self.state.work_branch, encoding="utf-8"
- )
- (metadata_dir / "target_commit.txt").write_text(
- self.state.target_commit, encoding="utf-8"
- )
- (metadata_dir / "current_step.txt").write_text(
- str(self.state.current_step), encoding="utf-8"
- )
- (metadata_dir / "total_steps.txt").write_text(
- str(self.state.total_steps), encoding="utf-8"
- )
- is_last = (self.state.current_step >= self.state.total_steps - 1)
- (metadata_dir / "is_last_step.txt").write_text(
- str(is_last).lower(), encoding="utf-8"
- )
- print_info(f"Metadata written to {metadata_dir} "
- f"(step {self.state.current_step + 1}/{self.state.total_steps}, "
- f"is_last={is_last})")
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Mode: single-step — per-step merge → IR → build → test → fix
- # ═══════════════════════════════════════════════════════════════════════════
-
- def _run_single_step_mode(self, inputs: dict | None) -> str:
- """Single-step mode: each planned step runs the full pipeline independently.
-
- For each step:
- 1. Merge upstream commits + resolve conflicts
- 2. If LLVM hash changed: IR analysis → patch → rebuild LLVM
- 3. Build Triton-Ascend + AI fix compile errors
- 4. Run tests + AI fix test failures
- 5. Commit step progress
-
- After all steps: finalize + push + create PR.
-
- Controlled by TA_SINGLE_STEP_MODE=true env var.
- """
- # ── Apply inputs to state ──
- if inputs:
- for key, value in inputs.items():
- if hasattr(self.state, key):
- setattr(self.state, key, value)
-
- ascend_path = Path(self.state.triton_ascend_path)
-
- # ── Phase 0: Initialize ──
- self.initialize()
-
- # ── Phase 1: Detect commits & plan steps ──
- detect_result = self.detect_commits()
- if detect_result == HasNoNewCommits:
- print_info("No new commits — nothing to merge")
- self.state.summary_rows.append(
- ("Detect", "SKIP", "No new upstream commits"))
- self.state.final_status = UpgradeCompleted
- return UpgradeCompleted
-
- print_header("Single-Step Mode — Per-Step Full Pipeline")
- print_key_value("Total steps", str(self.state.total_steps))
- print_info("Each step: merge → [IR patch] → build → fix → test → fix → commit")
-
- # ── Phase 1.5: Build baseline LLVM (pre-merge, with Ascend patch) ──
- if not self._build_baseline_llvm():
- print_error("Baseline LLVM build failed — cannot proceed")
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- # ── Phase 2: Per-step loop ──
- while self.state.current_step < self.state.total_steps:
- step = self.state.steps[self.state.current_step]
+ while ctx.current_step < ctx.total_steps:
+ step = ctx.steps[ctx.current_step]
step_id = step["id"]
- self.state.retry_count = 0
+ ctx = ctx.copy_with(retry_count=0)
- print_header(
- f"Single-Step {self.state.current_step + 1}/{self.state.total_steps}: {step_id}"
+ log.header(
+ f"Step {ctx.current_step + 1}/{ctx.total_steps}: {step_id}"
)
- print_key_value("commits in step", str(step["commit_count"]))
- print_key_value("end commit", step["end_commit"][:12])
+ log.key_value("commits in step", str(step["commit_count"]))
+ log.key_value("end commit", step["end_commit"][:12])
reason = step.get("reason", "line_budget")
- print_key_value("step reason", reason)
-
- self._print_workspace_info(f"Single-Step Mode — {step_id}")
+ log.key_value("step reason", reason)
# Record ascend HEAD before this step
- self.state.step_start_ascend_head = run_git(
- ascend_path, "rev-parse", "HEAD"
- ).strip()
-
- # ── Step A: git merge ──
- merge_result = self._do_step_merge(step)
- if merge_result == UpgradeFailed:
- self._backup_code_state(f"failed-merge-{step_id}")
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- # ── Step B: AI resolve conflicts ──
- if self.state.merge_has_conflicts:
- if not self._do_resolve_conflicts():
- self._backup_code_state(f"failed-conflict-{step_id}")
- self.state.final_status = UpgradeFailed
+ ascend_path = Path(ctx.triton_ascend_path)
+ step_start_head = run_git(ascend_path, "rev-parse", "HEAD").strip()
+ ctx = ctx.copy_with(step_start_ascend_head=step_start_head)
+
+ # ── Step A: Merge ───────────────────────────────────
+ with timed("merge"):
+ ctx = merge_upstream_commit(ctx, self.config)
+ if ctx.merge_has_conflicts:
+ log.status(False, f"Merge has {len(ctx.conflict_files)} conflict(s)")
+ else:
+ log.status(True, "Merge clean")
+
+ # ── Step B: Resolve conflicts ───────────────────────
+ if ctx.merge_has_conflicts:
+ with timed("resolve"):
+ ctx = resolve_conflicts(ctx, self.config)
+ if ctx.merge_has_conflicts:
+ log.error(f"Conflicts unresolved for {step_id}")
+ ctx = ctx.copy_with(final_status=UpgradeFailed)
return UpgradeFailed
-
- # ── Step C: IR patch if LLVM hash changed in this step ──
- if reason == "llvm_version":
- print_section(f"LLVM Version Change in {step_id} — IR Patch Pipeline")
- if not self._do_per_step_ir_patch(step):
- self.state.final_status = UpgradeFailed
+ log.status(True, "Conflicts resolved")
+
+ # ── Step C: Build/Test — IR patch or standard ───────
+ llvm_hash_changed = (reason == "llvm_version")
+ if not llvm_hash_changed:
+ llvm_hash_changed = llvm_hash_changed_after_merge(ctx)
+ if llvm_hash_changed:
+ if reason != "llvm_version":
+ log.info(f"[{step_id}] LLVM hash changed during merge "
+ f"(post-merge detection) — routing to IR patch pipeline")
+ log.section(f"LLVM Version Change in {step_id} — IR Patch Pipeline")
+ with timed("ir-patch"):
+ ctx = per_step_ir_patch(ctx, self.config, step)
+ if not ctx.build_passed:
+ log.error(f"IR patch pipeline failed for {step_id}")
+ ctx = ctx.copy_with(final_status=UpgradeFailed)
+ return UpgradeFailed
+ else:
+ # Standard build + fix
+ log.section(f"Build & Fix — {step_id}")
+ with timed("build"):
+ ctx = build_and_fix_loop(ctx, self.config)
+ if not ctx.build_passed:
+ log.error(f"Build failed for {step_id}")
+ ctx = ctx.copy_with(final_status=UpgradeFailed)
return UpgradeFailed
- # ── Step D: Build + AI fix compile errors ──
- print_section(f"Build & Fix — {step_id}")
- if not self._do_build_and_fix_loop():
- self._backup_code_state(f"failed-build-{step_id}")
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- # ── Step E: Test + AI fix test failures ──
- if not self._do_test_and_fix_loop():
- self._backup_code_state(f"failed-test-{step_id}")
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
+ # Standard test + fix
+ log.section(f"Test & Fix — {step_id}")
+ with timed("test"):
+ ctx = test_and_fix_loop(ctx, self.config)
+ if not ctx.test_passed:
+ log.error(f"Tests failed for {step_id}")
+ ctx = ctx.copy_with(final_status=UpgradeFailed)
+ return UpgradeFailed
- # ── Step F: Commit step progress ──
- self._do_commit_step(step)
+ # ── Step D: Commit ──────────────────────────────────
+ with timed("commit"):
+ ctx = commit_step(ctx, self.config)
# Record step description for PR body
desc = (
@@ -718,3171 +173,61 @@ def _run_single_step_mode(self, inputs: dict | None) -> str:
f"source lines={step.get('source_changed_lines', '?')}, "
f"reason={reason}"
)
- self.state.step_pr_descriptions.append(desc)
+ ctx.step_pr_descriptions.append(desc)
# Record per-step detail for sync report
- self.state.step_details.append({
+ ctx.step_details.append({
"step_id": step_id,
- "step_index": self.state.current_step + 1,
+ "step_index": ctx.current_step + 1,
"commits": step["commit_count"],
"end_commit": step["end_commit"][:12],
"source_lines": step.get("source_changed_lines", 0),
- "conflict_files": len(self.state.conflict_files),
- "build_fixes": self.state.build_fix_count,
- "test_fixes": self.state.test_fix_count,
- "retries": self.state.retry_count,
+ "conflict_files": len(ctx.conflict_files),
+ "build_fixes": ctx.build_fix_count,
+ "test_fixes": ctx.test_fix_count,
+ "retries": ctx.retry_count,
"reason": reason,
})
- self.state.current_step += 1
- print_status(True, f"Step {step_id} completed "
- f"({self.state.current_step}/{self.state.total_steps})")
+ # Advance to next step
+ ctx = ctx.copy_with(current_step=ctx.current_step + 1)
+ log.status(
+ True,
+ f"Step {step_id} completed "
+ f"({ctx.current_step}/{ctx.total_steps})",
+ )
- # ── Phase 3: Finalize ──
- print_header("Finalize — Generate Summary & Push")
- self._do_finalize()
+ # ── Phase 4: Finalize ───────────────────────────────────────
+ log.header("Phase 4: Finalize")
+ with timed("finalize"):
+ ctx = finalize(ctx, self.config)
- # ── Phase 4: Push to GitHub + create PR ──
- self.push_to_github()
+ # ── Phase 5: Push PR ────────────────────────────────────────
+ if self.config.push_to_github:
+ self._push_pr(ctx)
- self.state.summary_rows.append(
+ ctx.summary_rows.append(
("Single-Step Sync", "PASS",
- f"{self.state.total_steps} step(s), branch: {self.state.work_branch}")
+ f"{ctx.total_steps} step(s), branch: {ctx.work_branch}")
)
- print_summary_table(self.state.summary_rows)
- print_elapsed_total()
+ log.table(ctx.summary_rows)
+ log.elapsed(total_elapsed())
- self.state.final_status = UpgradeCompleted
+ ctx = ctx.copy_with(final_status=UpgradeCompleted)
return UpgradeCompleted
- # ═══════════════════════════════════════════════════════════════════════════
- # Mode: fix — AI fix on existing work branch
- # ═══════════════════════════════════════════════════════════════════════════
-
- def _run_fix_mode(self, inputs: dict | None) -> str:
- """AI fix on an existing work branch.
-
- Reads error logs from TA_ERROR_LOGS_PATH, calls the AI fix engine
- (_do_ai_fix), commits & pushes fixes. Used in CI after NPU tests fail.
- """
- ascend_path_str = (
- (inputs or {}).get("triton_ascend_path")
- or os.getenv("TRITON_ASCEND_PATH")
- or str(Path.cwd())
- )
- work_branch = os.getenv("TA_WORK_BRANCH", "")
- error_logs_path = os.getenv("TA_ERROR_LOGS_PATH", "")
- attempt = int(os.getenv("TA_FIX_ATTEMPT", "1"))
- target_commit = (
- (inputs or {}).get("target_commit")
- or os.getenv("TRITON_TARGET_COMMIT", "")
- )
-
- if not work_branch:
- print_error("TA_WORK_BRANCH is required for fix mode")
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- ascend_path = Path(ascend_path_str)
-
- # ── Setup ──
- print_header(f"Fix Mode — Attempt {attempt}")
- print_key_value("work branch", work_branch)
- print_key_value("error logs", error_logs_path or "")
- print_key_value("target commit", target_commit[:12] if target_commit else "")
- print_key_value("repo path", str(ascend_path))
-
- self._print_workspace_info(f"Fix Mode — Attempt {attempt}")
-
- # Clean old workspace
- if WORKSPACE_DIR.exists():
- shutil.rmtree(WORKSPACE_DIR)
- WORKSPACE_DIR.mkdir(parents=True)
-
- # Populate minimal state
- self.state.triton_ascend_path = str(ascend_path)
- self.state.triton_path = os.getenv("TRITON_PATH", str(ascend_path))
- self.state.target_commit = target_commit
- self.state.work_branch = work_branch
- self.state.original_branch = work_branch
- self.state.current_step = 0
- self.state.total_steps = 1
- self.state.steps = [{
- "index": 1,
- "id": "fix-step-1",
- "commit_count": 0,
- "end_commit": target_commit or "",
- "source_changed_lines": 0,
- }]
-
- # ── Checkout work branch ──
- print_section("Checkout Work Branch")
- try:
- run_git(ascend_path, "fetch", "origin", work_branch)
- except Exception as e:
- print_warn(f"Could not fetch {work_branch}: {e}")
- run_git(ascend_path, "checkout", work_branch)
- print_status(True, f"Checked out {work_branch}")
+ def _push_pr(self, ctx: WorkflowContext) -> None:
+ """Push work branch and create GitHub PR."""
+ from TA_main2main_workflow.pipeline.push_pr import push_and_create_pr
- # Pull latest (in case previous fix attempts pushed)
try:
- run_git(ascend_path, "pull", "origin", work_branch)
- print_info("Pulled latest changes")
- except Exception:
- print_warn("Could not pull latest — continuing with local state")
-
- # ── Collect error logs ──
- fix_errors: list[str] = []
- if error_logs_path:
- error_path = Path(error_logs_path)
- if error_path.exists():
- if error_path.is_dir():
- fix_errors = sorted(
- str(p) for p in error_path.rglob("*") if p.is_file()
- )
- print_info(f"Found {len(fix_errors)} error log file(s)")
- else:
- fix_errors = [str(error_path)]
- print_info(f"Using error log: {error_path}")
-
- if not fix_errors:
- print_warn("No error logs found — AI will analyze the codebase directly")
- # Create a stub so _do_ai_fix has something to work with
- stub_log = WORKSPACE_DIR / "no-error-logs.txt"
- stub_log.write_text(
- "No specific error logs were provided from the NPU CI run.\n"
- "Please analyze the triton-ascend codebase for potential issues\n"
- f"that could cause build or test failures after merging upstream triton.\n"
- f"Target upstream commit: {target_commit}\n"
- f"Work branch: {work_branch}\n"
+ pr_url = push_and_create_pr(
+ ascend_path=ctx.ascend_path,
+ github_repo=self.config.github_repo,
+ summary_path=WORKSPACE_DIR / "final_summary.md",
+ target_commit=ctx.target_commit,
+ work_branch=ctx.work_branch,
)
- fix_errors = [str(stub_log)]
-
- self.state.fix_errors = fix_errors
-
- # ── Set up step directory ──
- step_dir = WORKSPACE_DIR / "fix-step-1"
- step_dir.mkdir(parents=True, exist_ok=True)
-
- # ── Run AI fix ──
- print_header("AI Fix Analysis")
- print_info(f"Error sources ({len(fix_errors)}):")
- for e in fix_errors[:10]:
- print(f" • {e}")
- if len(fix_errors) > 10:
- print(f" ... and {len(fix_errors) - 10} more")
-
- try:
- fix_ok = self._do_ai_fix(ascend_path, step_dir, attempt)
+ log.status(True, f"PR created: {pr_url}")
except Exception as e:
- print_error(f"AI fix crashed: {e}")
- import traceback
- traceback.print_exc()
- fix_ok = False
-
- if not fix_ok:
- print_error("AI fix did not produce any changes")
- self.state.summary_rows.append(
- ("AI fix", "FAIL", "No changes produced")
- )
- print_summary_table(self.state.summary_rows)
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- # ── Commit and push ──
- print_section("Commit & Push Fixes")
-
- # ── Commit submodule changes first ──
- self.state.retry_count = attempt - 1
- self._commit_submodule_if_needed()
-
- # ── Clean temp artifacts BEFORE staging ──
- cleanup_temp_files(ascend_path)
-
- status = run_git(ascend_path, "status", "--porcelain").strip()
- if status:
- run_git(ascend_path, "add", "-A")
- staged = run_git(ascend_path, "diff", "--cached", "--name-only").strip()
- if staged:
- print_info(f"Files staged ({len(staged.splitlines())}):")
- for f in staged.splitlines()[:10]:
- print_info(f" - {f}")
- commit_target = target_commit[:12] if target_commit else "upstream"
- commit_msg = (
- f"[Sync](fix) AI-generated build/test failures fix "
- f"for merging {commit_target}\n\n"
- f"Upstream target: {commit_target}\n"
- f"Fix attempt: {attempt}\n"
- f"Work branch: {work_branch}\n"
- )
- run_git(ascend_path, "commit", "-s", "-m", commit_msg)
- print_status(True, "Committed AI fix")
-
- # ── Push AscendNPU-IR submodule first ──
- self._push_submodule_if_needed()
-
- run_git(ascend_path, "push", "origin", work_branch)
- print_status(True, f"Pushed to origin/{work_branch}")
- self.state.summary_rows.append(
- ("AI fix", "PASS", f"Attempt {attempt}")
- )
- else:
- print_info("No changes to commit after AI fix")
- self.state.summary_rows.append(
- ("AI fix", "NOOP", "No changes needed")
- )
-
- print_header("Fix Phase Complete!")
- print_key_value("work branch", work_branch)
- print_key_value("attempt", str(attempt))
- print_info(f"Next: re-trigger NPU tests on branch '{work_branch}'")
-
- print_summary_table(self.state.summary_rows)
- self.state.final_status = UpgradeCompleted
- return UpgradeCompleted
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Phase 0: Initialize
- # ═══════════════════════════════════════════════════════════════════════════
-
- @start()
- def initialize(self):
- start_timer("flow-total")
-
- print_header("Triton-Ascend Upstream Sync — Main2Main Flow")
- print(f" Started: {time.strftime('%Y-%m-%d %H:%M:%S')}", flush=True)
- print(f" AI Backend: {os.getenv('AI_BACKEND', 'auto-detect')}", flush=True)
- print(f" Max Retries: {self.state.max_retries}", flush=True)
-
- if WORKSPACE_DIR.exists():
- shutil.rmtree(WORKSPACE_DIR)
- WORKSPACE_DIR.mkdir(parents=True)
-
- raw_ascend = (
- self.state.triton_ascend_path
- or os.getenv("TRITON_ASCEND_PATH")
- or str(Path.cwd())
- )
- raw_triton = (
- self.state.triton_path
- or os.getenv("TRITON_PATH")
- or str(Path.cwd())
- )
-
- self.state.triton_ascend_path = raw_ascend
- self.state.triton_path = os.path.expanduser(raw_triton)
- self.state.target_commit = (
- self.state.target_commit or os.getenv("TRITON_TARGET_COMMIT", "")
- )
- self.state.llvm_prefix = os.getenv("LLVM_INSTALL_PREFIX", "")
- self.state.conda_env = os.getenv("CONDA_ENV", "ta-upgrade")
- self.state.num_procs = int(os.getenv("NUM_PROCS", "16"))
-
- if not self.state.test_log_dir:
- self.state.test_log_dir = str(WORKSPACE_DIR / "test-logs")
-
- ascend_path = Path(self.state.triton_ascend_path)
-
- # ── safety: abort any stale merge ──
- merge_head = ascend_path / ".git" / "MERGE_HEAD"
- if merge_head.exists():
- print_warn("Found stale MERGE_HEAD from previous run, aborting it")
- try:
- run_git(ascend_path, "merge", "--abort")
- print_info("Stale merge aborted successfully")
- except Exception:
- print_warn("merge --abort failed, trying reset --hard")
- try:
- run_git(ascend_path, "reset", "--hard", "HEAD")
- except Exception:
- pass
- for stale in [".git/MERGE_MODE", ".git/MERGE_MSG", ".git/CHERRY_PICK_HEAD"]:
- p = ascend_path / stale
- if p.exists():
- p.unlink()
-
- ascend_branch = run_git(ascend_path, "branch", "--show-current").strip()
- self.state.original_branch = ascend_branch or run_git(
- ascend_path, "rev-parse", "HEAD"
- ).strip()
-
- # ── Use the configured base branch for patch diffs ──
- # The work branch is created from the base branch, so all diffs should
- # be computed against it, not the checkout HEAD.
- base_branch = os.getenv(ENV_BASE_BRANCH, "main")
- base_ref = get_base_branch_ref()
- try:
- run_git(ascend_path, "fetch", "origin", base_branch)
- except Exception:
- print_warn(f"Could not fetch {base_ref}, using checkout HEAD as base")
- try:
- self.state.ascend_head = run_git(
- ascend_path, "rev-parse", base_ref).strip()
- except Exception:
- self.state.ascend_head = run_git(ascend_path, "rev-parse", "HEAD").strip()
-
- print_section("Repository Configuration")
- print_key_value("triton-ascend", self.state.triton_ascend_path)
- print_key_value("upstream triton", self.state.triton_path)
- print_key_value("target commit", self.state.target_commit or "")
- print_key_value("original branch", self.state.original_branch)
- print_key_value(f"base ({base_ref})", self.state.ascend_head[:12])
-
- self._print_workspace_info("Phase 0: Initialize")
-
- self.state.summary_rows = []
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Phase 1: Detect upstream commits
- # ═══════════════════════════════════════════════════════════════════════════
-
- @router(initialize)
- def detect_commits(self) -> Literal["HasNewCommits", "HasNoNewCommits"]:
- start_timer("detect")
- print_header("Phase 1: Detect Upstream Commits & Plan Steps")
-
- self._print_workspace_info("Phase 1: Detect Commits")
-
- ascend_path = Path(self.state.triton_ascend_path)
- triton_path = Path(self.state.triton_path)
-
- result, has_new = detect(
- ascend_path,
- triton_path,
- self.state.target_commit or None,
- )
-
- self.state.merge_base = result["merge_base"]
- self.state.target_commit = result["target_commit"]
- self.state.upstream_commits_count = result["upstream_commits_count"]
-
- print_key_value("merge_base", self.state.merge_base[:12])
- print_key_value("target", self.state.target_commit[:12])
- print_key_value("upstream commits", str(self.state.upstream_commits_count))
- print_key_value("changed files", str(result["changed_files_count"]))
- print_key_value("changed lines", str(result["changed_lines"]["total"]))
-
- commits = result.get("upstream_commits", [])
- if commits:
- print_info(f"Commits to merge ({len(commits)}):")
- for c in commits[:20]:
- print(f" {c['sha'][:8]} {c['subject'][:80]}")
- if len(commits) > 20:
- print(f" ... and {len(commits) - 20} more")
-
- if not has_new:
- print_status(True, "Already up to date — nothing to merge")
- self.state.summary_rows.append(("Detect commits", "PASS", "No new commits"))
- stop_timer("detect")
- return HasNoNewCommits
-
- # ── Check if progressive merge is enabled ──
- progressive_env = os.getenv("TA_PROGRESSIVE_MERGE", "true").lower()
- self.state.progressive_merge = progressive_env != "false"
-
- # ── Plan steps: split commits into chunks based on line budget ──
- if self.state.progressive_merge and self.state.upstream_commits_count > 1:
- print_section("Step Planning")
- line_budget = int(os.getenv("TA_LINE_BUDGET", str(LINE_BUDGET)))
- print_key_value("line budget", str(line_budget))
-
- plan = run_plan(
- triton_path,
- self.state.merge_base,
- self.state.target_commit,
- line_budget=line_budget,
- )
- self.state.steps = plan["steps"]
- self.state.total_steps = len(plan["steps"])
-
- # ── Guard: if planner produced 0 steps (e.g., all commits filtered
- # out), fall back to single-step mode so something still gets merged ──
- if self.state.total_steps == 0:
- print_warn("Plan returned 0 steps — falling back to single-step merge")
- self.state.total_steps = 1
- self.state.steps = [{
- "index": 1,
- "id": "step-1",
- "commit_count": self.state.upstream_commits_count,
- "start_commit": self.state.merge_base,
- "end_commit": self.state.target_commit,
- "source_changed_lines": result["changed_lines"]["total"],
- }]
-
- print_status(True, f"Planned {self.state.total_steps} step(s) "
- f"from {plan['total_source_commits']} source-touching commits "
- f"({plan['total_commits']} total upstream commits)")
- else:
- # Single-step mode: treat everything as one step
- self.state.total_steps = 1
- self.state.steps = [{
- "index": 1,
- "id": "step-1",
- "commit_count": self.state.upstream_commits_count,
- "start_commit": self.state.merge_base,
- "end_commit": self.state.target_commit,
- "source_changed_lines": result["changed_lines"]["total"],
- }]
- if not self.state.progressive_merge:
- print_info("TA_PROGRESSIVE_MERGE=false — using single-step mode")
- else:
- print_info("Only 1 upstream commit — using single-step mode")
-
- stop_timer("detect")
- print_status(True, f"Found {self.state.upstream_commits_count} upstream commits to merge "
- f"across {self.state.total_steps} step(s)")
- self.state.summary_rows.append(
- ("Detect commits", "PASS",
- f"{self.state.upstream_commits_count} commits, {self.state.total_steps} step(s)")
- )
- return HasNewCommits
-
- @listen(HasNoNewCommits)
- def has_no_commits(self):
- print_header("Sync Complete — Already Up To Date")
- print_elapsed_total()
- print_summary_table(self.state.summary_rows)
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Phase 2: Execute Sync (orchestrates merge → resolve → build → test → fix)
- # ═══════════════════════════════════════════════════════════════════════════
- #
- # This is the core loop. It runs as a SINGLE @router node to avoid
- # CrewAI @listen → @listen signal chaining issues. All sub-steps are
- # internal method calls, not CrewAI routing targets.
-
- @router(detect_commits)
- def execute_sync(self) -> Literal["UpgradeCompleted", "UpgradeFailed"]:
- """Orchestrate the full sync pipeline — progressively or single-step.
-
- When progressive_merge is True (default), each planned step is merged
- and validated independently before moving to the next. This keeps
- AI conflict resolution and fix scopes small and manageable.
-
- The internal per-step call chain is:
- _do_step_merge → _do_resolve_conflicts → _do_build_and_fix_loop → _do_commit_step → _push_step_progress
- """
- try:
- return self._execute_sync_inner()
- except Exception as exc:
- print_error(f"Unexpected error in execute_sync: {exc}")
- import traceback
- traceback.print_exc()
- # Backup code before failing so partial work is preserved
- self._backup_code_state(f"crash-step{self.state.current_step + 1}")
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- def _execute_sync_inner(self) -> Literal["UpgradeCompleted", "UpgradeFailed"]:
- """Inner body of execute_sync — wrapped by try/except for crash backup."""
-
- # ── Iterate over each planned step ──
- while self.state.current_step < self.state.total_steps:
- step = self.state.steps[self.state.current_step]
- step_id = step["id"]
- self.state.retry_count = 0
-
- print_header(f"Step {self.state.current_step + 1}/{self.state.total_steps}: {step_id}")
- print_key_value("commits in step", str(step["commit_count"]))
- print_key_value("end commit", step["end_commit"][:12])
- if "source_changed_lines" in step:
- print_key_value("source lines", str(step["source_changed_lines"]))
-
- self._print_workspace_info(f"Phase 2: Execute Sync — {step_id}")
-
- ascend_path = Path(self.state.triton_ascend_path)
-
- # ── Work-branch guard: verify we're on the right branch ──
- if self.state.current_step > 0 and self.state.work_branch:
- current_branch = run_git(ascend_path, "branch", "--show-current").strip()
- if current_branch != self.state.work_branch:
- print_warn(f"Expected work branch '{self.state.work_branch}' "
- f"but currently on '{current_branch}' — switching back")
- run_git(ascend_path, "checkout", self.state.work_branch)
- print_info(f"Same work branch: '{self.state.work_branch}' "
- f"(step {self.state.current_step + 1}/{self.state.total_steps})")
-
- # Record ascend HEAD before this step (for per-step patch generation)
- self.state.step_start_ascend_head = run_git(
- ascend_path, "rev-parse", "HEAD"
- ).strip()
-
- # ── Step A: git merge this step's end commit ──
- result = self._do_step_merge(step)
- if result == UpgradeFailed:
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- # ── Step B: AI resolve conflict (if merge had conflicts) ──
- if self.state.merge_has_conflicts:
- if not self._do_resolve_conflicts():
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- # ── Step C: build → test → AI fix bug loop ──
- try:
- build_ok = self._do_build_and_fix_loop()
- except Exception as exc:
- print_error(f"_do_build_and_fix_loop crashed: {exc}")
- import traceback
- traceback.print_exc()
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- if not build_ok:
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- # ── Step D: commit step progress ──
- self._do_commit_step(step)
-
- # ── Record step description for final PR body ──
- desc = (
- f"✅ **{step_id}**: {step['commit_count']} commits, "
- f"end_commit=`{step['end_commit'][:12]}`, "
- f"source lines={step.get('source_changed_lines', '?')}"
- )
- self.state.step_pr_descriptions.append(desc)
-
- # ── Record per-step detail for sync report ──
- conflicts_in_step = len(self.state.conflict_files)
- self.state.step_details.append({
- "step_id": step_id,
- "step_index": self.state.current_step + 1,
- "commits": step["commit_count"],
- "end_commit": step["end_commit"][:12],
- "source_lines": step.get("source_changed_lines", 0),
- "conflict_files": conflicts_in_step,
- "build_fixes": self.state.build_fix_count,
- "test_fixes": self.state.test_fix_count,
- "retries": self.state.retry_count,
- })
-
- # Move to next step
- self.state.current_step += 1
- print_status(True, f"Step {step_id} completed successfully "
- f"({self.state.current_step}/{self.state.total_steps})")
-
- # ── Phase 3+4: IR compatibility patches + pytest ut test ──
- ir_ok = self._do_ir_patch_loop()
- if not ir_ok:
- print_error("IR patch loop did not converge — sync failed")
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- # ── Finalize: generate cumulative patch & summary ──
- self._do_finalize()
- self.state.final_status = UpgradeCompleted
- return UpgradeCompleted
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Internal step implementations
- # ═══════════════════════════════════════════════════════════════════════════
-
- def _do_step_merge(self, step: dict) -> Literal["HasNewCommits"] | Literal["UpgradeFailed"]:
- """Merge this step's end_commit into triton-ascend.
-
- The first step creates a fresh work branch from upstream-ascend/main
- and merges its end_commit. Subsequent steps merge their end_commit
- on top of the SAME work branch — git handles the incremental merge
- automatically by computing the diff between the previous end_commit
- and the new one.
-
- ALL steps share ONE work branch. This is critical: we accumulate
- changes on a single branch so the final PR contains the full history.
- """
- start_timer("merge")
- step_id = step["id"]
- is_first_step = self.state.current_step == 0
-
- ascend_path = Path(self.state.triton_ascend_path)
- triton_path = Path(self.state.triton_path)
-
- # ── Verify / log work branch consistency ──
- if is_first_step:
- print_info(f"No work branch yet — will create one for step {step_id}")
- else:
- current_branch = run_git(ascend_path, "branch", "--show-current").strip()
- if current_branch != self.state.work_branch:
- print_warn(f"Expected work branch '{self.state.work_branch}' "
- f"but currently on '{current_branch}' — switching back")
- run_git(ascend_path, "checkout", self.state.work_branch)
- print_info(f"Continuing on work branch: '{self.state.work_branch}' "
- f"(verified same branch as step 1)")
-
- print_flow_progress("merge", f"[{step_id}] merging {step['end_commit'][:12]}")
-
- try:
- if is_first_step:
- # First step: create work branch and do full merge
- merge_result = run_merge(
- ascend_path,
- triton_path,
- step["end_commit"],
- )
- self.state.work_branch = merge_result["work_branch"]
- print_info(f"Created work branch: '{self.state.work_branch}' "
- f"(all {self.state.total_steps} step(s) will use this branch)")
- else:
- # Subsequent step: merge on top of existing work branch
- # fetch the new target if it's not already present
- try:
- run_git(ascend_path, "fetch", "upstream-triton", "--prune")
- except Exception:
- print_info("Could not fetch upstream-triton, assuming target is reachable")
-
- merge_result = run_merge_incremental(
- ascend_path,
- triton_path,
- step["end_commit"],
- self.state.work_branch,
- )
- except Exception as exc:
- print_error(f"Merge failed with exception: {exc}")
- stop_timer("merge")
- self.state.summary_rows.append(
- (f"Merge step {step_id}", "FAIL", str(exc)[:60])
- )
- return UpgradeFailed
-
- self.state.merge_has_conflicts = merge_result["has_conflicts"]
- self.state.conflict_files = merge_result.get("conflict_files", [])
-
- print_key_value("work branch", self.state.work_branch)
- print_key_value("has conflicts", str(self.state.merge_has_conflicts))
- print_key_value("exit code", str(merge_result["merge_exit_code"]))
- print_key_value("step", f"{self.state.current_step + 1}/{self.state.total_steps}")
-
- # If merge had non-zero exit but no conflict markers, that's a hard failure
- if merge_result["merge_exit_code"] != 0 and not self.state.merge_has_conflicts:
- print_error(f"Merge exited with code {merge_result['merge_exit_code']} "
- f"but no conflict markers found — this is an unexpected failure")
- stop_timer("merge")
- self.state.summary_rows.append(
- (f"Merge step {step_id}", "FAIL",
- f"exit code {merge_result['merge_exit_code']}")
- )
- return UpgradeFailed
-
- if self.state.merge_has_conflicts:
- print_conflict_list(self.state.conflict_files)
- stop_timer("merge")
- self.state.summary_rows.append(
- (f"Merge step {step_id}", "WARN", f"{len(self.state.conflict_files)} conflicts")
- )
- else:
- stop_timer("merge")
- print_status(True, f"Step {step_id} merge succeeded with no conflicts")
- self.state.summary_rows.append(
- (f"Merge step {step_id}", "PASS", f"{step['commit_count']} commits")
- )
-
- return HasNewCommits
-
- def _do_resolve_conflicts(self) -> bool:
- """AI-driven merge conflict resolution with retry loop.
-
- For each attempt (up to max_retries):
- 1. Refresh the conflict file list from git
- 2. Call opencode/claude with the conflict snapshots
- 3. Check if all conflicts are resolved
- 4. If not, retry with refreshed conflict list
-
- AI context includes: step index (N/total), is_last_step flag,
- previous_step_id and previous_step_summary_path for continuity
- (matching vllm-ascend's main2main_flow pattern).
-
- After all conflicts are resolved:
- - git commit the resolution
- - Run pre-CI checks (conflict markers, temp files, syntax)
- - Write step summary and cumulative patch
-
- Returns True if all conflicts resolved, False otherwise.
- """
- start_timer("resolve")
- print_header("Phase 3: AI Conflict Resolution")
-
- ascend_path = Path(self.state.triton_ascend_path)
-
- step = self.state.steps[self.state.current_step] if self.state.steps else None
- current_step_id = step["id"] if step else "step-0"
- is_last_step = self.state.current_step == self.state.total_steps - 1
-
- # Use step-specific directory in progressive mode, fall back to step-0
- if self.state.total_steps > 1 and self.state.steps:
- step_dir = WORKSPACE_DIR / STEPS_DIR / current_step_id
- else:
- step_dir = WORKSPACE_DIR / "step-0"
- step_dir.mkdir(parents=True, exist_ok=True)
-
- # ── Previous step context (matching vllm-ascend pattern) ──
- previous_step = (
- self.state.steps[self.state.current_step - 1]
- if self.state.current_step > 0 and self.state.steps else None
- )
- previous_step_id = previous_step["id"] if previous_step else ""
- previous_step_summary_path = (
- str(WORKSPACE_DIR / STEPS_DIR / previous_step_id / EACH_STEP_SUMMARY_FILE)
- if previous_step_id else ""
- )
-
- conflict_dir = WORKSPACE_DIR / CONFLICT_LOG_DIR
-
- # AI resolve conflict: check if AI is disabled
- if os.getenv("SKIP_AI_ANALYSIS", "false").lower() == "true":
- print_warn("SKIP_AI_ANALYSIS=true — skipping AI conflict resolution!")
- print_warn("Conflicts will NOT be resolved automatically.")
- print_conflict_list(self.state.conflict_files)
- print_info("To resolve: manually edit conflicted files, then run:")
- print_info(f" cd {ascend_path} && git add -u && git commit --no-edit")
- self.state.summary_rows.append(("AI resolve conflicts", "SKIP", "SKIP_AI_ANALYSIS set"))
- return False
-
- # AI resolve conflict: detect backend (opencode / claude)
- try:
- from TA_main2main_workflow.agent.opencode_adapter import _detect_backend
- backend = _detect_backend()
- print_info(f"AI backend detected: {backend}")
- except RuntimeError as e:
- print_error(f"AI backend not available: {e}")
- print_info("Install 'opencode' or 'claude' CLI, or set AI_BACKEND env var.")
- self.state.summary_rows.append(("AI resolve conflicts", "FAIL", str(e)[:50]))
- return False
-
- resolved_all = False
- ai_result: AIResult | None = None
- conflict_files = list(self.state.conflict_files)
- original_conflict_count = len(conflict_files)
-
- # AI resolve conflict: retry loop (up to max_retries)
- for attempt in range(1, self.state.max_retries + 1):
- print_step(attempt, self.state.max_retries, "AI conflict resolution")
-
- conflict_files = get_conflict_files(ascend_path)
- if not conflict_files:
- print_status(True, "No conflicts detected — already resolved!")
- resolved_all = True
- break
-
- print_info(f"Files with conflicts: {len(conflict_files)}")
- for f in conflict_files:
- print(f" • {f}")
-
- print_ai_call_info(
- backend=backend,
- mode="conflict",
- attempt=attempt,
- max_attempts=self.state.max_retries,
- )
-
- # AI resolve conflict: invoke opencode/claude
- # Context matches vllm-ascend pattern: is_last_step,
- # previous_step_id, previous_step_summary_path, step index
- try:
- ai_result = run_opencode_adapter({
- "step_id": f"{current_step_id}-conflict-{attempt}",
- "previous_step_id": previous_step_id,
- "previous_step_summary_path": previous_step_summary_path,
- "is_last_step": str(is_last_step).lower(),
- "step_index": f"{self.state.current_step + 1}/{self.state.total_steps}",
- "step_dir": str(step_dir),
- "conflict_dir": str(conflict_dir),
- "ascend_path": str(ascend_path),
- "triton_path": self.state.triton_path,
- "reference_dir": _REFERENCE_DIR,
- "mode": "conflict",
- "error_logs": json.dumps(conflict_files, ensure_ascii=False),
- "target_commit": self.state.target_commit,
- })
- except Exception as e:
- print_error(f"AI call failed: {e}")
- if attempt < self.state.max_retries:
- print_info(f"Retrying... ({attempt}/{self.state.max_retries})")
- continue
- break
-
- if not has_merge_conflicts(ascend_path):
- print_status(True, f"All conflicts resolved! (attempt {attempt})")
- self.state.conflict_files_resolved += original_conflict_count
- resolved_all = True
- break
- else:
- still_conflicted = len(get_conflict_files(ascend_path))
- print_status(False, f"{still_conflicted} conflict(s) remain after attempt {attempt}")
- conflict_files = get_conflict_files(ascend_path)
-
- if not resolved_all:
- remaining = get_conflict_files(ascend_path)
- print_error(f"Failed to resolve all conflicts after {self.state.max_retries} attempts")
- print_conflict_list(remaining)
- stop_timer("resolve")
- self.state.summary_rows.append(("AI resolve conflicts", "FAIL", "Conflicts remain"))
- return False
-
- # AI resolve conflict: git commit the resolution
- # Clean temp artifacts first, then use git add -A to ensure
- # AI-created files are NOT dropped.
- cleanup_temp_files(ascend_path)
- try:
- run_git(ascend_path, "add", "-A")
- staged = run_git(ascend_path, "diff", "--cached", "--name-only").strip()
- if staged:
- print_info(f"Files staged ({len(staged.splitlines())}):")
- for f in staged.splitlines()[:10]:
- print_info(f" - {f}")
- run_git(ascend_path, "commit", "--no-edit", "-s")
- print_status(True, "Committed conflict resolution")
- except subprocess.CalledProcessError as e:
- stderr = (e.stderr or "").strip() if hasattr(e, 'stderr') else str(e)
- if "nothing to commit" in stderr.lower():
- print_info("Nothing to commit — resolution may already be committed")
- else:
- print_warn(f"Commit may have failed: {stderr[-200:]}")
-
- # pre-CI check: scan for leftover conflict markers, temp files, syntax errors
- print_info("Running pre-CI check after conflict resolution...")
- pre_ci_result = run_pre_ci_check(ascend_path, step_id="conflict-resolution")
- if not pre_ci_result["all_passed"]:
- print_warn("Pre-CI check found issues — review before proceeding")
- self.state.summary_rows.append(
- ("Pre-CI check", "PASS" if pre_ci_result["all_passed"] else "WARN",
- f"{pre_ci_result.get('modified_files_count', 0)} files checked")
- )
-
- # ── Write step summary ──
- summary_path = step_dir / EACH_STEP_SUMMARY_FILE
- if ai_result and ai_result.step_summary and not summary_path.exists():
- summary_path.write_text(ai_result.step_summary, encoding="utf-8")
-
- # ── Generate step patch ──
- try:
- patch = run_git(ascend_path, "diff", self.state.ascend_head, "HEAD")
- (step_dir / EACH_STEP_TARGET_PATCH_FILE).write_text(patch, encoding="utf-8")
- except Exception:
- pass
-
- stop_timer("resolve")
- elapsed = ai_result.elapsed_seconds if ai_result else 0
- print_status(True, f"Conflict resolution complete ({elapsed:.0f}s AI time)")
- self.state.summary_rows.append(
- ("AI resolve conflicts", "PASS", f"{elapsed:.0f}s" if elapsed else "done")
- )
- self.state.merge_has_conflicts = False
- return True
-
- def _do_build_and_fix_loop(self) -> bool:
- """build → AI fix compile-error loop (up to max_retries rounds).
-
- Only handles compilation errors. Tests are deferred to after all
- upstream commits are merged and the final build passes.
- """
- ascend_path = Path(self.state.triton_ascend_path)
- step = self.state.steps[self.state.current_step] if self.state.steps else None
- current_step_id = step["id"] if step else "step-0"
-
- # Use step-specific directory in progressive mode, fall back to step-0
- if self.state.total_steps > 1 and self.state.steps:
- step_dir = WORKSPACE_DIR / STEPS_DIR / current_step_id
- else:
- step_dir = WORKSPACE_DIR / "step-0"
- step_dir.mkdir(parents=True, exist_ok=True)
-
- build_passed = False
-
- for attempt in range(self.state.max_retries + 1):
- is_fix_attempt = attempt > 0
- self.state.retry_count = attempt
-
- # AI fix compile errors (skip on first round)
- if is_fix_attempt:
- print_header(f"Fix Attempt {attempt}/{self.state.max_retries} (build)")
- ai_ok = self._do_ai_fix(ascend_path, step_dir, attempt)
- # Collect fix detail
- modified_files: list[str] = []
- ai_summary = ""
- if hasattr(self, '_last_ai_result') and self._last_ai_result:
- modified_files = self._last_ai_result.get("modified_files", [])
- ai_summary = self._last_ai_result.get("step_summary", "")
- # Read error log snippet for context
- error_snippet = ""
- for err_path in self.state.fix_errors:
- try:
- content = Path(err_path).read_text(encoding="utf-8", errors="replace")
- error_snippet += content[-2000:] if len(content) > 2000 else content
- except Exception:
- pass
- self.state.fix_attempts.append({
- "step_id": current_step_id,
- "attempt": attempt,
- "fix_type": "build",
- "error_logs": list(self.state.fix_errors),
- "error_snippet": error_snippet[-1500:],
- "modified_files": modified_files,
- "ai_summary": (ai_summary or "")[:2000],
- "ai_ok": ai_ok,
- })
- if not ai_ok:
- pass
-
- # build triton-ascend
- if not self._do_build(ascend_path, clean=(attempt == 0)):
- if os.getenv("SKIP_AI_ANALYSIS", "false").lower() == "true":
- return False
- self.state.fix_errors = [str(WORKSPACE_DIR / BUILD_RESULT_FILE)]
- self.state.build_fix_count += 1
- print_warn(f"Build failed (attempt {attempt + 1}/{self.state.max_retries + 1}) — "
- f"will retry after AI fix")
- print_info(f"Build log: {WORKSPACE_DIR / BUILD_LOG_FILE}")
- continue
-
- # Build passed — tests are deferred to after all merges complete
- build_passed = True
- break
-
- if not build_passed:
- print_error(f"All {self.state.max_retries} fix attempts exhausted — build still failing")
- self.state.summary_rows.append(
- ("AI fix", "FAIL", f"Failed after {self.state.max_retries} attempts")
- )
- return False
-
- # Commit build fixes
- self._commit_fixes(ascend_path, step_dir)
-
- return True
-
- def _commit_submodule_if_needed(self) -> None:
- """Commit uncommitted changes inside the AscendNPU-IR submodule.
-
- Must be called BEFORE parent 'git add -A' so that the submodule
- pointer update is picked up by the parent commit.
- """
- ascend_path = Path(self.state.triton_ascend_path)
- if not submodule_has_changes(ascend_path):
- return
-
- target_short = self.state.target_commit[:12]
- commit_msg = (
- f"[Sync](fix) AI-generated build/test failures fix "
- f"for merging {target_short}\n\n"
- f"Upstream target: {target_short}\n"
- f"Fix attempt: {self.state.retry_count}\n"
- f"Work branch: {self.state.work_branch}\n"
- )
- commit_submodule(ascend_path, commit_msg)
-
- def _commit_fixes(self, ascend_path: Path, step_dir: Path) -> None:
- """Commit AI bug fixes with a meaningful message.
-
- Only commits if there are uncommitted changes. Commits submodule
- changes first (AscendNPU-IR), then returns to triton-ascend for the
- parent commit. Uses git add -A so AI-created files are not dropped.
-
- Commit message priority:
- 1. AI-written commit_message.txt (one-line subject)
- 2. First line of step_summary.md
- 3. Default generic message
- """
- # ── Commit submodule changes first (inside AscendNPU-IR) ──
- self._commit_submodule_if_needed()
-
- # ── Clean temp artifacts BEFORE staging ──
- # Clean first, then check status — otherwise temp files that
- # AI fixes didn't touch would cause a false-positive "need to commit".
- cleanup_temp_files(ascend_path)
-
- status = run_git(ascend_path, "status", "--porcelain").strip()
- if not status:
- print_info("No uncommitted fix changes — nothing to commit")
- return
-
- print_section("Commit Bug Fixes")
-
- target_short = self.state.target_commit[:12]
-
- # ── Read AI-written commit message ──
- commit_msg_path = step_dir / "commit_message.txt"
- if commit_msg_path.exists():
- commit_summary = commit_msg_path.read_text(encoding="utf-8").strip()
- # Take first line only for the subject
- commit_summary = commit_summary.split("\n")[0].strip()[:72]
- print_info(f"Using AI-written commit message: {commit_summary}")
- else:
- # Fallback: first line of step_summary.md
- summary_path = step_dir / EACH_STEP_SUMMARY_FILE
- if summary_path.exists():
- summary_text = summary_path.read_text(encoding="utf-8").strip()
- commit_summary = summary_text.split("\n")[0].lstrip("#").strip()[:72]
- else:
- commit_summary = f"Resolve build/test failures for merging {target_short}"
- commit_msg = (
- f"[Sync](fix) {commit_summary}\n\n"
- f"Upstream target: {target_short}\n"
- f"Fix attempt: {self.state.retry_count}\n"
- f"Work branch: {self.state.work_branch}\n"
- f"Co-Authored-By: Claude \n"
- )
-
- # ── Stage and commit (already changed to -A above via replace_all) ──
- try:
- staged_before = run_git(ascend_path, "diff", "--cached", "--name-only").strip()
- if not staged_before:
- # git add -A was already called; if no files staged yet, stage now
- run_git(ascend_path, "add", "-A")
- staged = run_git(ascend_path, "diff", "--cached", "--name-only").strip()
- if staged:
- print_info(f"Files staged for commit ({len(staged.splitlines())}):")
- for f in staged.splitlines()[:15]:
- print_info(f" - {f}")
- if len(staged.splitlines()) > 15:
- print_info(f" ... and {len(staged.splitlines()) - 15} more")
- run_git(ascend_path, "commit", "-s", "-m", commit_msg)
- print_status(True, f"Committed fix: {commit_summary[:60]}")
- self.state.summary_rows.append(("Commit fixes", "PASS", commit_summary[:40]))
- except subprocess.CalledProcessError as e:
- stderr = (e.stderr or "").strip() if hasattr(e, 'stderr') else str(e)
- if "nothing to commit" in stderr.lower():
- print_info("Nothing to commit (AI made no changes)")
- self.state.summary_rows.append(("Commit fixes", "PASS", "No changes"))
- else:
- print_warn(f"Could not commit fixes: {stderr[-200:]}")
- self.state.summary_rows.append(("Commit fixes", "WARN", stderr[:40]))
-
- def _do_build(self, ascend_path: Path, clean: bool = False,
- python_exe: str = "python3") -> bool:
- start_timer("build")
- print_section("Build Triton-Ascend")
-
- if os.getenv("SKIP_BUILD", "false").lower() == "true":
- print_info("SKIP_BUILD=true — skipping build")
- self.state.build_passed = True
- stop_timer("build")
- self.state.summary_rows.append(("Build", "SKIP", "SKIP_BUILD set"))
- return True
-
- build_result = build_triton_ascend(
- ascend_path,
- llvm_prefix=self.state.llvm_prefix,
- conda_env=self.state.conda_env,
- clean_build=clean,
- python_exe=python_exe,
- )
- self.state.build_passed = build_result["all_passed"]
- stop_timer("build")
-
- if not self.state.build_passed:
- print_error("Build FAILED")
- self.state.summary_rows.append(("Build", "FAIL", "See build log"))
- return False
-
- print_status(True, "Build passed")
- self.state.summary_rows.append(("Build", "PASS", ""))
- return True
-
- def _do_test(self, ascend_path: Path, python_exe: str = "") -> bool | None:
- start_timer("test")
- print_section("Run Tests")
-
- if os.getenv("SKIP_E2E_TEST", "false").lower() == "true":
- print_info("SKIP_E2E_TEST=true — treating tests as passed")
- self.state.test_passed = True
- stop_timer("test")
- self.state.summary_rows.append(("Tests", "SKIP", "SKIP_E2E_TEST set"))
- return None
-
- test_dir_path = ascend_path / self.state.test_dir
- py_label = python_exe or os.getenv("PYTHON", "python3.10")
- print_info(f"Test directory: {test_dir_path}")
- print_info(f"Python: {py_label}, procs: {self.state.num_procs}")
-
- try:
- test_result = run_tests(
- ascend_path,
- test_dir=self.state.test_dir,
- num_procs=self.state.num_procs,
- conda_env=self.state.conda_env,
- python_exe=python_exe,
- )
- except Exception as exc:
- print_error(f"run_tests raised exception: {exc}")
- import traceback
- traceback.print_exc()
- self.state.test_passed = False
- stop_timer("test")
- self.state.summary_rows.append(("Tests", "FAIL", f"Exception: {exc}"))
- return False
-
- self.state.test_passed = test_result["passed"]
- stop_timer("test")
-
- if test_result["passed"]:
- passed_count = test_result.get("passed_count", "?")
- print_status(True, f"All tests passed ({passed_count} passed)")
- self.state.summary_rows.append(("Tests", "PASS", f"{passed_count} passed"))
- return True
- else:
- failed_count = test_result.get("failed_count", "?")
- error_count = test_result.get("error_count", 0)
- error_msg = test_result.get("error", "")
- if error_msg:
- print_error(f"Tests FAILED — {error_msg}")
- else:
- print_error(f"Tests FAILED ({failed_count} failed, {error_count} errors)")
- self.state.summary_rows.append(
- ("Tests", "FAIL", f"{failed_count} failed, {error_count} errors")
- )
- return False
-
- def _do_ai_fix(self, ascend_path: Path, step_dir: Path, attempt: int) -> bool:
- """AI fix bug: invoke opencode/claude to fix build/test failures.
-
- AI context includes: step index, is_last_step, previous_step_summary
- (matching vllm-ascend's main2main_flow pattern).
- """
- print_step(attempt, self.state.max_retries, "AI fix attempt")
-
- step = self.state.steps[self.state.current_step] if self.state.steps else None
- current_step_id = step["id"] if step else "step-0"
- is_last_step = self.state.current_step == self.state.total_steps - 1
-
- # ── Previous step context (matching vllm-ascend pattern) ──
- previous_step = (
- self.state.steps[self.state.current_step - 1]
- if self.state.current_step > 0 and self.state.steps else None
- )
- previous_step_id = previous_step["id"] if previous_step else ""
- previous_step_summary_path = (
- str(WORKSPACE_DIR / STEPS_DIR / previous_step_id / EACH_STEP_SUMMARY_FILE)
- if previous_step_id else ""
- )
-
- # Per-attempt fix directory for logs/artifacts. The step_dir is the
- # canonical per-step directory (matching vllm-ascend pattern).
- fix_dir = WORKSPACE_DIR / FIX_LOG_DIR / f"{current_step_id}-fix-{attempt}"
- fix_dir.mkdir(parents=True, exist_ok=True)
-
- print_info(f"Error sources ({len(self.state.fix_errors)}):")
- for e in self.state.fix_errors:
- print(f" • {e}")
-
- # AI fix bug: detect backend
- try:
- from TA_main2main_workflow.agent.opencode_adapter import _detect_backend
- backend = _detect_backend()
- except RuntimeError as e:
- print_error(f"AI backend not available: {e}")
- self._last_ai_result = None
- return False
-
- print_ai_call_info(
- backend=backend,
- mode="fix",
- attempt=attempt,
- max_attempts=self.state.max_retries,
- )
-
- # AI fix bug: invoke opencode/claude with error logs
- # Context matches vllm-ascend pattern: is_last_step,
- # previous_step_id, previous_step_summary_path, step index.
- # step_dir points to the canonical step directory (like vllm-ascend);
- # fix_dir captures per-attempt fix artifacts separately.
- error_logs = json.dumps(self.state.fix_errors, ensure_ascii=False)
- try:
- ai_result = run_opencode_adapter({
- "step_id": f"{current_step_id}-fix-{attempt}",
- "previous_step_id": previous_step_id,
- "previous_step_summary_path": previous_step_summary_path,
- "is_last_step": str(is_last_step).lower(),
- "step_index": f"{self.state.current_step + 1}/{self.state.total_steps}",
- "step_dir": str(step_dir),
- "fix_dir": str(fix_dir),
- "conflict_dir": "",
- "ascend_path": str(ascend_path),
- "triton_path": self.state.triton_path,
- "reference_dir": _REFERENCE_DIR,
- "mode": "fix",
- "error_logs": error_logs,
- "target_commit": self.state.target_commit,
- })
-
- print_ai_result(
- ok=bool(ai_result.modified_files),
- modified_files=ai_result.modified_files,
- summary=(ai_result.step_summary or "")[:500],
- )
-
- # Store result for caller to capture fix details
- self._last_ai_result = {
- "modified_files": ai_result.modified_files,
- "step_summary": ai_result.step_summary or "",
- "is_noop": ai_result.is_noop,
- "elapsed_seconds": ai_result.elapsed_seconds,
- }
-
- print_info("Running pre-CI check after fix...")
- run_pre_ci_check(ascend_path, step_id=f"fix-{attempt}")
-
- return bool(ai_result.modified_files)
-
- except Exception as e:
- print_error(f"AI fix call failed: {e}")
- self._last_ai_result = None
- return False
-
- def _do_commit_step(self, step: dict) -> None:
- """Commit the current step's progress with a descriptive message.
-
- Only commits if there are uncommitted changes. Uses "git add -u" to
- avoid staging test artifacts or transient files.
-
- Commits AscendNPU-IR submodule changes first (if any), so the parent
- repo records the updated submodule pointer.
- """
- ascend_path = Path(self.state.triton_ascend_path)
- step_id = step["id"]
-
- # ── Commit submodule changes first ──
- self._commit_submodule_if_needed()
-
- status = run_git(ascend_path, "status", "--porcelain").strip()
-
- if not status:
- print_info(f"[{step_id}] No uncommitted changes — nothing to commit")
- self.state.summary_rows.append(
- (f"Commit {step_id}", "PASS", "No changes (clean merge)")
- )
- return
-
- print_section(f"Commit Step {step_id}")
-
- # Clean up temp artifacts before staging to avoid committing them
- cleanup_temp_files(ascend_path)
-
- end_commit_short = step["end_commit"][:12]
- commit_msg = (
- f"sync: merge upstream commits for step {step_id}\n\n"
- f"Upstream range: {step.get('start_commit', '?')[:12]}..{end_commit_short}\n"
- f"Step: {self.state.current_step + 1}/{self.state.total_steps}\n"
- f"Commits in step: {step['commit_count']}\n"
- f"Work branch: {self.state.work_branch}\n"
- f"All steps on single branch: {self.state.work_branch}\n"
- )
-
- try:
- run_git(ascend_path, "add", "-A")
- staged = run_git(ascend_path, "diff", "--cached", "--name-only").strip()
- if staged:
- print_info(f"Files staged ({len(staged.splitlines())}):")
- for f in staged.splitlines()[:10]:
- print_info(f" - {f}")
- if len(staged.splitlines()) > 10:
- print_info(f" ... and {len(staged.splitlines()) - 10} more")
- run_git(ascend_path, "commit", "-s", "-m", commit_msg)
- print_status(True, f"Committed step {step_id}")
- self.state.summary_rows.append(
- (f"Commit {step_id}", "PASS", f"{step['commit_count']} commits")
- )
- except subprocess.CalledProcessError as e:
- stderr = (e.stderr or "").strip() if hasattr(e, 'stderr') else str(e)
- if "nothing to commit" in stderr.lower():
- print_info(f"[{step_id}] Nothing to commit (clean merge)")
- self.state.summary_rows.append(
- (f"Commit {step_id}", "PASS", "No changes (clean merge)"))
- else:
- print_warn(f"Could not commit step {step_id}: {stderr[-200:]}")
- self.state.summary_rows.append(
- (f"Commit {step_id}", "WARN", stderr[:40]))
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Phase 3+4: IR Compatibility Patch Loop
- # ═══════════════════════════════════════════════════════════════════════════
-
- def _llvm_hash_did_change(self) -> bool:
- """Check if cmake/llvm-hash.txt differs from the Ascend baseline LLVM.
-
- The Ascend backend OP usage is based on a fixed baseline LLVM version.
- If the target LLVM hash differs from the baseline, IR compatibility
- patches need to be generated.
- """
- ascend_path = Path(self.state.triton_ascend_path)
- try:
- current_hash = (ascend_path / "cmake" / "llvm-hash.txt") \
- .read_text(encoding="utf-8").strip()
- except Exception:
- return False
-
- old_hash = _ASCEND_BASELINE_LLVM_HASH
- changed = old_hash != current_hash
- if changed:
- print_info(f"LLVM hash changed from baseline: "
- f"{old_hash[:12]} → {current_hash[:12]}")
- else:
- print_info("LLVM hash matches baseline — skipping IR patch phase")
- return changed
-
- def _do_ir_patch_loop(self) -> bool:
- """Phase 3+4 outer loop: IR analysis → patch → rebuild → test → fix.
-
- Runs AFTER all progressive merge steps have completed. If
- cmake/llvm-hash.txt didn't change, skips IR patches and goes
- directly to pytest.
-
- Outer loop (max IR_MAX_ITERATIONS rounds):
- [3.1-3.3] AI: analyze OPs, analyze changes, generate patches
- [3.4-3.5] Apply patches to LLVM + rebuild
- [4.1-4.2] Build TA + run pytest
- [4.3] If failures, AI classifies (IR vs code)
- → IR issues: loop back to modify patches
- → Code issues: AI fix inner loop
- Returns True if all tests pass, False on exhaustion.
- """
- ascend_path = Path(self.state.triton_ascend_path)
-
- # ── Skip IR patch phase via env var ──
- if os.getenv("SKIP_IR_PATCH", "false").lower() == "true":
- print_header("Phase 3+4: IR Patch + Pytest — SKIPPED (SKIP_IR_PATCH=true)")
- self.state.summary_rows.append(
- ("IR Patch", "SKIP", "SKIP_IR_PATCH set"))
- return True
-
- # ── Skip if LLVM hash unchanged ──
- self.state.llvm_hash_changed = self._llvm_hash_did_change()
- if not self.state.llvm_hash_changed:
- print_header("Phase 4: Pytest (LLVM unchanged)")
- print_info("LLVM hash unchanged — skipping IR analysis and patch generation")
- return self._do_pytest()
-
- print_header("Phase 3: IR Compatibility Patch Auto-Generation")
- print_info(f"LLVM hash changed — IR compatibility analysis required")
- print_key_value("Baseline LLVM", _ASCEND_BASELINE_LLVM_HASH[:12])
-
- self._print_workspace_info("Phase 3: IR Patch Loop")
- print_key_value("Max IR iterations", str(self.state.ir_max_iterations))
-
- for iteration in range(self.state.ir_max_iterations):
- self.state.ir_patch_iteration = iteration
- print_header(
- f"IR Patch Loop — Iteration {iteration + 1}/"
- f"{self.state.ir_max_iterations}"
- )
-
- # ── [3.1 + 3.2] Analysis (only on first iteration) ──
- if iteration == 0:
- print_info("First iteration — running full OP analysis pipeline")
- if not self._do_ir_op_analysis():
- return False
- if not self._do_ir_change_analysis():
- return False
- else:
- # On retry, re-analyze changes (patches from previous
- # iteration may have altered the picture)
- print_info("Re-analyzing OP changes after patch retry...")
- if not self._do_ir_change_analysis():
- return False
-
- # ── [3.3] Generate patches ──
- print_info("Step 3.3: Invoking AI to generate IR compatibility patches...")
- if not self._do_ir_generate_patches():
- return False
-
- # ── [3.4 + 3.5] Apply patches + rebuild LLVM ──
- print_info("Step 3.4-3.5: Applying patches and rebuilding LLVM (this may take a while)...")
- if not self._do_ir_apply_patches_and_rebuild():
- return False
-
- # ── [4.1] Build TA ──
- print_info("Step 4.1: Building Triton-Ascend with patched LLVM...")
- if not self._do_build(ascend_path, clean=True):
- if os.getenv("SKIP_AI_ANALYSIS", "false").lower() == "true":
- return False
- self.state.fix_errors = [str(WORKSPACE_DIR / BUILD_RESULT_FILE)]
- self.state.build_fix_count += 1
- print_warn("Build failed after IR patches — will attempt AI fix")
- self._do_ai_fix(ascend_path, WORKSPACE_DIR, 1)
- continue
-
- # ── [4.2] Pytest ──
- print_info("Step 4.2: Running pytest suite...")
- if self._do_pytest():
- print_status(True, "All tests pass!")
- self._commit_fixes(ascend_path, WORKSPACE_DIR)
- self.state.ir_loop_details.append({
- "iteration": iteration + 1,
- "result": "ALL_PASS",
- })
- return True
-
- # ── [4.3] Diagnose failures ──
- print_info("Step 4.3: Invoking AI to classify test failures (IR vs code)...")
- has_ir_issues = self._do_ir_diagnose_failures()
- if has_ir_issues:
- self.state.ir_issues_found += 1
- print_warn(
- f"IR compatibility issues found in iteration "
- f"{iteration + 1} — retrying with modified patches"
- )
- self.state.ir_loop_details.append({
- "iteration": iteration + 1,
- "result": "IR_RETRY",
- "ir_issues": self.state.ir_issues_found,
- })
- continue
-
- # ── [4.4] Non-IR issues → AI fix inner loop ──
- print_info("Non-IR failures detected — entering AI fix loop")
- print_key_value("Max fix attempts", str(self.state.max_retries))
- for fix_attempt in range(1, self.state.max_retries + 1):
- print_header(f"AI Fix Attempt {fix_attempt}/{self.state.max_retries}")
- self.state.retry_count = fix_attempt
- self._do_ai_fix(ascend_path, WORKSPACE_DIR, fix_attempt)
- if not self._do_build(ascend_path, clean=False):
- self.state.fix_errors = [str(WORKSPACE_DIR / BUILD_RESULT_FILE)]
- continue
- if self._do_pytest():
- self._commit_fixes(ascend_path, WORKSPACE_DIR)
- self.state.ir_loop_details.append({
- "iteration": iteration + 1,
- "result": "PASS_AFTER_FIX",
- "fix_attempts": fix_attempt,
- })
- return True
-
- print_error(f"All {self.state.max_retries} fix attempts exhausted "
- f"in iteration {iteration + 1}")
- self.state.ir_loop_details.append({
- "iteration": iteration + 1,
- "result": "FIX_EXHAUSTED",
- })
-
- print_error(f"IR patch loop exhausted {self.state.ir_max_iterations} "
- f"iterations")
- return False
-
- def _do_ir_op_analysis(self) -> bool:
- """[3.1] AI analyzes which MLIR OPs the Ascend backend uses."""
- print_header("Phase 3.1: IR OP Analysis")
- ascend_path = Path(self.state.triton_ascend_path)
-
- self._print_workspace_info("Phase 3.1: IR OP Analysis")
-
- ir_dir = WORKSPACE_DIR / IR_ANALYSIS_DIR
- ir_dir.mkdir(parents=True, exist_ok=True)
-
- from TA_main2main_workflow.agent.opencode_adapter import _detect_backend
- backend = _detect_backend()
- print_info(f"AI backend: {backend}")
- print_key_value("Triton-Ascend", str(ascend_path))
- print_key_value("Output dir", str(ir_dir))
-
- # ── Pre-scan: find candidate files with MLIR OP usage ──
- print_info("Pre-scanning Ascend backend for MLIR OP patterns...")
- candidate_files: list[str] = []
- scan_dirs = [
- ascend_path / "third_party" / "ascend" / "lib",
- ascend_path / "lib" / "Target" / "Ascend",
- ]
- op_patterns = [
- r'::create\b', r'::get\b', r'\.match\b', r'\.walk\b',
- r'isa<', r'cast<', r'dyn_cast<',
- ]
- for sd in scan_dirs:
- if not sd.exists():
- print_warn(f"Scan dir not found: {sd}")
- continue
- for pattern in op_patterns:
- try:
- result = subprocess.run(
- ["grep", "-rl", pattern, str(sd)],
- capture_output=True, text=True, timeout=30,
- )
- for f in result.stdout.splitlines():
- if f not in candidate_files:
- candidate_files.append(f)
- except (subprocess.TimeoutExpired, Exception):
- pass
-
- candidate_files.sort()
- print_info(f"Found {len(candidate_files)} candidate files with MLIR OP patterns")
- for f in candidate_files[:15]:
- print_info(f" - {Path(f).relative_to(ascend_path)}")
- if len(candidate_files) > 15:
- print_info(f" ... and {len(candidate_files) - 15} more files")
-
- # Write candidate file list for AI reference
- hint_path = ir_dir / "candidate_files.txt"
- hint_path.write_text("\n".join(candidate_files), encoding="utf-8")
- print_info(f"Candidate file list written to {hint_path}")
-
- print_info("AI will scan candidate files for MLIR OP usage and output structured JSON")
- print_info("Invoking AI for IR OP analysis (this may take several minutes)...")
-
- try:
- ai_result = run_opencode_adapter({
- "step_id": "ir-analyze-ops",
- "previous_step_id": "",
- "previous_step_summary_path": "",
- "is_last_step": "true",
- "step_index": "ir",
- "step_dir": str(ir_dir),
- "fix_dir": str(ir_dir),
- "conflict_dir": "",
- "ascend_path": str(ascend_path),
- "triton_path": self.state.triton_path,
- "reference_dir": _REFERENCE_DIR,
- "mode": "ir_analyze_ops",
- "error_logs": "[]",
- "target_commit": self.state.target_commit,
- "llvm_project_path": str(_llvm_project_path()),
- })
- _ = ai_result
- except Exception as e:
- print_error(f"IR OP analysis failed: {e}")
- self.state.summary_rows.append(("IR OP Analysis", "FAIL", str(e)[:60]))
- return False
-
- ops_report = ir_dir / IR_OPS_REPORT_FILE
- if ops_report.exists():
- try:
- data = json.loads(ops_report.read_text(encoding="utf-8"))
- # ── Content validation: must have 'ops' array with real OP data ──
- ops_list = data.get("ops", [])
- if not ops_list or not isinstance(ops_list, list):
- print_error(
- f"AI output is NOT a valid OP report! "
- f"Missing or empty 'ops' array. "
- f"Top-level keys: {list(data.keys())}")
- print_warn(
- f"AI may have produced a merge analysis instead of IR OP scan. "
- f"Check {ops_report} for content.")
- self.state.summary_rows.append(
- ("IR OP Analysis", "FAIL",
- f"No 'ops' array — AI produced wrong output type"))
- return False
- # Check that ops have expected fields
- valid_ops = [o for o in ops_list if isinstance(o, dict) and "name" in o]
- if len(valid_ops) < len(ops_list):
- print_warn(
- f"{len(ops_list) - len(valid_ops)} entries missing 'name' field — filtered")
- if not valid_ops:
- print_error("No valid OP entries with 'name' field found!")
- self.state.summary_rows.append(
- ("IR OP Analysis", "FAIL", "No valid OP entries"))
- return False
-
- self.state.ir_ops_report = data
- dialects = data.get("dialects", [])
- print_status(True,
- f"OP analysis complete: "
- f"{data.get('total_ops', len(valid_ops))} OPs, "
- f"{len(dialects)} dialects — "
- f"{', '.join(dialects[:10])}")
- self.state.summary_rows.append(
- ("IR OP Analysis", "PASS",
- f"{data.get('total_ops', len(valid_ops))} OPs"))
- return True
- except Exception as e:
- print_warn(f"Could not parse ops report: {e}")
-
- self.state.summary_rows.append(("IR OP Analysis", "FAIL", "No report"))
- return False
-
- def _do_ir_change_analysis(self) -> bool:
- """[3.2] AI analyzes OP definition changes between LLVM versions."""
- print_header("Phase 3.2: IR OP Change Analysis")
- ascend_path = Path(self.state.triton_ascend_path)
-
- self._print_workspace_info("Phase 3.2: IR Change Analysis")
-
- ir_dir = WORKSPACE_DIR / IR_ANALYSIS_DIR
- ir_dir.mkdir(parents=True, exist_ok=True)
-
- baseline_hash = _ASCEND_BASELINE_LLVM_HASH
- llvm_project = _llvm_project_path()
-
- # Read target LLVM hash from ascend repo
- llvm_hash_file = ascend_path / "cmake" / "llvm-hash.txt"
- if not llvm_hash_file.exists():
- print_error(f"llvm-hash.txt not found at {llvm_hash_file}")
- self.state.summary_rows.append(
- ("IR Change Analysis", "FAIL", "llvm-hash.txt missing"))
- return False
- target_hash = llvm_hash_file.read_text(encoding="utf-8").strip()
-
- print_key_value("Input ops report", str(ir_dir / IR_OPS_REPORT_FILE))
- print_key_value("LLVM project", str(llvm_project))
- print_key_value("Baseline LLVM", f"{baseline_hash[:12]} ({baseline_hash})")
- print_key_value("Target LLVM", f"{target_hash[:12]} ({target_hash})")
-
- # ── Pre-flight: verify both commits exist in llvm-project ──
- if not llvm_project.exists():
- print_error(f"llvm-project not found at {llvm_project}")
- self.state.summary_rows.append(
- ("IR Change Analysis", "FAIL", "llvm-project not found"))
- return False
-
- print_info("Verifying LLVM commits are available in llvm-project...")
- for label, h in [("Baseline", baseline_hash), ("Target", target_hash)]:
- try:
- result = subprocess.run(
- ["git", "cat-file", "-t", h],
- cwd=str(llvm_project),
- capture_output=True, text=True, timeout=30,
- )
- if result.returncode == 0:
- print_status(True, f"{label} commit {h[:12]} — found in llvm-project")
- continue
-
- # ── Commit not found locally — try fetching from origin ──
- print_warn(
- f"{label} commit {h[:12]} NOT found locally — "
- f"fetching from origin...")
- fetched = False
- for attempt in range(1, 7):
- fetch_proc = subprocess.run(
- ["git", "fetch", "origin", h, "--no-tags"],
- cwd=str(llvm_project),
- capture_output=True, text=True, timeout=300,
- )
- if fetch_proc.returncode == 0:
- fetched = True
- print_status(True,
- f"{label} commit {h[:12]} — fetched (attempt {attempt})")
- break
- print_warn(
- f"Fetch attempt {attempt}/6 for {label} commit "
- f"{h[:12]} failed — retrying...")
- if not fetched:
- print_error(
- f"{label} commit {h[:12]} NOT found in llvm-project "
- f"after 6 fetch attempts! "
- f"(git cat-file -t returned: {result.stderr.strip()})")
- self.state.summary_rows.append(
- ("IR Change Analysis", "FAIL",
- f"{label} commit {h[:12]} not in llvm-project"))
- return False
- except subprocess.TimeoutExpired:
- print_error(f"Timeout checking {label} commit {h[:12]}")
- return False
- except Exception as e:
- print_error(f"Failed to verify {label} commit: {e}")
- return False
-
- # ── Pre-flight: show MLIR .td file changes between the two commits ──
- print_info("Scanning MLIR .td file changes between baseline and target...")
- try:
- diff_result = subprocess.run(
- ["git", "diff", "--name-only", baseline_hash, target_hash,
- "--", "mlir/include/"],
- cwd=str(llvm_project),
- capture_output=True, text=True, timeout=60,
- )
- if diff_result.returncode == 0:
- changed_files = [f for f in diff_result.stdout.splitlines()
- if f.endswith(".td")]
- print_info(f"Found {len(changed_files)} changed .td files in mlir/include/ "
- f"between baseline and target")
- for f in changed_files[:20]:
- print_info(f" - {f}")
- if len(changed_files) > 20:
- print_info(f" ... and {len(changed_files) - 20} more .td files")
- else:
- print_warn(f"git diff returned non-zero: {diff_result.stderr.strip()}")
- except subprocess.TimeoutExpired:
- print_warn("git diff timed out after 60s — continuing anyway")
- except Exception as e:
- print_warn(f"Could not run git diff for .td files: {e}")
-
- # ── Pre-flight: show ops report summary for AI context ──
- ops_report_path = ir_dir / IR_OPS_REPORT_FILE
- if ops_report_path.exists():
- try:
- ops = json.loads(ops_report_path.read_text(encoding="utf-8"))
- print_info(
- f"Ops report: {ops.get('total_ops', '?')} OPs across "
- f"{len(ops.get('dialects', []))} dialects — "
- f"{', '.join(ops.get('dialects', [])[:8])}")
- except Exception:
- print_warn("Could not read ops_report.json for summary")
- else:
- print_warn(f"Ops report not found at {ops_report_path} — "
- f"AI will need to discover OPs on its own")
-
- print_info("AI will compare each OP's .td definition with:")
- print_info(f" git show {baseline_hash[:12]}:mlir/include/.../.td")
- print_info(f" git show {target_hash[:12]}:mlir/include/.../.td")
- print_info("Invoking AI for OP change analysis (this may take several minutes)...")
-
- try:
- ai_result = run_opencode_adapter({
- "step_id": "ir-analyze-changes",
- "previous_step_id": "ir-analyze-ops",
- "previous_step_summary_path": str(ir_dir / IR_OPS_REPORT_FILE),
- "is_last_step": "true",
- "step_index": "ir",
- "step_dir": str(ir_dir),
- "fix_dir": str(ir_dir),
- "conflict_dir": "",
- "ascend_path": str(ascend_path),
- "triton_path": self.state.triton_path,
- "reference_dir": _REFERENCE_DIR,
- "mode": "ir_analyze_changes",
- "error_logs": json.dumps(
- [str(ir_dir / IR_OPS_REPORT_FILE)], ensure_ascii=False),
- "target_commit": self.state.target_commit,
- "llvm_project_path": str(_llvm_project_path()),
- "baseline_llvm_hash": _ASCEND_BASELINE_LLVM_HASH,
- "target_llvm_hash": target_hash,
- })
- _ = ai_result
- except Exception as e:
- print_error(f"IR change analysis failed: {e}")
- self.state.summary_rows.append(
- ("IR Change Analysis", "FAIL", str(e)[:60]))
- return False
-
- changes_report = ir_dir / IR_CHANGES_REPORT_FILE
- if changes_report.exists():
- try:
- data = json.loads(changes_report.read_text(encoding="utf-8"))
- # ── Content validation: must have 'changes' array and 'summary' ──
- changes_list = data.get("changes", [])
- summary = data.get("summary", {})
- if not changes_list or not isinstance(changes_list, list):
- print_error(
- f"AI output is NOT a valid changes report! "
- f"Missing or empty 'changes' array. "
- f"Top-level keys: {list(data.keys())}")
- print_warn(
- f"AI may have produced a merge analysis instead of "
- f"OP change comparison. Check {changes_report} for content.")
- self.state.summary_rows.append(
- ("IR Change Analysis", "FAIL",
- "No 'changes' array — AI produced wrong output type"))
- return False
-
- self.state.ir_changes_report = data
- print_status(True,
- f"Change analysis: {summary.get('total_ops_analyzed', '?')} "
- f"OPs, {summary.get('ops_needing_patch', '?')} need patch, "
- f"{summary.get('renamed_ops', 0)} renamed, "
- f"{summary.get('signature_changes', 0)} signature changes")
- self.state.summary_rows.append(
- ("IR Change Analysis", "PASS",
- f"{summary.get('ops_needing_patch', '?')} OPs need patch"))
- # If no OPs need patching, still return True (Phase 3 is a no-op)
- return True
- except Exception as e:
- print_warn(f"Could not parse changes report: {e}")
-
- self.state.summary_rows.append(
- ("IR Change Analysis", "FAIL", "No report"))
- return False
-
- def _do_ir_generate_patches(self) -> bool:
- """[3.3] AI modifies the Ascend LLVM patch for IR compatibility.
-
- The AI directly edits the existing patch file at
- ``third_party/ascend/patch/llvm_patch_f6ded0b.patch`` rather than
- creating a new file from scratch — this lets it start from a known-
- working baseline and only adjust the parts that need changing for
- the current LLVM version.
- """
- print_header("Phase 3.3: IR Patch Generation")
- ascend_path = Path(self.state.triton_ascend_path)
-
- self._print_workspace_info("Phase 3.3: IR Patch Generation")
-
- ir_dir = WORKSPACE_DIR / IR_ANALYSIS_DIR
-
- # The patch file that AI modifies in-place
- ascend_patch = (ascend_path / "third_party" / "ascend" / "patch"
- / "llvm_patch_f6ded0b.patch")
- print_key_value("Target patch", str(ascend_patch))
-
- changes_report = ir_dir / IR_CHANGES_REPORT_FILE
- if changes_report.exists():
- try:
- report = json.loads(changes_report.read_text(encoding="utf-8"))
- summary = report.get("summary", {})
- print_info(f"Changes report: {summary.get('total_ops_analyzed', '?')} OPs analyzed, "
- f"{summary.get('ops_needing_patch', '?')} need patches")
- except Exception:
- pass
- print_info("Invoking AI to modify the Ascend LLVM compatibility patch...")
-
- try:
- ai_result = run_opencode_adapter({
- "step_id": "ir-generate-patch",
- "previous_step_id": "ir-analyze-changes",
- "previous_step_summary_path": str(ir_dir / IR_CHANGES_REPORT_FILE),
- "is_last_step": "true",
- "step_index": "ir",
- "step_dir": str(ascend_patch.parent),
- "fix_dir": str(ascend_patch.parent),
- "conflict_dir": "",
- "ascend_path": str(ascend_path),
- "triton_path": self.state.triton_path,
- "reference_dir": _REFERENCE_DIR,
- "mode": "ir_generate_patch",
- "error_logs": json.dumps(
- [str(ir_dir / IR_CHANGES_REPORT_FILE)], ensure_ascii=False),
- "target_commit": self.state.target_commit,
- "llvm_project_path": str(_llvm_project_path()),
- "baseline_llvm_hash": _ASCEND_BASELINE_LLVM_HASH,
- "ascend_patch_file": str(ascend_patch),
- })
- _ = ai_result
- except Exception as e:
- print_error(f"IR patch generation failed: {e}")
- self.state.summary_rows.append(
- ("IR Patch Gen", "FAIL", str(e)[:60]))
- return False
-
- # Check the ascend patch was modified
- if ascend_patch.exists():
- print_status(True, f"Modified {ascend_patch.name}")
- self.state.ir_patches = [str(ascend_patch)]
- self.state.summary_rows.append(
- ("IR Patch Gen", "PASS", ascend_patch.name))
- return True
-
- # No changes needed — valid if changes_report showed no issues
- print_info(f"{ascend_patch.name} unchanged — "
- "IR compatibility may already be satisfied")
- self.state.summary_rows.append(
- ("IR Patch Gen", "PASS", "No changes needed"))
- return True
-
- def _do_ir_apply_patches_and_rebuild(self) -> bool:
- """[3.4 + 3.5] Apply the Ascend LLVM patch and rebuild."""
- print_header("Phase 3.4-3.5: Apply Patches + Rebuild LLVM")
- ascend_path = Path(self.state.triton_ascend_path)
-
- self._print_workspace_info("Phase 3.4-3.5: Apply Patches + Rebuild LLVM")
-
- llvm_project = _llvm_project_path()
- # The in-repo Ascend LLVM patch (modified by AI in step 3.3)
- ascend_patch = (ascend_path / "third_party" / "ascend" / "patch"
- / "llvm_patch_f6ded0b.patch")
- print_key_value("LLVM project", str(llvm_project))
- print_key_value("Patch file", str(ascend_patch))
-
- if not llvm_project.exists():
- print_error(f"LLVM project not found at {llvm_project}")
- self.state.summary_rows.append(
- ("IR Apply+Rebuild", "FAIL", "llvm-project not found"))
- return False
-
- # ── Ensure llvm-project workspace is clean before checkout + patch ──
- if not self._ensure_llvm_workspace_clean(reason="ir-apply-patches"):
- print_error("Cannot clean llvm-project workspace — aborting IR patch rebuild")
- self.state.summary_rows.append(
- ("IR Apply+Rebuild", "FAIL", "workspace not clean"))
- return False
-
- # ── [3.4] Apply patch to llvm-project ──
- # Deterministic: clean → checkout → apply. No AI involved.
- from TA_main2main_workflow.scripts.build_test import apply_llvm_patches
-
- # Read the target LLVM hash from triton-ascend
- llvm_hash_file = ascend_path / "cmake" / "llvm-hash.txt"
- target_llvm_hash = ""
- if llvm_hash_file.exists():
- target_llvm_hash = llvm_hash_file.read_text(encoding="utf-8").strip()
- print_key_value("Target LLVM hash", target_llvm_hash[:12])
-
- print_info(f"Step 3.4: Applying {ascend_patch.name} to llvm-project...")
- patch_result = apply_llvm_patches(
- ascend_patch.parent, llvm_project,
- target_hash=target_llvm_hash, patch_file=ascend_patch)
- if not patch_result["all_ok"]:
- failed = patch_result["failed"]
- print_error(f"LLVM patch apply failed: "
- f"{failed[0]['error'][:200] if failed else 'unknown'}")
- self.state.summary_rows.append(
- ("IR Apply+Rebuild", "FAIL", "patch did not apply cleanly"))
- return False
-
- print_status(True, f"{ascend_patch.name} applied to llvm-project")
-
- # ── Show git status after patch for debugging ──
- status_proc = subprocess.run(
- ["git", "status", "--short"],
- cwd=str(llvm_project), capture_output=True, text=True, timeout=30,
- )
- if status_proc.stdout.strip():
- print_info("llvm-project git status after patch:")
- for line in status_proc.stdout.strip().splitlines():
- print(f" {line}")
- else:
- print_info("llvm-project working tree is clean after patch")
-
- # ── [3.5] Rebuild LLVM ──
- from TA_main2main_workflow.scripts.build_test import build_llvm
- try:
- print_info("Step 3.5: Rebuilding LLVM (this takes ~15-30 minutes)...")
- # Patch is already applied — just build, no hash/checkout logic.
- llvm_prefix = build_llvm(
- llvm_project,
- Path(os.path.expanduser(
- os.getenv("LLVM_INSTALL_PREFIX_SYNC", "~/llvm-install-sync"))),
- required_hash=target_llvm_hash,
- )
- if llvm_prefix and not self.state.llvm_prefix:
- self.state.llvm_prefix = llvm_prefix
- print_status(True, "LLVM rebuild complete")
- self.state.summary_rows.append(
- ("LLVM Patch Apply+Rebuild", "PASS", "patch applied, LLVM rebuilt"))
- return True
- except Exception as e:
- print_error(f"LLVM rebuild failed: {e}")
- self.state.summary_rows.append(
- ("LLVM Patch Apply+Rebuild", "FAIL", str(e)[:60]))
- return False
-
- def _do_pytest(self) -> bool:
- """[4.2] Build TA and run pytest.
-
- Returns True when all tests pass.
- """
- ascend_path = Path(self.state.triton_ascend_path)
- py_exe = os.getenv("PYTHON", "python3.10")
-
- import shutil
- if not shutil.which(py_exe):
- print_warn(f"{py_exe} not found on PATH — skipping tests")
- self.state.pytest_passed = False
- self.state.summary_rows.append(
- ("Pytest", "SKIP", f"{py_exe} not found"))
- return False
-
- print_section(f"Pytest ({py_exe})")
- print_key_value("Ascend path", str(ascend_path))
-
- # Build with test python
- print_info(f"Building Triton-Ascend with {py_exe}...")
- if not self._do_build(ascend_path, clean=True, python_exe=py_exe):
- print_error(f"Build failed ({py_exe})")
- self.state.pytest_passed = False
- self.state.summary_rows.append(
- ("Pytest", "FAIL", "Build failed"))
- return False
-
- # Run tests
- result = self._do_test(ascend_path, python_exe=py_exe)
- if result is None:
- passed = True # SKIP_E2E_TEST
- else:
- passed = bool(result)
-
- self.state.pytest_passed = passed
- if not passed:
- print_error(f"Pytest FAILED ({py_exe})")
-
- self.state.summary_rows.append(
- ("Pytest", "PASS" if passed else "FAIL", py_exe))
- return passed
-
- def _do_ir_diagnose_failures(self) -> bool:
- """[4.3] AI classifies test failures: IR compatibility vs code issues.
-
- Returns True if IR issues are found (triggering outer loop retry).
- Returns False if failures are all code/environment issues.
- """
- print_header("Phase 4.3: IR Failure Diagnosis")
-
- ir_dir = WORKSPACE_DIR / IR_ANALYSIS_DIR
- ir_dir.mkdir(parents=True, exist_ok=True)
- print_key_value("Diagnosis output", str(ir_dir / IR_DIAGNOSIS_FILE))
-
- # Collect test failure logs from both Python runs
- error_log_paths: list[str] = []
- test_log_dir = WORKSPACE_DIR / "test-logs"
- if test_log_dir.exists():
- for log_file in sorted(test_log_dir.rglob("*.log")):
- error_log_paths.append(str(log_file))
- # Also include test result files
- test_result = WORKSPACE_DIR / TEST_RESULT_FILE
- if test_result.exists():
- error_log_paths.append(str(test_result))
-
- if not error_log_paths:
- print_warn("No test failure logs found — assuming code issues")
- return False
-
- print_info(f"Collected {len(error_log_paths)} log file(s) for AI diagnosis")
- for p in error_log_paths[:5]:
- print_info(f" - {p}")
- if len(error_log_paths) > 5:
- print_info(f" ... and {len(error_log_paths) - 5} more")
- print_info("Invoking AI to classify failures (IR compatibility vs code vs environment)...")
-
- try:
- ai_result = run_opencode_adapter({
- "step_id": "ir-diagnose",
- "previous_step_id": "ir-generate-patch",
- "previous_step_summary_path": str(ir_dir / IR_CHANGES_REPORT_FILE),
- "is_last_step": "true",
- "step_index": "ir",
- "step_dir": str(ir_dir),
- "fix_dir": str(ir_dir),
- "conflict_dir": "",
- "ascend_path": str(Path(self.state.triton_ascend_path)),
- "triton_path": self.state.triton_path,
- "reference_dir": _REFERENCE_DIR,
- "mode": "ir_diagnose",
- "error_logs": json.dumps(error_log_paths, ensure_ascii=False),
- "target_commit": self.state.target_commit,
- })
- _ = ai_result
- except Exception as e:
- print_error(f"IR diagnosis failed: {e}")
- return False
-
- diagnosis_path = ir_dir / IR_DIAGNOSIS_FILE
- if not diagnosis_path.exists():
- print_warn("No diagnosis report generated")
- return False
-
- try:
- diagnosis = json.loads(
- diagnosis_path.read_text(encoding="utf-8"))
- summary = diagnosis.get("summary", {})
- has_ir = summary.get("has_ir_issues", False)
- print_key_value("total failures", str(summary.get("total_failures", "?")))
- print_key_value("IR issues", str(summary.get("ir_issues", "?")))
- print_key_value("code issues", str(summary.get("code_issues", "?")))
- print_key_value("env issues", str(summary.get("environment_issues", "?")))
- self.state.summary_rows.append(
- ("IR Diagnosis", "PASS",
- f"IR={summary.get('ir_issues', '?')} "
- f"code={summary.get('code_issues', '?')} "
- f"env={summary.get('environment_issues', '?')}"))
- return bool(has_ir)
- except Exception as e:
- print_warn(f"Could not parse diagnosis: {e}")
- return False
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Per-step IR patch pipeline (single-step mode)
- # ═══════════════════════════════════════════════════════════════════════════
-
- def _do_per_step_ir_patch(self, step: dict) -> bool:
- """Per-step IR compatibility patch generation + LLVM rebuild.
-
- Called from _run_single_step_mode() when a step's merge included an
- LLVM hash change. Reuses the shared IR analysis / patch-generation
- methods but adds a patch→rebuild retry loop specific to the
- single-step context.
-
- Pipeline:
- 1. Verify LLVM hash changed
- 2. Create llvm_change_analysis// workspace
- 3. IR analysis → patch → apply → rebuild LLVM (max 3 iterations)
- 4. On patch failure: stash/drop, loop back for AI to fix
- """
- step_id = step["id"]
-
- # ── Guard: check LLVM hash actually changed ──
- if not self._llvm_hash_did_change():
- print_info(f"[{step_id}] LLVM hash unchanged — skipping IR patch")
- return True
-
- # ── Create per-step analysis workspace ──
- analysis_dir = WORKSPACE_DIR / LLVM_CHANGE_ANALYSIS_DIR / step_id
- analysis_dir.mkdir(parents=True, exist_ok=True)
- print_key_value("IR analysis dir", str(analysis_dir))
-
- # ── IR analysis → patch → rebuild loop ──
- for iteration in range(self.state.ir_max_iterations):
- self.state.ir_patch_iteration = iteration
- print_header(
- f"Per-Step IR Patch — {step_id} "
- f"(iter {iteration + 1}/{self.state.ir_max_iterations})"
- )
-
- # [3.1 + 3.2] Analysis (first iteration only for OP scan)
- if iteration == 0:
- print_info("First iteration — running full OP analysis pipeline")
- if not self._do_ir_op_analysis():
- return False
- if not self._do_ir_change_analysis():
- return False
- else:
- print_info("Re-analyzing OP changes after patch retry...")
- if not self._do_ir_change_analysis():
- return False
-
- # [3.3] Generate patches
- if not self._do_ir_generate_patches():
- return False
-
- # [3.4 + 3.5] Apply patches + rebuild LLVM (with retry for patch failures)
- rebuild_ok = False
- for patch_attempt in range(IR_MAX_ITERATIONS):
- print_info(
- f"Patch apply attempt {patch_attempt + 1}/{IR_MAX_ITERATIONS}"
- )
- if self._do_ir_apply_patches_and_rebuild():
- rebuild_ok = True
- break
- # Patch failed — stash/drop, let AI regenerate
- print_warn(
- f"LLVM rebuild failed (patch attempt {patch_attempt + 1}) — "
- f"will stash changes and retry patch generation"
- )
- self._stash_and_drop_llvm_patch()
- if not self._do_ir_generate_patches():
- break
-
- if rebuild_ok:
- print_status(True, f"IR patch + LLVM rebuild OK for {step_id}")
- self.state.ir_loop_details.append({
- "step_id": step_id,
- "iteration": iteration + 1,
- "result": "PASS",
- })
- return True
-
- print_warn(f"IR patch iteration {iteration + 1} exhausted "
- f"— retrying outer loop")
-
- print_error(f"IR patch loop exhausted {self.state.ir_max_iterations} "
- f"iterations for {step_id}")
- return False
-
- def _build_baseline_llvm(self) -> bool:
- """Build baseline LLVM (pre-merge state) before any merge steps.
-
- Called once at the start of _run_single_step_mode(). Reads the
- current cmake/llvm-hash.txt from triton-ascend, checks out that
- commit in llvm-project, applies the Ascend backend LLVM patch,
- builds LLVM, then stashes + drops the patch to leave a clean tree.
-
- The baseline LLVM must be built before merging because the Ascend
- backend code depends on it for compilation.
- """
- print_header("Build Baseline LLVM (pre-merge)")
- ascend_path = Path(self.state.triton_ascend_path)
-
- self._print_workspace_info("Build Baseline LLVM")
-
- # ── Allow skipping baseline LLVM build for debugging ──
- if os.getenv("SKIP_BASELINE_LLVM", "false").lower() == "true":
- print_info("SKIP_BASELINE_LLVM=true — skipping baseline LLVM build")
- print_warn("Ensure LLVM is already built at LLVM_INSTALL_PREFIX_SYNC")
- if not self.state.llvm_prefix:
- self.state.llvm_prefix = str(_llvm_install_prefix())
- self.state.summary_rows.append(
- ("Baseline LLVM", "SKIP", "SKIP_BASELINE_LLVM set"))
- return True
-
- llvm_project = _llvm_project_path()
- llvm_install = _llvm_install_prefix()
-
- if not llvm_project.exists():
- print_error(f"llvm-project not found at {llvm_project}")
- return False
-
- # ── 1. Read LLVM hash from base branch (work branch base) ──
- # Use git show to get the hash from the base branch, NOT the checkout
- # filesystem — the checkout may be on a stale branch.
- base_branch = os.getenv(ENV_BASE_BRANCH, "main")
- base_ref = get_base_branch_ref()
- try:
- run_git(ascend_path, "fetch", "origin", base_branch)
- except Exception:
- print_warn(f"[baseline-llvm] Could not fetch {base_ref}, using local ref")
- try:
- llvm_hash = run_git(
- ascend_path, "show", f"{base_ref}:cmake/llvm-hash.txt"
- ).strip()
- except Exception:
- # Fallback: read from checkout filesystem
- llvm_hash_file = ascend_path / "cmake" / "llvm-hash.txt"
- if not llvm_hash_file.exists():
- print_error(f"LLVM hash file not found: {llvm_hash_file}")
- return False
- llvm_hash = llvm_hash_file.read_text(encoding="utf-8").strip()
- print_warn(f"[baseline-llvm] Using checkout llvm-hash.txt ({base_ref} not available)")
- if not llvm_hash:
- print_error("LLVM hash is empty")
- return False
- print_key_value("LLVM commit", llvm_hash[:12])
- print_info(f" (from {base_ref})")
-
- # ── Ensure llvm-project workspace is clean before checkout ──
- if not self._ensure_llvm_workspace_clean(reason="baseline-llvm-build"):
- print_error("Cannot clean llvm-project workspace — aborting baseline build")
- return False
-
- # ── 2. Checkout the LLVM commit ──
- print_info(f"Checking out LLVM commit {llvm_hash[:12]} in llvm-project...")
- try:
- # Fetch the specific commit with retries
- for attempt in range(1, 7):
- fetch_proc = subprocess.run(
- ["git", "fetch", "origin", llvm_hash],
- cwd=str(llvm_project), capture_output=True, text=True, timeout=2000,
- )
- if fetch_proc.returncode == 0:
- break
- print_warn(f"git fetch attempt {attempt}/6 failed: "
- f"{fetch_proc.stderr.strip()[-150:]}")
- else:
- raise RuntimeError(
- f"Failed to fetch LLVM commit {llvm_hash[:12]} after 6 attempts")
-
- subprocess.run(
- ["git", "checkout", llvm_hash],
- cwd=str(llvm_project), check=True, capture_output=True, text=True,
- timeout=2000,
- )
- print_status(True, f"Checked out {llvm_hash[:12]}")
- except Exception as e:
- print_error(f"Failed to checkout LLVM commit: {e}")
- log_proc = subprocess.run(
- ["git", "log", "--oneline", "-5"],
- cwd=str(llvm_project), capture_output=True, text=True, timeout=10,
- )
- print_info(f"llvm-project HEAD and recent commits:\n{log_proc.stdout.strip()}")
- return False
-
- # ── 3. Apply Ascend backend LLVM patch ──
- ascend_patch = ascend_path / "third_party" / "ascend" / "patch" / "llvm_patch_f6ded0b.patch"
- if ascend_patch.exists():
- print_info(f"Applying Ascend LLVM patch: {ascend_patch.name}")
- # Dry-run first
- dry_run = subprocess.run(
- ["git", "apply", "--check", str(ascend_patch)],
- cwd=str(llvm_project), capture_output=True, text=True, timeout=30,
- )
- if dry_run.returncode != 0:
- print_error(f"Patch does not apply cleanly: {dry_run.stderr.strip()[-400:]}")
- return False
- try:
- subprocess.run(
- ["git", "apply", str(ascend_patch)],
- cwd=str(llvm_project), check=True, capture_output=True, text=True, timeout=30,
- )
- print_status(True, "Ascend LLVM patch applied")
- except Exception as e:
- print_error(f"Failed to apply patch: {e}")
- return False
- else:
- print_warn(f"Ascend LLVM patch not found at {ascend_patch} — continuing without it")
-
- # ── 4. Build LLVM ──
- llvm_build_log = WORKSPACE_DIR / "llvm_build_baseline.log"
- llvm_build_log.parent.mkdir(parents=True, exist_ok=True)
-
- build_dir = llvm_project / "build"
- if build_dir.exists():
- import shutil
- shutil.rmtree(build_dir)
- build_dir.mkdir()
-
- cmake_cmd = [
- "cmake", str(llvm_project / "llvm"),
- "-G", "Ninja",
- "-DCMAKE_BUILD_TYPE=Release",
- "-DLLVM_ENABLE_ASSERTIONS=ON",
- "-DLLVM_ENABLE_PROJECTS=mlir;llvm;lld",
- "-DLLVM_TARGETS_TO_BUILD=host;NVPTX;AMDGPU",
- f"-DCMAKE_INSTALL_PREFIX={llvm_install}",
- "-DCMAKE_C_COMPILER=clang",
- "-DCMAKE_CXX_COMPILER=clang++",
- ]
-
- # ── Helper: run a command with live output streaming ──
- def _stream_cmd(cmd: list[str], cwd: Path, log_fh, timeout: int,
- label: str) -> int:
- """Stream subprocess output line-by-line to console and log file.
- Returns the process exit code."""
- print_info(f"{label} (streaming to {llvm_build_log.name})...")
- proc = subprocess.Popen(
- cmd, cwd=str(cwd),
- stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True,
- )
- assert proc.stdout is not None
- last_line = ""
- for line in proc.stdout:
- log_fh.write(line)
- stripped = line.rstrip()
- if stripped:
- last_line = stripped
- # \r returns to line start, \033[K clears trailing residue
- print(f"\r {stripped[:140]}\033[K", end="", flush=True)
- proc.wait(timeout=timeout)
- if last_line:
- print() # final newline after \r lines
- return proc.returncode
-
- # ── cmake configure ──
- with llvm_build_log.open("w", encoding="utf-8") as fh:
- fh.write(f"=== cmake ===\n{' '.join(cmake_cmd)}\n\n")
- fh.flush()
- rc = _stream_cmd(cmake_cmd, build_dir, fh, timeout=300,
- label="Configuring LLVM with cmake")
- if rc != 0:
- print_error(f"cmake failed (exit {rc}) — see {llvm_build_log}")
- return False
- print_status(True, "cmake configure OK")
-
- # ── ninja build + install ──
- print_info("Building LLVM with ninja (this may take ~0.5 hours)...")
- with llvm_build_log.open("a", encoding="utf-8") as fh:
- fh.write(f"\n=== ninja install ===\n")
- fh.flush()
- rc = _stream_cmd(["ninja", "install"], build_dir, fh, timeout=7200,
- label="ninja install")
- if rc != 0:
- print_error(f"ninja install failed (exit {rc}) — see {llvm_build_log}")
- return False
- print_status(True, "ninja install OK")
-
- # Copy FileCheck
- import shutil
- filecheck_src = build_dir / "bin" / "FileCheck"
- filecheck_dst = llvm_install / "bin" / "FileCheck"
- if filecheck_src.exists():
- filecheck_dst.parent.mkdir(parents=True, exist_ok=True)
- shutil.copy2(filecheck_src, filecheck_dst)
- print_info("Copied FileCheck to install prefix")
-
- # Write hash cache
- hash_cache = llvm_install / ".llvm_hash"
- llvm_install.mkdir(parents=True, exist_ok=True)
- hash_cache.write_text(llvm_hash, encoding="utf-8")
- print_status(True, "Baseline LLVM build complete")
-
- # Store llvm_prefix for later use
- if not self.state.llvm_prefix:
- self.state.llvm_prefix = str(llvm_install)
-
- # ── 5. Stash + drop the patch to leave a clean tree ──
- print_info("Stashing and dropping Ascend LLVM patch to clean working tree...")
- try:
- subprocess.run(
- ["git", "stash", "push", "-u", "-m", "ta-baseline-llvm-patch"],
- cwd=str(llvm_project), capture_output=True, text=True, timeout=30,
- )
- subprocess.run(
- ["git", "stash", "drop", "stash@{0}"],
- cwd=str(llvm_project), capture_output=True, text=True, timeout=30,
- )
- print_status(True, "LLVM working tree clean (patch stashed + dropped)")
- except Exception as e:
- print_warn(f"Stash/drop failed: {e} — forcing clean with checkout")
- subprocess.run(
- ["git", "checkout", "--", "."],
- cwd=str(llvm_project), capture_output=True, text=True, timeout=60,
- )
- subprocess.run(
- ["git", "clean", "-fd"],
- cwd=str(llvm_project), capture_output=True, text=True, timeout=60,
- )
-
- self.state.summary_rows.append(
- ("Baseline LLVM", "PASS", f"Built {llvm_hash[:12]}"))
- return True
-
- def _ensure_llvm_workspace_clean(self, reason: str = "") -> bool:
- """Ensure the llvm-project working tree is clean before building.
-
- Checks git status; if dirty, stashes and drops all uncommitted
- changes (including untracked files). Falls back to 'git checkout
- -- .' + 'git clean -fd' if stash fails.
-
- Returns True if the workspace is clean (or was cleaned successfully).
- """
- llvm_project = _llvm_project_path()
- if not llvm_project.exists():
- print_warn("[llvm-clean] llvm-project not found — cannot verify workspace")
- return True # nothing to clean
-
- # ── Check if working tree is dirty ──
- try:
- status = subprocess.run(
- ["git", "status", "--porcelain"],
- cwd=str(llvm_project),
- capture_output=True, text=True, timeout=15,
- ).stdout.strip()
- except Exception as e:
- print_warn(f"[llvm-clean] Could not check git status: {e}")
- return True # proceed and let the build step surface errors
-
- if not status:
- print_info(f"[llvm-clean] llvm-project workspace is clean"
- f"{f' ({reason})' if reason else ''}")
- return True
-
- # ── Workspace is dirty — clean it ──
- dirty_files = status.splitlines()
- print_warn(f"[llvm-clean] llvm-project has {len(dirty_files)} uncommitted"
- f" file(s){f' ({reason})' if reason else ''} — cleaning...")
- for f in dirty_files[:10]:
- print(f" {f}")
- if len(dirty_files) > 10:
- print(f" ... and {len(dirty_files) - 10} more")
-
- try:
- subprocess.run(
- ["git", "stash", "push", "-u", "-m",
- f"ta-auto-clean{': ' + reason if reason else ''}"],
- cwd=str(llvm_project),
- capture_output=True, text=True, timeout=30,
- )
- subprocess.run(
- ["git", "stash", "drop", "stash@{0}"],
- cwd=str(llvm_project),
- capture_output=True, text=True, timeout=30,
- )
- print_status(True, "[llvm-clean] Workspace cleaned (stash + drop)")
- return True
- except Exception as e:
- print_warn(f"[llvm-clean] Stash/drop failed: {e} — "
- f"forcing clean with checkout")
- try:
- subprocess.run(
- ["git", "checkout", "--", "."],
- cwd=str(llvm_project),
- capture_output=True, text=True, timeout=60,
- )
- subprocess.run(
- ["git", "clean", "-fd"],
- cwd=str(llvm_project),
- capture_output=True, text=True, timeout=60,
- )
- print_status(True, "[llvm-clean] Workspace cleaned (checkout + clean)")
- return True
- except Exception as e2:
- print_error(f"[llvm-clean] Failed to clean workspace: {e2}")
- return False
-
- def _stash_and_drop_llvm_patch(self) -> None:
- """Deprecated: use _ensure_llvm_workspace_clean() instead."""
- self._ensure_llvm_workspace_clean(reason="ir-patch-failed")
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Per-step test + fix loop (single-step mode)
- # ═══════════════════════════════════════════════════════════════════════════
-
- def _do_test_and_fix_loop(self) -> bool:
- """Run tests + AI-fix loop for the current step.
-
- Returns True if all tests pass, False on exhaustion.
- """
- ascend_path = Path(self.state.triton_ascend_path)
- step = (self.state.steps[self.state.current_step]
- if self.state.steps else None)
- step_id = step["id"] if step else "step-0"
- step_dir = WORKSPACE_DIR / STEPS_DIR / step_id
- step_dir.mkdir(parents=True, exist_ok=True)
-
- test_passed = False
-
- for attempt in range(self.state.max_retries + 1):
- is_fix_attempt = attempt > 0
- self.state.retry_count = attempt
-
- # AI fix test failures (skip on first round)
- if is_fix_attempt:
- print_header(f"Fix Attempt {attempt}/{self.state.max_retries} (test)")
- self.state.fix_errors = self._collect_test_error_logs()
- if self.state.fix_errors:
- ai_ok = self._do_ai_fix(ascend_path, step_dir, attempt)
- # Record fix attempt
- modified_files: list[str] = []
- ai_summary = ""
- if hasattr(self, '_last_ai_result') and self._last_ai_result:
- modified_files = self._last_ai_result.get("modified_files", [])
- ai_summary = self._last_ai_result.get("step_summary", "")
- error_snippet = ""
- for err_path in self.state.fix_errors:
- try:
- content = Path(err_path).read_text(
- encoding="utf-8", errors="replace")
- error_snippet += (content[-2000:]
- if len(content) > 2000 else content)
- except Exception:
- pass
- self.state.fix_attempts.append({
- "step_id": step_id,
- "attempt": attempt,
- "fix_type": "test",
- "error_logs": list(self.state.fix_errors),
- "error_snippet": error_snippet[-1500:],
- "modified_files": modified_files,
- "ai_summary": (ai_summary or "")[:2000],
- "ai_ok": ai_ok,
- })
- self.state.test_fix_count += 1
- else:
- print_warn("No test error logs found — cannot fix")
-
- # Rebuild after fix (skip on first attempt since build_and_fix already built)
- if is_fix_attempt:
- if not self._do_build(ascend_path, clean=False):
- print_warn(f"Build failed after test fix (attempt {attempt})")
- continue
-
- # Run tests
- test_result = self._do_test(ascend_path)
- if test_result is None:
- # SKIP_E2E_TEST — treat as pass
- test_passed = True
- break
- if test_result:
- test_passed = True
- break
-
- print_warn(f"Tests failed (attempt {attempt + 1}/"
- f"{self.state.max_retries + 1})")
-
- if os.getenv("SKIP_AI_ANALYSIS", "false").lower() == "true":
- print_warn("SKIP_AI_ANALYSIS=true — stopping test fix loop")
- break
-
- if test_passed:
- self.state.test_passed = True
- # Commit test fixes if any were applied
- if self.state.retry_count > 0:
- self._commit_fixes(ascend_path, step_dir)
- self.state.summary_rows.append(
- ("Tests", "PASS", f"{step_id}"))
- else:
- print_error(f"All {self.state.max_retries} fix attempts exhausted "
- f"— tests still failing")
- self.state.summary_rows.append(
- ("Tests", "FAIL", f"After {self.state.max_retries} attempts"))
- self.state.test_passed = False
-
- return test_passed
-
- def _collect_test_error_logs(self) -> list[str]:
- """Collect test failure log paths for AI fix context.
-
- Returns a list of file paths pointing to test logs and test result
- files in the workspace.
- """
- error_logs: list[str] = []
-
- # Test log directory — includes raw logs and JUnit XML reports
- test_log_dir = WORKSPACE_DIR / "test-logs"
- if test_log_dir.exists():
- for log_file in sorted(test_log_dir.rglob("*.log")):
- error_logs.append(str(log_file))
- for xml_file in sorted(test_log_dir.rglob("*.xml")):
- error_logs.append(str(xml_file))
-
- # Test result JSON
- test_result_path = WORKSPACE_DIR / TEST_RESULT_FILE
- if test_result_path.exists():
- error_logs.append(str(test_result_path))
-
- # Build result JSON (may contain build errors that affect tests)
- build_result_path = WORKSPACE_DIR / BUILD_RESULT_FILE
- if build_result_path.exists():
- error_logs.append(str(build_result_path))
-
- if error_logs:
- print_info(f"Collected {len(error_logs)} error log(s) for AI fix")
- for p in error_logs[:5]:
- print_info(f" - {p}")
- if len(error_logs) > 5:
- print_info(f" ... and {len(error_logs) - 5} more")
-
- return error_logs
-
- def _backup_code_state(self, label: str = "snapshot") -> Path | None:
- """Backup triton-ascend working tree to workspace for CI artifact retention.
-
- Copies the entire working tree (tracked + untracked) excluding .git
- and build artifacts. Used both on success (label="final") and on
- failure (label="failed-step-N") so no AI fix or conflict resolution
- work is ever lost.
- """
- ascend_path = Path(self.state.triton_ascend_path)
- ts = time.strftime("%Y%m%d-%H%M%S")
- backup_dir = WORKSPACE_DIR / "code-backups" / f"{label}_{ts}"
- backup_dir.parent.mkdir(parents=True, exist_ok=True)
-
- _ignore_patterns = shutil.ignore_patterns(
- ".git", "__pycache__", "*.pyc", "*.pyo",
- "*.o", "*.a", "*.so", "*.dylib",
- "build", "dist", "*.egg-info",
- ".mypy_cache", ".pytest_cache", ".ruff_cache",
- "result_profiling", "*.lock",
- )
- try:
- shutil.copytree(str(ascend_path), str(backup_dir),
- ignore=_ignore_patterns, symlinks=False)
- file_count = sum(1 for _ in backup_dir.rglob("*") if _.is_file())
- print_info(f"Code backup [{label}]: {backup_dir} ({file_count} files)")
-
- # ── Also record git state snapshot ──
- try:
- head = run_git(ascend_path, "rev-parse", "HEAD").strip()
- branch = run_git(ascend_path, "branch", "--show-current").strip()
- status = run_git(ascend_path, "status", "--porcelain").strip()
- info = (
- f"# Backup: {label}\n"
- f"# Time: {ts}\n"
- f"# Branch: {branch}\n"
- f"# HEAD: {head}\n"
- f"# Uncommitted changes: {'yes' if status else 'none'}\n"
- )
- (backup_dir / "_BACKUP_INFO.txt").write_text(info, encoding="utf-8")
- except Exception:
- pass
-
- return backup_dir
- except Exception as e:
- print_warn(f"Could not create code backup [{label}]: {e}")
- return None
-
- def _do_finalize(self):
- """Generate patch, summary & print final report.
-
- Does NOT restore the original branch — the work branch must stay
- checked out so push_to_github can push it. Branch restore happens
- at the end of push_to_github (or handle_failure).
- """
- print_header("Phase Final: Finalize & Summary")
-
- self._print_workspace_info("Phase Final: Finalize")
-
- ascend_path = Path(self.state.triton_ascend_path)
-
- # ── Generate final summary ──
- print_section("Generate Final Summary")
- final_summary_path = WORKSPACE_DIR / FINAL_SUMMARY_FILE
-
- # Collect step summaries if available
- steps_dir = WORKSPACE_DIR / STEPS_DIR
- if self.state.total_steps > 1 and steps_dir.exists():
- summaries = []
- for step in self.state.steps:
- step_dir = steps_dir / step["id"]
- summary_file = step_dir / EACH_STEP_SUMMARY_FILE
- if summary_file.exists():
- summaries.append(
- f"## {step['id']}\n\n"
- f"{summary_file.read_text(encoding='utf-8').strip()}"
- )
- if summaries:
- final_summary_path.write_text("\n\n".join(summaries), encoding="utf-8")
- else:
- final_summary_path.write_text(
- f"# Triton-Ascend Upstream Sync\n\n"
- f"- **Target**: `{self.state.target_commit[:12]}`\n"
- f"- **Steps**: {self.state.total_steps}\n"
- f"- **Work branch**: `{self.state.work_branch}`\n"
- f"- **Status**: Success\n"
- f"- **Date**: {time.strftime('%Y-%m-%d %H:%M:%S')}\n",
- encoding="utf-8",
- )
- else:
- step_dir = WORKSPACE_DIR / "step-0"
- last_summary_path = step_dir / EACH_STEP_SUMMARY_FILE
- if last_summary_path.exists():
- shutil.copy2(last_summary_path, final_summary_path)
- else:
- final_summary_path.write_text(
- f"# Triton-Ascend Upstream Sync\n\n"
- f"- **Target**: `{self.state.target_commit[:12]}`\n"
- f"- **Work branch**: `{self.state.work_branch}`\n"
- f"- **Status**: Success\n"
- f"- **Upstream commits merged**: {self.state.upstream_commits_count}\n"
- f"- **Date**: {time.strftime('%Y-%m-%d %H:%M:%S')}\n",
- encoding="utf-8",
- )
-
- print_info(f"Final summary: {final_summary_path}")
-
- # ── Generate cumulative patch (from original ascend HEAD to latest) ──
- try:
- patch = run_git(ascend_path, "diff", self.state.ascend_head, "HEAD")
- patch_path = WORKSPACE_DIR / FINAL_TARGET_PATCH_FILE
- patch_path.write_text(patch, encoding="utf-8")
- print_info(f"Cumulative patch: {patch_path} ({len(patch)} bytes)")
- except Exception as e:
- print_warn(f"Could not generate final patch: {e}")
-
- # ── Backup work branch code ──
- self._backup_code_state("final")
-
- self.state.summary_rows.append(
- ("Finalize", "PASS", f"{self.state.total_steps} step(s) completed")
- )
-
- # ── Print final summary table ──
- print_header("Sync Complete — Success!")
- print_elapsed_total()
- # Add IR loop metrics if applicable
- if self.state.llvm_hash_changed:
- self.state.summary_rows.append(
- ("IR Loop", "PASS",
- f"{len(self.state.ir_loop_details)} iteration(s)"))
- # Add pytest result
- pytest_status = "PASS" if self.state.pytest_passed else "N/A"
- self.state.summary_rows.append(("Pytest", pytest_status, ""))
- self.state.summary_rows.append(("OVERALL", "PASS", f"{self.state.total_steps} step(s) completed"))
- print_summary_table(self.state.summary_rows)
-
- # ── Generate sync report ──
- self._write_sync_report()
-
- print_section("Output Files")
- for f in sorted(WORKSPACE_DIR.rglob("*")):
- if f.is_file() and ".git" not in str(f):
- print(f" {f.relative_to(WORKSPACE_DIR)}")
- print_info(f"Work branch preserved: {self.state.work_branch}")
- print_info(f"To inspect: cd {ascend_path} && git checkout {self.state.work_branch}")
-
- def _write_sync_report(self) -> None:
- """Generate SYNC_REPORT.md via AI — let Claude Code write the report.
-
- Collects all sync data (fix attempt details, step summaries, error logs,
- modified files) into a context file, then calls the AI backend to produce
- a comprehensive, human-readable sync report.
- """
- report_path = WORKSPACE_DIR / "SYNC_REPORT.md"
-
- # ── Collect context for AI ──
- context = self._build_report_context()
- context_path = WORKSPACE_DIR / "report-context.json"
- context_path.write_text(
- json.dumps(context, indent=2, ensure_ascii=False), encoding="utf-8"
- )
- print_info(f"Report context written to {context_path}")
-
- # ── Build report prompt ──
- prompt = self._build_report_prompt(context)
-
- # ── Call AI backend to generate the report ──
- try:
- from TA_main2main_workflow.agent.opencode_adapter import (
- _detect_backend,
- )
- backend = _detect_backend()
- print_info(f"AI backend for report: {backend}")
-
- # Write prompt file for debugging
- prompt_path = WORKSPACE_DIR / "report-prompt.txt"
- prompt_path.write_text(prompt, encoding="utf-8")
-
- print_header("AI Report Generation")
- print_info("Calling AI backend to generate sync report...")
-
- ai_result = run_opencode_adapter({
- "step_id": "sync-report",
- "previous_step_id": "",
- "previous_step_summary_path": "",
- "is_last_step": "true",
- "step_index": "final",
- "step_dir": str(WORKSPACE_DIR),
- "fix_dir": str(WORKSPACE_DIR / "report-fix"),
- "conflict_dir": "",
- "ascend_path": self.state.triton_ascend_path,
- "triton_path": self.state.triton_path,
- "reference_dir": _REFERENCE_DIR,
- "mode": "report",
- "error_logs": json.dumps([str(context_path)], ensure_ascii=False),
- "target_commit": self.state.target_commit,
- })
-
- # AI writes report to step_dir/step_summary.md; we read it from there.
- # (ai_result return value is not used directly — report is file-based.)
- _ = ai_result # suppress unused-var warning
- ai_report_path = WORKSPACE_DIR / EACH_STEP_SUMMARY_FILE
- if ai_report_path.exists():
- report_content = ai_report_path.read_text(encoding="utf-8")
- # Add metadata header
- header = (
- f"# Triton-Ascend Upstream Sync Report\n\n"
- f"- **Date**: {time.strftime('%Y-%m-%d %H:%M:%S')}\n"
- f"- **Target commit**: `{self.state.target_commit[:12]}`\n"
- f"- **Work branch**: `{self.state.work_branch}`\n"
- f"- **Upstream commits**: {self.state.upstream_commits_count}\n"
- f"- **Steps**: {self.state.total_steps}\n"
- f"- **Merge conflicts resolved**: {self.state.conflict_files_resolved}\n"
- f"- **Build errors fixed**: {sum(s['build_fixes'] for s in self.state.step_details)}\n"
- f"- **Test failures fixed**: {sum(s['test_fixes'] for s in self.state.step_details)}\n"
- f"- **Total AI fix rounds**: {sum(s['retries'] for s in self.state.step_details)}\n\n"
- f"---\n\n"
- )
- report_path.write_text(header + report_content, encoding="utf-8")
- print_status(True, f"AI-generated sync report: {report_path}")
- else:
- print_warn("AI did not produce a report — using fallback")
- self._write_sync_report_fallback()
- except Exception as e:
- print_error(f"AI report generation failed: {e}")
- print_info("Using fallback report generator...")
- self._write_sync_report_fallback()
-
- def _build_report_context(self) -> dict:
- """Collect all sync data into a structured context for AI report generation."""
- total_build_fixes = sum(s["build_fixes"] for s in self.state.step_details)
- total_test_fixes = sum(s["test_fixes"] for s in self.state.step_details)
- total_retries = sum(s["retries"] for s in self.state.step_details)
-
- # Collect step AI summaries
- step_summaries: dict[str, str] = {}
- steps_dir = WORKSPACE_DIR / STEPS_DIR
- if steps_dir.exists():
- for step in self.state.steps:
- step_dir = steps_dir / step["id"]
- parts = []
- for fname in ["analysis.md", "step_summary.md", "review.md"]:
- fp = step_dir / fname
- if fp.exists():
- parts.append(
- f"### {fname}\n\n"
- f"{fp.read_text(encoding='utf-8', errors='replace').strip()}"
- )
- if parts:
- step_summaries[step["id"]] = "\n\n".join(parts)
-
- return {
- "overview": {
- "date": time.strftime('%Y-%m-%d %H:%M:%S'),
- "target_commit": self.state.target_commit[:12],
- "work_branch": self.state.work_branch,
- "upstream_commits_count": self.state.upstream_commits_count,
- "total_steps": self.state.total_steps,
- "conflict_files_resolved": self.state.conflict_files_resolved,
- "build_fix_count": total_build_fixes,
- "test_fix_count": total_test_fixes,
- "total_retries": total_retries,
- },
- "step_details": self.state.step_details,
- "fix_attempts": self.state.fix_attempts,
- "step_ai_summaries": step_summaries,
- "step_pr_descriptions": self.state.step_pr_descriptions,
- }
-
- def _build_report_prompt(self, context: dict) -> str:
- """Build the AI prompt for generating the sync report."""
- summary_json = json.dumps(context, indent=2, ensure_ascii=False)
- return (
- "Generate a comprehensive sync report in Chinese (中文) based on "
- "the structured context below. The report should be written as "
- "step_summary.md in the output directory.\n\n"
- "The report MUST include:\n\n"
- "## 1. Executive Summary\n"
- "- Brief overview of this sync (how many upstream commits, "
- "how many steps, overall outcome)\n"
- "- Key metrics (conflicts resolved, build errors fixed, "
- "test failures fixed, AI fix rounds)\n\n"
- "## 2. Per-Step Analysis\n"
- "- For each step, explain:\n"
- " - Which upstream commits were merged and what areas they touched\n"
- " - What merge conflicts arose and how they were resolved\n"
- " - What build errors occurred, root causes, and how AI fixed them\n"
- " - What test failures occurred, root causes, and how AI fixed them\n"
- "- Include specific file paths and error messages where relevant\n\n"
- "## 3. Fix Pattern Analysis\n"
- "- Identify recurring patterns across fixes (e.g., API changes, "
- "missing includes, signature mismatches)\n"
- "- Highlight any fixes that required multiple attempts\n\n"
- "## 4. Recommendations\n"
- "- Suggest preventative measures for future syncs\n"
- "- Flag any areas of the codebase that are particularly fragile\n\n"
- "Rules:\n"
- "- Write in Chinese (中文)\n"
- "- Be specific — include file paths, error messages, commit ranges\n"
- "- Write the output to {step_dir}/step_summary.md\n"
- "- DO NOT modify any source code — this is a report-only task\n\n"
- f"CONTEXT DATA:\n\n{summary_json}"
- )
-
- def _write_sync_report_fallback(self) -> None:
- """Fallback: assemble report from template (no AI)."""
- report_path = WORKSPACE_DIR / "SYNC_REPORT.md"
- L: list[str] = []
-
- total_build_fixes = sum(s["build_fixes"] for s in self.state.step_details)
- total_test_fixes = sum(s["test_fixes"] for s in self.state.step_details)
- total_retries = sum(s["retries"] for s in self.state.step_details)
-
- L.append("# Triton-Ascend Upstream Sync Report\n")
- L.append(f"**Date**: {time.strftime('%Y-%m-%d %H:%M:%S')}")
- L.append(f"**Target commit**: `{self.state.target_commit[:12]}`")
- L.append(f"**Work branch**: `{self.state.work_branch}`")
- L.append(f"**Status**: Success\n")
-
- L.append("## Summary\n")
- L.append("| Metric | Count |")
- L.append("|--------|-------|")
- L.append(f"| Upstream commits synced | {self.state.upstream_commits_count} |")
- L.append(f"| Steps | {self.state.total_steps} |")
- L.append(f"| Merge conflicts resolved | {self.state.conflict_files_resolved} |")
- L.append(f"| Build errors fixed | {total_build_fixes} |")
- L.append(f"| Test failures fixed | {total_test_fixes} |")
- L.append(f"| AI fix rounds | {total_retries} |")
-
- if self.state.step_details:
- L.append("\n## Per-Step Breakdown\n")
- L.append("| Step | Commits | Lines | Conflicts | Build Fixes | Test Fixes | Retries |")
- L.append("|------|---------|-------|-----------|-------------|------------|---------|")
- for s in self.state.step_details:
- L.append(
- f"| {s['step_id']} ({s['step_index']}/{self.state.total_steps}) "
- f"| {s['commits']} | {s['source_lines']} | {s['conflict_files']} "
- f"| {s['build_fixes']} | {s['test_fixes']} | {s['retries']} |"
- )
-
- for fa in self.state.fix_attempts:
- ftype = fa["fix_type"].upper()
- L.append(
- f"\n### {fa['step_id']} — Fix {fa['attempt']} ({ftype})\n"
- )
- if fa["modified_files"]:
- L.append(f"**Files**: {', '.join(f'`{f}`' for f in fa['modified_files'])}")
- ai_sum = fa.get("ai_summary", "").strip()
- if ai_sum:
- L.append(f"\n{ai_sum}")
-
- steps_dir = WORKSPACE_DIR / STEPS_DIR
- if steps_dir.exists():
- for step in self.state.steps:
- step_dir = steps_dir / step["id"]
- for fname in ["analysis.md", "step_summary.md", "review.md"]:
- fp = step_dir / fname
- if fp.exists():
- L.append(
- f"\n### {step['id']} — {fname}\n\n"
- f"{fp.read_text(encoding='utf-8', errors='replace').strip()}\n"
- )
-
- L.append(f"\n---\n🤖 Generated at {time.strftime('%Y-%m-%d %H:%M:%S')}\n")
- report_path.write_text("\n".join(L), encoding="utf-8")
- print_info(f"Fallback sync report: {report_path}")
-
- # ═══════════════════════════════════════════════════════════════════════════
- # Terminal nodes (routed from execute_sync)
- # ═══════════════════════════════════════════════════════════════════════════
-
- @listen(UpgradeCompleted)
- def push_to_github(self):
- """Push work branch & create a single GitHub PR after ALL steps complete.
-
- In the vllm-ascend step-by-step merge style, all step commits accumulate
- on the work branch locally. Only after every step passes (merge →
- resolve → build → test → fix → commit) do we push and open one PR.
- """
- if os.getenv("PUSH_TO_GITHUB", "false").lower() != "true":
- print_info("PUSH_TO_GITHUB is not 'true' — skipping PR creation")
- print_info("To push manually:")
- print_info(f" cd {self.state.triton_ascend_path}")
- print_info(f" git checkout {self.state.work_branch}")
- print_info(f" git push -u origin {self.state.work_branch}")
- self.state.summary_rows.append(("Push & PR", "SKIP", "PUSH_TO_GITHUB not set"))
- return "SKIP_PUSH"
-
- print_header("Push to GitHub & Create PR")
-
- self._print_workspace_info("Push to GitHub & Create PR")
-
- github_repo = os.getenv("GITHUB_REPO", "triton-lang/triton-ascend")
- if not github_repo:
- print_error("GITHUB_REPO is empty — cannot create PR")
- self.state.summary_rows.append(("Push & PR", "FAIL", "GITHUB_REPO empty"))
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
-
- # ── Build a comprehensive PR body from step summaries ──
- pr_body_path = WORKSPACE_DIR / FINAL_SUMMARY_FILE
- self._build_pr_body(pr_body_path)
-
- # ── Push AscendNPU-IR submodule first ──
- self._push_submodule_if_needed()
-
- try:
- pr_url = push_and_create_pr(
- ascend_path=Path(self.state.triton_ascend_path),
- github_repo=github_repo,
- work_branch=self.state.work_branch,
- summary_path=pr_body_path,
- target_commit=self.state.target_commit,
- )
- self.state.pr_url = pr_url
- print_status(True, f"PR created: {pr_url}")
- self.state.summary_rows.append(("Push & PR", "PASS", pr_url))
- except Exception as e:
- print_error(f"Failed to push/create PR: {e}")
- # ── Print detailed failure diagnostics ──
- if isinstance(e, subprocess.CalledProcessError):
- print_section("Push/PR Failure Details")
- print_key_value("Command", " ".join(e.cmd) if e.cmd else "N/A")
- print_key_value("Exit code", str(e.returncode))
- if e.stdout:
- print_info(f"stdout:\n{e.stdout.strip()}")
- if e.stderr:
- print_error(f"stderr:\n{e.stderr.strip()}")
- else:
- import traceback
- print_info(f"Traceback:\n{traceback.format_exc()}")
- # Print git context for debugging
- ascend_path = Path(self.state.triton_ascend_path)
- print_section("Git Context at Failure")
- print_key_value("Work branch", self.state.work_branch)
- try:
- current_branch = run_git(ascend_path, "branch", "--show-current").strip()
- print_key_value("Current branch", current_branch)
- status_out = run_git(ascend_path, "status", "--short").strip()
- print_info(f"Git status:\n{status_out}" if status_out else "Git status: (clean)")
- log_out = run_git(ascend_path, "log", "--oneline", "-5")
- print_info(f"Recent commits:\n{log_out.strip()}")
- except Exception:
- pass
- self.state.summary_rows.append(("Push & PR", "FAIL", str(e)[:60]))
- self.state.final_status = UpgradeFailed
- # Still try to restore branch, then signal failure
- self._restore_branch()
- return UpgradeFailed
-
- # ── Restore original branch after push ──
- self._restore_branch()
- return self.state.pr_url if self.state.pr_url else "SKIP_PUSH"
-
- def _restore_branch(self) -> None:
- """Restore the original branch after all work is done."""
- ascend_path = Path(self.state.triton_ascend_path)
- print_section("Restore Original Branch")
- try:
- current = run_git(ascend_path, "branch", "--show-current").strip()
- if current != self.state.original_branch:
- run_git(ascend_path, "checkout", self.state.original_branch)
- print_status(True, f"Restored to '{self.state.original_branch}'")
- else:
- print_info(f"Already on '{self.state.original_branch}'")
- except Exception as e:
- print_warn(f"Could not restore branch: {e}")
- print_info(f"Work branch '{self.state.work_branch}' left checked out")
-
- def _build_pr_body(self, output_path: Path) -> None:
- """Build a comprehensive PR body from all step descriptions and summaries."""
- parts: list[str] = []
-
- # Title / overview
- parts.append(
- "# Triton-Ascend Upstream Sync\n\n"
- f"- **Target commit**: `{self.state.target_commit[:12]}`\n"
- f"- **Work branch**: `{self.state.work_branch}`\n"
- f"- **Steps completed**: {self.state.total_steps}\n"
- f"- **Upstream commits merged**: {self.state.upstream_commits_count}\n"
- f"- **Date**: {time.strftime('%Y-%m-%d %H:%M:%S')}\n"
- )
-
- # Per-step progress
- if self.state.step_pr_descriptions:
- parts.append("## Step Progress\n")
- for desc in self.state.step_pr_descriptions:
- parts.append(f"- {desc}\n")
-
- # Per-step AI summaries (if available)
- steps_dir = WORKSPACE_DIR / STEPS_DIR
- if self.state.total_steps > 1 and steps_dir.exists():
- parts.append("\n## Step Details\n")
- for step in self.state.steps:
- step_dir = steps_dir / step["id"]
- summary_file = step_dir / EACH_STEP_SUMMARY_FILE
- if summary_file.exists():
- parts.append(
- f"### {step['id']}\n\n"
- f"{summary_file.read_text(encoding='utf-8').strip()}\n\n"
- )
- else:
- parts.append(
- f"### {step['id']}\n\n"
- f"- Commits: {step['commit_count']}\n"
- f"- End commit: `{step['end_commit'][:12]}`\n"
- f"- Source lines changed: {step.get('source_changed_lines', '?')}\n\n"
- )
- elif steps_dir.exists():
- # Single step: include its summary
- step_dir = WORKSPACE_DIR / "step-0"
- summary_file = step_dir / EACH_STEP_SUMMARY_FILE
- if summary_file.exists():
- parts.append(
- "\n## Summary\n\n"
- f"{summary_file.read_text(encoding='utf-8').strip()}\n"
- )
- else:
- # Fallback: just the final summary
- fallback = WORKSPACE_DIR / FINAL_SUMMARY_FILE
- if fallback.exists():
- parts.append(fallback.read_text(encoding='utf-8'))
-
- parts.append(
- f"\n---\n"
- f"🤖 Generated with [TA_main2main_workflow]"
- f"(https://github.com/TecJesh/TA-AI-WorkFlow)"
- f" at {time.strftime('%Y-%m-%d %H:%M:%S')}\n"
- )
-
- output_path.write_text("".join(parts), encoding="utf-8")
- print_info(f"PR body written to {output_path}")
-
- @listen(UpgradeFailed)
- def handle_failure(self):
- """write FAILURE.md, print diagnostics & summary, suggest recovery commands."""
- print_header("Sync Failed — Diagnostics")
-
- self._print_workspace_info("Handle Failure")
-
- ascend_path = Path(self.state.triton_ascend_path)
-
- # ── Backup code state BEFORE anything else ──
- # Capture the working tree so AI fixes, conflict resolutions, and
- # partial merge progress are preserved as CI artifacts even on failure.
- failed_step = self.state.current_step + 1 if self.state.current_step < self.state.total_steps else self.state.total_steps
- self._backup_code_state(f"failed-step{failed_step}")
-
- print_error(f"Upgrade failed after {self.state.retry_count} retries")
-
- print_section("Failure Details")
- print_key_value("Target commit", self.state.target_commit[:12])
- print_key_value("Work branch", self.state.work_branch)
- print_key_value("Original branch", self.state.original_branch)
- print_key_value("Conflict files", ", ".join(self.state.conflict_files) if self.state.conflict_files else "none")
- print_key_value("Build passed", str(self.state.build_passed))
- print_key_value("Test passed", str(self.state.test_passed))
-
- failure_path = WORKSPACE_DIR / "FAILURE.md"
- failure_text = (
- f"# Upgrade Failed\n\n"
- f"- **Target**: `{self.state.target_commit[:12]}`\n"
- f"- **Work branch**: `{self.state.work_branch}`\n"
- f"- **Original branch**: `{self.state.original_branch}`\n"
- f"- **Retries**: {self.state.retry_count}/{self.state.max_retries}\n"
- f"- **Conflict files**: {', '.join(self.state.conflict_files) if self.state.conflict_files else 'none'}\n"
- f"- **Build passed**: {self.state.build_passed}\n"
- f"- **Test passed**: {self.state.test_passed}\n"
- f"- **Date**: {time.strftime('%Y-%m-%d %H:%M:%S')}\n\n"
- f"## Recovery\n\n"
- f"```bash\n"
- f"cd {ascend_path}\n"
- f"git checkout {self.state.original_branch}\n"
- f"# Work branch '{self.state.work_branch}' has the partial merge\n"
- f"# git branch -D {self.state.work_branch}\n"
- f"```\n"
- )
- failure_path.write_text(failure_text, encoding="utf-8")
- print_info(f"Failure report: {failure_path}")
-
- print_elapsed_total()
- self.state.summary_rows.append(("OVERALL", "FAIL", f"Failed after {self.state.retry_count} retries"))
- print_summary_table(self.state.summary_rows)
-
- print_section("Recovery")
- print_info(f"Work branch '{self.state.work_branch}' preserved for manual inspection")
- print_info(f"To restore: cd {ascend_path} && git checkout {self.state.original_branch}")
- print_info(f"To clean up: cd {ascend_path} && git branch -D {self.state.work_branch}")
-
- self.state.final_status = UpgradeFailed
- return UpgradeFailed
+ log.error(f"Failed to create PR: {e}")
diff --git a/src/TA_main2main_workflow/main.py b/src/TA_main2main_workflow/main.py
index d735cf2..80ca79a 100644
--- a/src/TA_main2main_workflow/main.py
+++ b/src/TA_main2main_workflow/main.py
@@ -1,32 +1,34 @@
#!/usr/bin/env python3
"""CLI entrypoint for TA_main2main_workflow — Triton-Ascend upstream sync.
-Commands:
- ta-kickoff Run the main2main sync flow (all output printed locally)
- ta-plot Generate a flow diagram (HTML)
+Single-step mode is the only supported mode. Each step runs the full
+pipeline: merge → resolve conflicts → build → fix → test → fix → commit.
Environment variables:
TRITON_ASCEND_PATH — path to triton-ascend repo (default: cwd)
TRITON_PATH — path to upstream triton repo (default: uses remote)
TRITON_TARGET_COMMIT — specific upstream commit to sync to (default: HEAD)
AI_BACKEND — "opencode" or "claude" (default: auto-detect)
- SKIP_AI_ANALYSIS — set to "true" to skip AI (NOT recommended)
+ SKIP_AI_ANALYSIS — set to "true" to skip AI calls
SKIP_BUILD — set to "true" to skip build step
SKIP_E2E_TEST — set to "true" to skip test step
PUSH_TO_GITHUB — set to "true" to auto-create PR after success
GITHUB_REPO — "owner/repo" for PR creation
LLVM_INSTALL_PREFIX — path to LLVM for building
+ LLVM_PROJECT_PATH — path to llvm-project repo (default: ~/llvm-project)
+ LLVM_INSTALL_PREFIX_SYNC — path to LLVM install (default: ~/llvm-install-sync)
CONDA_ENV — conda env name (default: ta-upgrade)
- NUM_PROCS — number of parallel pytest workers (default: 16)
-
- TA_MODE — Execution mode:
- full (default) Complete flow: merge → build → test → fix → PR
- merge Merge + AI resolve only, then push work branch & exit.
- Used by CI: runs on ubuntu-latest, then triggers NPU tests.
- fix AI fix only on an existing work branch. Requires:
- TA_WORK_BRANCH — work branch name
- TA_ERROR_LOGS_PATH — path to test failure logs (optional)
- TA_FIX_ATTEMPT — retry attempt number (optional)
+ BUILD_PROCS — number of parallel build workers (default: 32)
+ TEST_PROCS — number of parallel pytest workers (default: 16)
+ TA_LINE_BUDGET — max source lines per merge step (default: 1000)
+ TA_MAX_RETRIES — max AI fix retries (default: 10)
+ TA_BASE_BRANCH — base branch name (default: upstream_sync)
+ TA_TEST_DIR — primary test directory (default: third_party/ascend/unittest/pytest_ut)
+ TA_EXTRA_TEST_DIRS — extra test directories, comma/space separated (default: none)
+ TA_TEST_COMMAND — additional custom test command (default: none; runs after pytest ut)
+ TA_SINGLE_STEP_MODE — set to "true" for single-step mode (default: true)
+ TA_WORK_BRANCH_BASE — remote name for work branch base (default: upstream-ascend)
+ TA_RESUME — set to "true" to resume from cached step outputs
"""
import argparse
@@ -36,83 +38,23 @@
from TA_main2main_workflow.flow import TA_Main2MainFlow
from TA_main2main_workflow.utils import UpgradeFailed
+from TA_main2main_workflow.utils.config import TAConfig, _resolve_test_dirs
+from TA_main2main_workflow.utils.logging import get_logger
-
-def _print_startup_banner() -> None:
- skip_ai = os.getenv("SKIP_AI_ANALYSIS", "false").lower() == "true"
- skip_build = os.getenv("SKIP_BUILD", "false").lower() == "true"
- skip_test = os.getenv("SKIP_E2E_TEST", "false").lower() == "true"
- ai_backend = os.getenv("AI_BACKEND", "auto-detect")
- mode = os.getenv("TA_MODE", "full")
-
- print(f"╔{'═' * 60}╗")
- print(f"║ TA_main2main_workflow — Triton-Ascend Upstream Sync ║")
- print(f"╠{'═' * 60}╣")
- print(f"║ Mode: {mode:<44}║")
- print(f"║ AI Backend: {ai_backend:<44}║")
- print(f"║ AI Enabled: {'YES' if not skip_ai else 'NO (SKIP_AI_ANALYSIS=true)':<44}║")
- print(f"║ Skip Build: {str(skip_build):<44}║")
- print(f"║ Skip Test: {str(skip_test):<44}║")
- print(f"╚{'═' * 60}╝")
-
- if skip_ai:
- print()
- print(" ⚠ WARNING: SKIP_AI_ANALYSIS=true")
- print(" ⚠ AI will NOT be called to resolve conflicts or fix failures!")
- print(" ⚠ You must resolve conflicts and fix test failures manually.")
- print()
-
-
-def _is_failed(result) -> bool:
- """Check whether a kickoff result indicates workflow failure.
-
- Handles both plain string returns (merge/fix modes) and CrewAI
- CrewOutput objects (full mode).
- """
- if result is None:
- return False
- if isinstance(result, str):
- return result == UpgradeFailed
- # CrewAI CrewOutput / object with raw attribute
- if hasattr(result, 'raw'):
- return str(result.raw) == UpgradeFailed
- # Last resort: string representation
- return str(result) == UpgradeFailed
+log = get_logger(__name__)
def kickoff():
parser = argparse.ArgumentParser(
- description="Triton-Ascend Main2Main Upstream Sync Flow"
- )
- parser.add_argument(
- "--mode", default=None,
- choices=["full", "merge", "fix"],
- help="Execution mode: full (default), merge (merge+resolve only), "
- "fix (AI fix on existing work branch). "
- "Can also be set via TA_MODE env var."
- )
- parser.add_argument(
- "--work-branch", default=None,
- help="Work branch name (required for --mode=fix). "
- "Can also be set via TA_WORK_BRANCH env var."
- )
- parser.add_argument(
- "--error-logs-path", default=None,
- help="Path to test failure logs for AI fix (--mode=fix). "
- "Can also be set via TA_ERROR_LOGS_PATH env var."
- )
- parser.add_argument(
- "--fix-attempt", type=int, default=None,
- help="Retry attempt number (--mode=fix). "
- "Can also be set via TA_FIX_ATTEMPT env var."
+ description="Triton-Ascend Main2Main Upstream Sync (Single-Step Mode)"
)
parser.add_argument(
"--triton-ascend-path", default=None,
- help="Local path to the triton-ascend repository (default: current directory)"
+ help="Local path to the triton-ascend repository (default: TRITON_ASCEND_PATH env)"
)
parser.add_argument(
"--triton-path", default=None,
- help="Local path to the upstream triton repository (default: uses remote)"
+ help="Local path to the upstream triton repository (default: TRITON_PATH env)"
)
parser.add_argument(
"--target-commit", default=None,
@@ -127,67 +69,73 @@ def kickoff():
help="Conda environment name (default: ta-upgrade)"
)
parser.add_argument(
- "--num-procs", type=int, default=None,
- help="Number of parallel pytest workers (default: 16)"
+ "--build-procs", type=int, default=None,
+ help="Parallel build workers (default: 32)"
+ )
+ parser.add_argument(
+ "--test-procs", type=int, default=None,
+ help="Parallel pytest workers (default: 8)"
+ )
+ parser.add_argument(
+ "--extra-test-dirs", default=None,
+ help="Extra test directories (comma/space separated, appended to default pytest ut)"
+ )
+ parser.add_argument(
+ "--test-command", default=None,
+ help="Additional custom test command (runs after pytest ut)"
)
args = parser.parse_args()
- # ── Mode: CLI arg takes precedence over env var ──
- if args.mode:
- os.environ["TA_MODE"] = args.mode
- if args.work_branch:
- os.environ["TA_WORK_BRANCH"] = args.work_branch
- if args.error_logs_path:
- os.environ["TA_ERROR_LOGS_PATH"] = args.error_logs_path
- if args.fix_attempt is not None:
- os.environ["TA_FIX_ATTEMPT"] = str(args.fix_attempt)
-
- _print_startup_banner()
-
- inputs = {}
+ config = TAConfig.from_env()
if args.triton_ascend_path:
- inputs["triton_ascend_path"] = args.triton_ascend_path
+ config.triton_ascend_path = args.triton_ascend_path
if args.triton_path:
- inputs["triton_path"] = args.triton_path
+ config.triton_path = args.triton_path
if args.target_commit:
- inputs["target_commit"] = args.target_commit
+ config.target_commit = args.target_commit
if args.llvm_prefix:
- inputs["llvm_prefix"] = args.llvm_prefix
+ config.llvm_install_prefix = args.llvm_prefix
if args.conda_env:
- inputs["conda_env"] = args.conda_env
- if args.num_procs:
- inputs["num_procs"] = args.num_procs
-
- flow = TA_Main2MainFlow()
+ config.conda_env = args.conda_env
+ if args.build_procs is not None:
+ config.build_procs = args.build_procs
+ if args.test_procs is not None:
+ config.test_procs = args.test_procs
+ if args.extra_test_dirs is not None:
+ config.test_dirs = _resolve_test_dirs(
+ primary=config.test_dir,
+ extra=args.extra_test_dirs,
+ )
+ if args.test_command is not None:
+ config.test_command = args.test_command
+
+ _print_banner(config)
+
+ flow = TA_Main2MainFlow(config=config)
try:
- result = flow.kickoff(inputs=inputs if inputs else None)
+ result = flow.run()
except Exception as exc:
- print(f"\n{'=' * 60}")
- print(f" WORKFLOW CRASHED: {exc}")
- print(f"{'=' * 60}")
+ log.error(f"WORKFLOW CRASHED: {exc}")
+ import traceback
+ traceback.print_exc()
sys.exit(1)
- if _is_failed(result):
- print(f"\n{'=' * 60}")
- print(f" WORKFLOW FAILED — exiting with code 1")
- print(f"{'=' * 60}")
+ if result == UpgradeFailed:
+ log.error("WORKFLOW FAILED")
sys.exit(1)
- print(f"\n{'=' * 60}")
- print(f" WORKFLOW COMPLETED SUCCESSFULLY")
- print(f"{'=' * 60}")
-
+ log.info("WORKFLOW COMPLETED SUCCESSFULLY")
-def plot():
- import shutil
- output_dir = Path(__file__).resolve().parent / "output"
- output_dir.mkdir(parents=True, exist_ok=True)
- flow = TA_Main2MainFlow()
- tmp_html = Path(flow.plot(filename="flow.html", show=False))
- for f in tmp_html.parent.iterdir():
- shutil.copy2(f, output_dir / f.name)
- print(f"Flow plot saved to: {output_dir / tmp_html.name}")
+def _print_banner(config: TAConfig) -> None:
+ ai = "NO (SKIP_AI_ANALYSIS=true)" if config.skip_ai_analysis else "YES"
+ log.header("TA_main2main_workflow — Triton-Ascend Upstream Sync")
+ log.key_value("AI Backend", config.ai_backend)
+ log.key_value("AI Enabled", ai)
+ log.key_value("Skip Build", str(config.skip_build))
+ log.key_value("Skip Test", str(config.skip_e2e_test))
+ if config.skip_ai_analysis:
+ log.warning("SKIP_AI_ANALYSIS=true — AI will not be called!")
if __name__ == "__main__":
diff --git a/src/TA_main2main_workflow/pipeline/__init__.py b/src/TA_main2main_workflow/pipeline/__init__.py
new file mode 100644
index 0000000..3a2b64a
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/__init__.py
@@ -0,0 +1,9 @@
+"""Pipeline step functions.
+
+Each step is an independent function with signature::
+
+ def step_xxx(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext
+
+Steps read from *ctx*, perform their work, and return an updated
+``WorkflowContext`` (never mutating the input).
+"""
diff --git a/src/TA_main2main_workflow/pipeline/build.py b/src/TA_main2main_workflow/pipeline/build.py
new file mode 100644
index 0000000..0914918
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/build.py
@@ -0,0 +1,378 @@
+"""Pipeline step: Build Triton-Ascend with AI fix loops.
+
+Main entry point for single-step mode:
+
+ ``build_and_fix_loop(ctx, config)`` — Build TA with AI fix compile
+ errors loop. Used when LLVM hash has NOT changed (LLVM was
+ already built by the baseline step).
+
+LLVM build helpers (``build_llvm``, ``llvm_setup``) are public so the
+IR patch pipeline in ``ir_patch.py`` can reuse them for LLVM version
+changes.
+"""
+
+from __future__ import annotations
+
+import json
+import os
+import subprocess
+from pathlib import Path
+
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.tracker import timed
+from TA_main2main_workflow.utils.git import run_git, stream_cmd
+from TA_main2main_workflow.utils import (
+ BUILD_RESULT_FILE, STEPS_DIR, WORKSPACE_DIR,
+)
+from TA_main2main_workflow.pipeline.fix import ai_fix
+from TA_main2main_workflow.pipeline.pre_ci import cleanup_temp_files
+from TA_main2main_workflow.utils.submodule import commit_submodule, submodule_has_changes
+
+log = get_logger(__name__)
+
+
+def build_and_fix_loop(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Build TA with AI fix loop for compile errors.
+
+ Used in single-step mode when LLVM hash has NOT changed.
+ Does NOT rebuild LLVM — assumes baseline LLVM is already built.
+ """
+ if config.skip_build:
+ log.info("SKIP_BUILD=true — skipping build")
+ return ctx.copy_with(build_passed=True)
+
+ attempt = 0
+ while attempt <= config.max_retries:
+ ctx = ctx.copy_with(retry_count=attempt)
+
+ if attempt > 0:
+ log.header(f"Build Fix Attempt {attempt}/{config.max_retries}")
+ ctx = ai_fix(ctx, config, attempt=attempt, mode="fix")
+
+ with timed("build-triton"):
+ ctx = build_triton(ctx, config, clean=(attempt == 0))
+
+ if ctx.build_passed:
+ # Commit fixes if any
+ if attempt > 0:
+ commit_fixes(ctx, config)
+ return ctx.copy_with(
+ build_passed=True,
+ build_fix_count=ctx.build_fix_count + (1 if attempt > 0 else 0),
+ )
+
+ log.info(f"Triton build failed (attempt {attempt + 1}) — retrying")
+ attempt += 1
+
+ return ctx.copy_with(build_passed=False)
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# LLVM
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def llvm_setup(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Clone LLVM, checkout hash, apply patch. Idempotent."""
+ ascend_path = Path(ctx.triton_ascend_path)
+ llvm_project = config.llvm_project
+ llvm_hash_file = ascend_path / "cmake" / "llvm-hash.txt"
+
+ if not llvm_hash_file.exists():
+ log.info("No llvm-hash.txt — skipping LLVM rebuild")
+ return ctx
+ required_hash = llvm_hash_file.read_text(encoding="utf-8").strip()
+ if not required_hash:
+ return ctx
+
+ if not llvm_project.exists():
+ run_git(WORKSPACE_DIR, "clone", config.llvm_repo_url, str(llvm_project))
+
+ log.section(f"LLVM setup (hash: {required_hash[:12]})")
+ run_git(llvm_project, "fetch", "origin", required_hash)
+ run_git(llvm_project, "reset", "--hard", "HEAD")
+ run_git(llvm_project, "clean", "-fd")
+ run_git(llvm_project, "checkout", "-f", required_hash)
+
+ patch_dir = ascend_path / "third_party/ascend/patch"
+ patch_files = sorted(patch_dir.glob("*.patch")) if patch_dir.exists() else []
+ if patch_files:
+ if required_hash[:7] in patch_files[0].name:
+ log.info(f"Applying patch: {patch_files[0].name}")
+ run_git(llvm_project, "apply", str(patch_files[0]))
+ else:
+ log.info(
+ f"LLVM hash changed ({required_hash[:12]}), "
+ f"patch {patch_files[0].name} is for old version"
+ )
+ return ctx
+
+
+def build_llvm(ctx: WorkflowContext, num_procs: int = 32) -> WorkflowContext:
+ """cmake + ninja for LLVM. Pure build — no retry logic."""
+ llvm_project = Path(os.path.expanduser(
+ os.getenv("LLVM_PROJECT_PATH", "~/llvm-project")))
+ llvm_install = Path(os.path.expanduser(
+ os.getenv("LLVM_INSTALL_PREFIX_SYNC", "~/llvm-install-sync")))
+ ascend_path = Path(ctx.triton_ascend_path)
+ required_hash = (
+ (ascend_path / "cmake" / "llvm-hash.txt").read_text(encoding="utf-8").strip()
+ )
+
+ step_id = ctx.steps[ctx.current_step]["id"] if ctx.steps else "step-0"
+ step_dir = WORKSPACE_DIR / STEPS_DIR / step_id
+ step_dir.mkdir(parents=True, exist_ok=True)
+
+ build_dir = WORKSPACE_DIR / "llvm-build"
+ if build_dir.exists():
+ import shutil
+ shutil.rmtree(build_dir)
+ build_dir.mkdir(parents=True, exist_ok=True)
+
+ log_dir = step_dir
+ log_dir.mkdir(parents=True, exist_ok=True)
+ llvm_build_log = log_dir / "llvm-build.log"
+
+ # ── cmake configure ──
+ cmake_cmd = [
+ "cmake", str(llvm_project / "llvm"),
+ "-G", "Ninja",
+ "-DCMAKE_BUILD_TYPE=Release",
+ "-DLLVM_ENABLE_ASSERTIONS=ON",
+ "-DLLVM_ENABLE_PROJECTS=mlir;llvm;lld",
+ "-DLLVM_TARGETS_TO_BUILD=host;NVPTX;AMDGPU",
+ f"-DCMAKE_INSTALL_PREFIX={llvm_install}",
+ "-DCMAKE_C_COMPILER=clang",
+ "-DCMAKE_CXX_COMPILER=clang++",
+ ]
+ with open(llvm_build_log, "w", encoding="utf-8") as fh:
+ fh.write(f"=== cmake ===\n{' '.join(cmake_cmd)}\n\n")
+ fh.flush()
+ rc = stream_cmd(cmake_cmd, build_dir, fh, timeout=300,
+ label="Configuring LLVM with cmake")
+ if rc != 0:
+ log.error(f"LLVM cmake FAILED — see {llvm_build_log}")
+ return ctx.copy_with(
+ build_passed=False, fix_errors=[str(llvm_build_log)]
+ )
+ log.status(True, "cmake configure OK")
+
+ # ── ninja build + install ──
+ log.info(f"ninja -j{num_procs} install (this may take a while)...")
+ with open(llvm_build_log, "a", encoding="utf-8") as fh:
+ fh.write(f"\n=== ninja install ===\n")
+ fh.flush()
+ rc = stream_cmd(
+ ["ninja", "-j", str(num_procs), "install"],
+ build_dir, fh, timeout=7200, label="ninja install",
+ )
+ if rc != 0:
+ log.error(f"LLVM ninja FAILED — see {llvm_build_log}")
+ return ctx.copy_with(
+ build_passed=False, fix_errors=[str(llvm_build_log)]
+ )
+ log.status(True, "LLVM ninja install OK")
+
+ # Copy FileCheck
+ import shutil
+ fc = build_dir / "bin" / "FileCheck"
+ if fc.exists():
+ shutil.copy2(fc, llvm_install / "bin" / "FileCheck")
+
+ # Regenerate patch if AI fix modified the source
+ if ctx.retry_count > 0:
+ patch_dir = ascend_path / "third_party/ascend/patch"
+ patch_dir.mkdir(parents=True, exist_ok=True)
+ new_patch_file = patch_dir / f"llvm_patch_{required_hash[:7]}.patch"
+ new_patch = run_git(llvm_project, "diff", "HEAD")
+ new_patch_file.write_text(new_patch, encoding="utf-8")
+ for old in patch_dir.glob("*.patch"):
+ if old.name != new_patch_file.name:
+ old.unlink()
+ log.info(f"Updated patch: {new_patch_file.name} ({len(new_patch)} bytes)")
+
+ log.status(True, "LLVM build passed")
+ return ctx.copy_with(build_passed=True)
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Triton-Ascend
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def build_triton(
+ ctx: WorkflowContext, config: TAConfig,
+ clean: bool = False, python_exe: str = "",
+) -> WorkflowContext:
+ """Build triton-ascend. Pure build — no retry logic."""
+ ascend_path = Path(ctx.triton_ascend_path)
+ llvm_install = config.llvm_install
+ llvm_prefix = config.llvm_install_prefix or (
+ str(llvm_install) if llvm_install.exists() else ""
+ )
+ python_exe = python_exe or config.python_exe or os.getenv("PYTHON", "python3.10")
+
+ if clean:
+ build_dir_path = ascend_path / "build"
+ if build_dir_path.exists():
+ subprocess.run(["rm", "-rf", str(build_dir_path)], check=False)
+
+ build_env = {
+ "LLVM_BUILD_DIR": llvm_prefix,
+ "LLVM_SYSPATH": llvm_prefix,
+ "LLVM_INSTALL_PREFIX": llvm_prefix,
+ "TRITON_BUILD_WITH_CCACHE": "true",
+ "TRITON_BUILD_WITH_CLANG_LLD": "true",
+ "TRITON_BUILD_PROTON": "OFF",
+ "DEBUG": "1",
+ "TRITON_WHEEL_NAME": "triton-ascend",
+ "TRITON_APPEND_CMAKE_ARGS": "-DTRITON_BUILD_UT=OFF",
+ "MAX_JOBS": str(config.build_procs),
+ "CMAKE_BUILD_PARALLEL_LEVEL": str(config.build_procs),
+ }
+
+ log.key_value("LLVM prefix", llvm_prefix if llvm_prefix else "(empty)")
+
+ step_id = ctx.steps[ctx.current_step]["id"] if ctx.steps else "step-0"
+ step_dir = WORKSPACE_DIR / STEPS_DIR / step_id
+ step_dir.mkdir(parents=True, exist_ok=True)
+ build_log = step_dir / "build.log"
+
+ log.section("Build Triton-Ascend")
+ log.info(f"Running: {python_exe} setup.py install")
+ with open(build_log, "w", encoding="utf-8") as fh:
+ fh.write(f"=== setup.py install ===\n{' '.join(build_env.keys())}\n\n")
+ fh.flush()
+ rc = stream_cmd(
+ [python_exe, "setup.py", "install"],
+ cwd=ascend_path,
+ log_fh=fh,
+ timeout=1800,
+ env=build_env,
+ label="Building Triton-Ascend",
+ )
+ passed = rc == 0
+
+ result = {
+ "all_passed": passed,
+ "steps": [
+ {"step": "setup_py_install", "passed": passed, "exit_code": rc}
+ ],
+ }
+ (step_dir / BUILD_RESULT_FILE).write_text(
+ json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
+ )
+
+ if not passed:
+ log.error(f"Build FAILED — see {build_log}")
+ return ctx.copy_with(
+ build_passed=False, fix_errors=[str(build_log)]
+ )
+ log.status(True, "Build passed")
+ return ctx.copy_with(build_passed=True)
+
+
+def commit_fixes(ctx: WorkflowContext, config: TAConfig) -> None:
+ """Commit AI build/test fixes with AI-authored commit message.
+
+ Commit message priority:
+ 1. ``commit_message.txt`` written by AI (one-line subject)
+ 2. First line of ``step_summary.md`` written by AI
+ 3. Default generic message
+
+ Commits submodule changes first (AscendNPU-IR), then parent repo.
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ step = ctx.steps[ctx.current_step] if ctx.current_step < len(ctx.steps) else None
+ step_id = step["id"] if step else "step-0"
+ step_dir = WORKSPACE_DIR / STEPS_DIR / step_id
+ target_short = ctx.target_commit[:12]
+
+ # ── 1. Commit submodule changes first ──────────────────────────────
+ if submodule_has_changes(ascend_path):
+ commit_submodule(
+ ascend_path,
+ f"[Sync](fix) AI fix for {target_short}\n",
+ )
+
+ # ── 2. Clean temp files before staging ─────────────────────────────
+ cleanup_temp_files(ascend_path)
+
+ # ── 3. Check if there's anything to commit ─────────────────────────
+ status = run_git(ascend_path, "status", "--porcelain").strip()
+ if not status:
+ log.info("No uncommitted fix changes — nothing to commit")
+ return
+
+ # ── 4. Build commit message (AI-authored priority) ─────────────────
+ commit_subject = _read_ai_commit_subject(step_dir, target_short)
+ commit_msg = (
+ f"[Sync](fix) {commit_subject}\n\n"
+ f"Upstream target: {target_short}\n"
+ f"Fix attempt: {ctx.retry_count}\n"
+ f"Work branch: {ctx.work_branch}\n"
+ )
+
+ # ── 5. Stage and commit ────────────────────────────────────────────
+ log.section("Commit AI Fixes")
+ try:
+ staged_files = run_git(ascend_path, "diff", "--name-only", "HEAD").strip()
+ run_git(ascend_path, "add", "-A")
+ staged = run_git(ascend_path, "diff", "--cached", "--name-only").strip()
+ if staged:
+ files = staged.splitlines()
+ log.info(f"Files staged ({len(files)}):")
+ for f in files[:20]:
+ log.info(f" {f}")
+ if len(files) > 20:
+ log.info(f" ... and {len(files) - 20} more")
+ run_git(ascend_path, "commit", "-s", "-m", commit_msg)
+ log.status(True, f"Committed: {commit_subject[:60]}")
+ except Exception as e:
+ stderr = (getattr(e, "stderr", "") or "").strip()
+ if "nothing to commit" in stderr.lower():
+ log.info("Nothing to commit (AI made no changes)")
+ else:
+ log.warning(f"Failed to commit fixes: {stderr[-200:]}")
+
+
+def _read_ai_commit_subject(step_dir: Path, target_short: str) -> str:
+ """Extract AI-authored commit subject from fix output files.
+
+ Priority:
+ 1. ``commit_message.txt`` — AI-written one-line subject
+ 2. ``step_summary.md`` — first heading line of AI summary
+ 3. Default generic fallback
+ """
+ # Priority 1: AI-written commit_message.txt
+ cmt_file = step_dir / "commit_message.txt"
+ if cmt_file.exists():
+ try:
+ subject = cmt_file.read_text(encoding="utf-8").strip()
+ # Take first line only, cap at 72 chars
+ subject = subject.split("\n")[0].strip()[:72]
+ if subject:
+ log.info(f"Using AI-written commit message: {subject}")
+ return subject
+ except Exception:
+ pass
+
+ # Priority 2: First line of step_summary.md
+ summary_file = step_dir / "step_summary.md"
+ if summary_file.exists():
+ try:
+ first_line = summary_file.read_text(encoding="utf-8").strip().split("\n")[0]
+ # Strip leading # marks and whitespace
+ subject = first_line.lstrip("#").strip()[:72]
+ if subject:
+ log.info(f"Using step_summary.md first line: {subject}")
+ return subject
+ except Exception:
+ pass
+
+ # Priority 3: Default
+ subject = f"Resolve build/test failures for {target_short}"
+ log.info(f"Using default commit message: {subject}")
+ return subject
diff --git a/src/TA_main2main_workflow/pipeline/commit.py b/src/TA_main2main_workflow/pipeline/commit.py
new file mode 100644
index 0000000..4c31c51
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/commit.py
@@ -0,0 +1,77 @@
+"""Pipeline step: Commit step progress.
+
+Handles submodule commit first (AscendNPU-IR), then parent repo commit.
+"""
+
+from __future__ import annotations
+
+from pathlib import Path
+
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.git import run_git
+from TA_main2main_workflow.utils.submodule import (
+ commit_submodule,
+ submodule_has_changes,
+)
+from TA_main2main_workflow.pipeline.pre_ci import cleanup_temp_files
+
+log = get_logger(__name__)
+
+
+def commit_step(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Commit all changes for the current step.
+
+ Order:
+ 1. Commit AscendNPU-IR submodule if it has changes
+ 2. Clean temp files
+ 3. Stage and commit parent repo
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ step = ctx.steps[ctx.current_step]
+ step_id = step["id"]
+
+ # ── 1. Submodule first ────────────────────────────────────────────
+ if submodule_has_changes(ascend_path):
+ target_short = ctx.target_commit[:12] if ctx.target_commit else "HEAD"
+ commit_submodule(
+ ascend_path,
+ f"[Sync](fix) AI fix for {target_short}\n"
+ )
+
+ # ── 2. Clean temp files ───────────────────────────────────────────
+ cleanup_temp_files(ascend_path)
+
+ # ── 3. Stage and commit parent repo ───────────────────────────────
+ staged = run_git(ascend_path, "status", "--porcelain").strip()
+ if not staged:
+ log.info(f"[{step_id}] Nothing to commit")
+ return ctx
+
+ # Print staged files for visibility
+ staged_files = [l[3:] for l in staged.splitlines() if l.strip()]
+ log.info(f"Files staged ({len(staged_files)}):")
+ for f in staged_files[:30]:
+ log.info(f" {f}")
+ if len(staged_files) > 30:
+ log.info(f" ... and {len(staged_files) - 30} more")
+
+ start_short = step.get("start_commit", "?")[:12]
+ end_short = step["end_commit"][:12]
+ msg = (
+ f"[Sync](feat) Merge upstream commits for step {step_id}"
+ f"({start_short}..{end_short}, {step['commit_count']} commits)\n\n"
+ f"Upstream range: {start_short}..{end_short}\n"
+ f"Step: {ctx.current_step + 1}/{ctx.total_steps}\n"
+ f"Commits: {step['commit_count']}\n"
+ )
+ try:
+ run_git(ascend_path, "add", "-A")
+ run_git(ascend_path, "commit", "-s", "-m", msg)
+ log.status(True, f"Committed step {step_id}")
+ except Exception as e:
+ if "nothing to commit" not in str(getattr(e, "stderr", "")):
+ log.warning(f"Commit failed: {e}")
+
+ return ctx
diff --git a/src/TA_main2main_workflow/pipeline/detect.py b/src/TA_main2main_workflow/pipeline/detect.py
new file mode 100644
index 0000000..c6bd6f4
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/detect.py
@@ -0,0 +1,162 @@
+"""Pipeline step 1: Detect upstream commits to merge.
+
+Entry point: ``run_detect(ctx, config)`` — handles resume from
+``detect.json`` or runs full detection from scratch.
+
+Calculates the commit gap between triton-ascend and upstream Triton:
+ - Finds merge-base
+ - Lists upstream commits since merge-base
+ - Counts changed files and lines
+"""
+
+from __future__ import annotations
+
+import json
+from pathlib import Path
+
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils import DETECT_FILE, WORKSPACE_DIR
+from TA_main2main_workflow.utils.git import run_git
+from TA_main2main_workflow.utils.logging import get_logger
+
+log = get_logger(__name__)
+
+
+def run_detect(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Detect upstream commits (with resume support).
+
+ If ``config.resume`` is set and ``detect.json`` exists, loads cached
+ results. Otherwise runs full detection via :func:`_detect_commits`.
+ """
+ detect_file = WORKSPACE_DIR / DETECT_FILE
+ if config.resume and detect_file.exists():
+ log.info("Resume: detect.json exists, skipping detect")
+ data = json.loads(detect_file.read_text(encoding="utf-8"))
+ return ctx.copy_with(
+ merge_base=data["merge_base"],
+ target_commit=data["target_commit"],
+ upstream_commits=data.get("upstream_commits", []),
+ upstream_commits_count=data["upstream_commits_count"],
+ changed_files_count=data.get("changed_files_count", 0),
+ changed_lines_total=data.get("changed_lines", 0),
+ has_new_commits=True,
+ ascend_head=data.get("ascend_head", ""),
+ )
+
+ return _detect_commits(ctx, config)
+
+
+def _detect_commits(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Detect upstream commits that need to be merged.
+
+ Assumes ``prepare`` has already run (remotes configured, fetched, and
+ ascend_head / target_commit resolved).
+
+ Returns updated ctx with merge_base, upstream_commits, has_new_commits,
+ changed_files_count, changed_lines_total.
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ ascend_head = ctx.ascend_head
+ target = ctx.target_commit
+
+ # Compute merge-base
+ try:
+ merge_base = run_git(ascend_path, "merge-base", ascend_head, target).strip()
+ except Exception:
+ raise RuntimeError(
+ f"No common ancestor between ascend HEAD ({ascend_head[:12]}) "
+ f"and target ({target[:12]}).\n"
+ f"Ensure triton-upstream remote points to the correct triton repo:\n"
+ f" cd {ascend_path} && git remote -v\n"
+ f"Current upstream URL: {config.triton_upstream_url}"
+ )
+ log.info(f"merge_base: {merge_base[:12]} target: {target[:12]}")
+
+ commits = _list_upstream_commits(ascend_path, merge_base, target)
+ has_new = len(commits) > 0 and merge_base != target
+
+ changed_files = _changed_files(ascend_path, merge_base, target)
+ changed_lines_total = _count_changed_lines(ascend_path, merge_base, target)
+
+ # ── Print detection summary ──
+ log.key_value("upstream commits", str(len(commits)))
+ log.key_value("changed files", str(len(changed_files)))
+ log.key_value("changed lines", str(changed_lines_total))
+ if commits:
+ log.info(f"First {min(20, len(commits))} upstream commits:")
+ for c in commits[:20]:
+ log.info(f" {c['sha'][:8]} {c['subject'][:100]}")
+
+ result = {
+ "ascend_head": ascend_head,
+ "target_commit": target,
+ "merge_base": merge_base,
+ "upstream_commits_count": len(commits),
+ "upstream_commits": commits,
+ "changed_lines": changed_lines_total,
+ "changed_files": changed_files,
+ "changed_files_count": len(changed_files),
+ }
+
+ # Write detect.json
+ (WORKSPACE_DIR / DETECT_FILE).write_text(
+ json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
+ )
+
+ return ctx.copy_with(
+ merge_base=merge_base,
+ target_commit=target,
+ ascend_head=ascend_head,
+ upstream_commits=commits,
+ upstream_commits_count=len(commits),
+ changed_files_count=result["changed_files_count"],
+ changed_lines_total=changed_lines_total,
+ has_new_commits=has_new,
+ )
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Internal helpers
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _list_upstream_commits(repo: Path, merge_base: str, target: str) -> list[dict]:
+ output = run_git(
+ repo, "log", "--reverse", "--format=%H%x1f%s", f"{merge_base}..{target}"
+ )
+ commits: list[dict] = []
+ for line in output.strip().splitlines():
+ if not line.strip():
+ continue
+ parts = line.split("\x1f", 1)
+ commits.append(
+ {
+ "sha": parts[0].strip(),
+ "subject": parts[1].strip() if len(parts) > 1 else "",
+ }
+ )
+ return commits
+
+
+def _count_changed_lines(repo: Path, merge_base: str, target: str) -> int:
+ """Return total lines changed between *merge_base* and *target*."""
+ try:
+ output = run_git(repo, "diff", "--shortstat", merge_base, target)
+ except Exception:
+ return 0
+ # " 97 files changed, 1234 insertions(+), 567 deletions(-)"
+ total = 0
+ for part in output.split(","):
+ part = part.strip()
+ if "insertion" in part or "deletion" in part:
+ try:
+ total += int(part.split()[0])
+ except ValueError:
+ pass
+ return total
+
+
+def _changed_files(repo: Path, merge_base: str, target: str) -> list[str]:
+ output = run_git(repo, "diff", "--name-only", merge_base, target)
+ return sorted(f for f in output.strip().splitlines() if f)
diff --git a/src/TA_main2main_workflow/pipeline/finalize.py b/src/TA_main2main_workflow/pipeline/finalize.py
new file mode 100644
index 0000000..09eb2e8
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/finalize.py
@@ -0,0 +1,231 @@
+"""Pipeline step: Finalize — generate cumulative patch, PR description, and sync report.
+
+PR description is AI-generated via the ``report`` mode and saved as
+``final_summary.md``. Falls back to a basic template if AI is
+unavailable or skipped.
+"""
+
+from __future__ import annotations
+
+import json
+import time
+from pathlib import Path
+
+from TA_main2main_workflow.agent.opencode_adapter import run_opencode_adapter
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.git import run_git
+from TA_main2main_workflow.utils.tracker import total_elapsed
+from TA_main2main_workflow.utils import (
+ FINAL_SUMMARY_FILE,
+ FINAL_TARGET_PATCH_FILE,
+ STEPS_DIR,
+ WORKSPACE_DIR,
+)
+
+log = get_logger(__name__)
+_REF = str(Path(__file__).parent.parent / "reference")
+
+
+def finalize(ctx: WorkflowContext, config: TAConfig | None = None) -> WorkflowContext:
+ """Generate final summary, cumulative patch, and sync report.
+
+ Uses AI (report mode) for the PR description when available.
+ """
+ log.header("Finalize & Summary")
+ ascend_path = Path(ctx.triton_ascend_path)
+
+ # ── Cumulative patch ──────────────────────────────────────────────
+ try:
+ patch = run_git(ascend_path, "diff", ctx.ascend_head, "HEAD")
+ patch_path = WORKSPACE_DIR / FINAL_TARGET_PATCH_FILE
+ patch_path.write_text(patch, encoding="utf-8")
+ log.info(f"Cumulative patch: {len(patch)} bytes → {patch_path}")
+ except Exception as e:
+ log.warning(f"Could not generate patch: {e}")
+
+ # ── PR description (AI-generated) ─────────────────────────────────
+ summary_path = WORKSPACE_DIR / FINAL_SUMMARY_FILE
+ if config and not config.skip_ai_analysis:
+ try:
+ _generate_pr_description(ctx, config, summary_path)
+ except Exception as e:
+ log.warning(f"AI PR description failed: {e} — using fallback")
+ _write_summary_fallback(ctx, summary_path)
+ else:
+ _write_summary_fallback(ctx, summary_path)
+
+ # ── Sync Report ───────────────────────────────────────────────────
+ _write_sync_report(ctx)
+
+ # ── Print final table ─────────────────────────────────────────────
+ elapsed = total_elapsed()
+ log.elapsed(elapsed)
+ rows = list(ctx.summary_rows or [])
+ rows.append(("Finalize", "PASS", f"{ctx.total_steps} step(s)"))
+ rows.append(("OVERALL", "PASS", f"{ctx.total_steps} step(s)"))
+ log.table(rows)
+
+ return ctx
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Internal
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _generate_pr_description(
+ ctx: WorkflowContext, config: TAConfig, summary_path: Path,
+) -> None:
+ """Invoke AI (report mode) to write the PR description."""
+ step = ctx.steps[-1] if ctx.steps else {"id": "step-0"}
+ step_dir = WORKSPACE_DIR / STEPS_DIR / step["id"]
+ step_dir.mkdir(parents=True, exist_ok=True)
+
+ # Build context for the AI: collect per-step data + fix records
+ context = _build_report_context(ctx)
+
+ log.section("AI PR Description (report mode)")
+ result = run_opencode_adapter({
+ "step_id": "finalize-report",
+ "previous_step_id": "",
+ "previous_step_summary_path": "",
+ "is_last_step": "true",
+ "step_dir": str(step_dir),
+ "fix_dir": str(step_dir),
+ "conflict_dir": str(WORKSPACE_DIR / "conflicts"),
+ "ascend_path": str(Path(ctx.triton_ascend_path)),
+ "triton_path": ctx.triton_ascend_path,
+ "reference_dir": _REF,
+ "mode": "report",
+ "error_logs": json.dumps(context, ensure_ascii=False, default=str),
+ "target_commit": ctx.target_commit,
+ "step_index": f"{ctx.total_steps}/{ctx.total_steps}",
+ "upstream_commits_count": str(ctx.upstream_commits_count),
+ "total_steps": str(ctx.total_steps),
+ "conflict_files_resolved": str(ctx.conflict_files_resolved),
+ "build_fix_count": str(ctx.build_fix_count),
+ "test_fix_count": str(ctx.test_fix_count),
+ "final_status": ctx.final_status or "Success",
+ "ascend_npu_ir_fix": "false",
+ "ascend_npu_ir_compat_ref": "",
+ })
+
+ # AI writes to step_dir/step_summary.md; copy to final location
+ ai_summary = step_dir / "step_summary.md"
+ if ai_summary.exists():
+ content = ai_summary.read_text(encoding="utf-8")
+ summary_path.write_text(content, encoding="utf-8")
+ log.status(True, f"AI PR description: {summary_path} ({len(content)} bytes)")
+ elif result.step_summary:
+ summary_path.write_text(result.step_summary, encoding="utf-8")
+ log.status(True, f"AI PR description (from output): {summary_path}")
+ else:
+ log.warning("AI produced no summary — using fallback")
+ _write_summary_fallback(ctx, summary_path)
+
+
+def _build_report_context(ctx: WorkflowContext) -> dict:
+ """Collect all sync data for the AI report."""
+ steps_dir = WORKSPACE_DIR / STEPS_DIR
+
+ # Collect per-step summaries and fix details
+ step_data: list[dict] = []
+ for s in ctx.steps:
+ sd: dict = {
+ "id": s["id"],
+ "commits": s["commit_count"],
+ "start_commit": s.get("start_commit", "")[:12],
+ "end_commit": s.get("end_commit", "")[:12],
+ "source_lines": s.get("source_changed_lines", 0),
+ "reason": s.get("reason", "line_budget"),
+ }
+ # Include step summary if AI wrote one
+ step_dir = steps_dir / s["id"]
+ summary_file = step_dir / "step_summary.md"
+ if summary_file.exists():
+ try:
+ sd["summary"] = summary_file.read_text(encoding="utf-8")[:4000]
+ except Exception:
+ pass
+ # Include commit list
+ commits_file = step_dir / "commits.txt"
+ if commits_file.exists():
+ try:
+ sd["commit_list"] = commits_file.read_text(encoding="utf-8")[:2000]
+ except Exception:
+ pass
+ step_data.append(sd)
+
+ return {
+ "target_commit": ctx.target_commit[:12],
+ "upstream_commits_count": ctx.upstream_commits_count,
+ "total_steps": ctx.total_steps,
+ "work_branch": ctx.work_branch,
+ "conflict_files_resolved": ctx.conflict_files_resolved,
+ "build_fix_count": ctx.build_fix_count,
+ "test_fix_count": ctx.test_fix_count,
+ "final_status": ctx.final_status or "Success",
+ "steps": step_data,
+ "step_pr_descriptions": ctx.step_pr_descriptions,
+ "step_details": ctx.step_details,
+ "ir_analysis_done": ctx.ir_analysis_done,
+ }
+
+
+def _write_summary_fallback(ctx: WorkflowContext, summary_path: Path) -> None:
+ """Write a basic PR description when AI is unavailable."""
+ parts = [
+ "## Summary",
+ f"- **Target**: `{ctx.target_commit[:12]}`",
+ f"- **Steps**: {ctx.total_steps}",
+ f"- **Upstream commits**: {ctx.upstream_commits_count}",
+ f"- **Status**: Success",
+ f"- **Date**: {time.strftime('%Y-%m-%d %H:%M:%S')}",
+ f"- **Work branch**: `{ctx.work_branch}`",
+ ]
+ if ctx.step_details:
+ parts.append("\n## Changes\n")
+ for d in ctx.step_details:
+ parts.append(
+ f"- **{d['step_id']}**: {d['commits']} commits, "
+ f"end=`{d.get('end_commit', '?')[:12]}`, "
+ f"build_fixes={d.get('build_fixes', 0)}, "
+ f"test_fixes={d.get('test_fixes', 0)}"
+ )
+ if ctx.step_pr_descriptions:
+ parts.append("\n## Details\n")
+ for desc in ctx.step_pr_descriptions:
+ parts.append(f"- {desc}")
+
+ summary_path.write_text("\n".join(parts) + "\n", encoding="utf-8")
+ log.info(f"PR description (fallback): {summary_path}")
+
+
+def _write_sync_report(ctx: WorkflowContext) -> None:
+ """Generate a human-readable sync report (fallback, no AI)."""
+ report_path = WORKSPACE_DIR / "SYNC_REPORT.md"
+ try:
+ parts = [
+ f"# Triton-Ascend 上游同步报告\n",
+ f"## 基本信息\n",
+ f"- 目标提交: `{ctx.target_commit[:12]}`",
+ f"- 步骤数: {ctx.total_steps}",
+ f"- 上游提交数: {ctx.upstream_commits_count}",
+ f"- 工作分支: `{ctx.work_branch}`",
+ f"- 状态: 成功",
+ ]
+ if ctx.step_details:
+ parts.append(f"\n## 步骤详情\n")
+ for d in ctx.step_details:
+ parts.append(
+ f"### {d['step_id']}\n"
+ f"- 提交数: {d['commits']}\n"
+ f"- 构建修复: {d.get('build_fixes', 0)}\n"
+ f"- 测试修复: {d.get('test_fixes', 0)}\n"
+ )
+ report_path.write_text("\n".join(parts), encoding="utf-8")
+ log.info(f"Sync report: {report_path}")
+ except Exception as e:
+ log.warning(f"Could not write sync report: {e}")
diff --git a/src/TA_main2main_workflow/pipeline/fix.py b/src/TA_main2main_workflow/pipeline/fix.py
new file mode 100644
index 0000000..dd39df8
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/fix.py
@@ -0,0 +1,177 @@
+"""Pipeline step: AI fix build/test failures with fix validation gate."""
+
+from __future__ import annotations
+
+import json
+from pathlib import Path
+
+from TA_main2main_workflow.agent.opencode_adapter import run_opencode_adapter
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.git import run_git
+from TA_main2main_workflow.utils import FIX_LOG_DIR, STEPS_DIR, WORKSPACE_DIR
+
+log = get_logger(__name__)
+_REF = str(Path(__file__).parent.parent / "reference")
+
+
+def ai_fix(ctx: WorkflowContext, config: TAConfig, attempt: int = 1,
+ mode: str = "fix") -> WorkflowContext:
+ """Invoke AI to fix build or test failures.
+
+ Args:
+ ctx: Current workflow context
+ config: Workflow configuration
+ attempt: Fix attempt number (1-based)
+ mode: AI mode — ``"fix"`` for build/test failures,
+ ``"ir_patch"`` for IR patch adjustments
+
+ Returns updated context. On success, the AI will have modified files
+ on disk; caller is responsible for committing and rebuilding/retesting.
+ """
+ if config.skip_ai_analysis:
+ log.info("SKIP_AI_ANALYSIS=true — skipping AI fix")
+ return ctx
+
+ ascend_path = Path(ctx.triton_ascend_path)
+ step = ctx.steps[ctx.current_step] if ctx.current_step < len(ctx.steps) else None
+ step_id = step["id"] if step else "step-0"
+ step_dir = WORKSPACE_DIR / STEPS_DIR / step_id
+ step_dir.mkdir(parents=True, exist_ok=True)
+ fix_dir = WORKSPACE_DIR / FIX_LOG_DIR / f"{step_id}-fix-{attempt}"
+ fix_dir.mkdir(parents=True, exist_ok=True)
+
+ # ── Compute AI context: previous step info ──────────────────────────
+ prev_step_id = ""
+ prev_summary_path = ""
+ if ctx.current_step > 0 and ctx.current_step <= len(ctx.steps):
+ prev = ctx.steps[ctx.current_step - 1]
+ prev_step_id = prev["id"]
+ prev_summary = WORKSPACE_DIR / STEPS_DIR / prev_step_id / "step_summary.md"
+ prev_summary_path = str(prev_summary) if prev_summary.exists() else ""
+ is_last_step = (ctx.current_step >= ctx.total_steps - 1)
+ ascend_npu_ir_fix = _detect_ascend_npu_ir_errors(ascend_path, step_id)
+ ascend_npu_ir_compat_ref = str(
+ Path(__file__).parent.parent / "reference" / "AscendNPU-IR_LLVM_VERSION_COMPAT.md"
+ )
+ conflict_dir = str(WORKSPACE_DIR / "conflicts")
+
+ log.step(attempt, config.max_retries, f"AI {mode}")
+ try:
+ # Record pre-fix file list for validation
+ pre_files = _list_tracked_files(ascend_path)
+
+ result = run_opencode_adapter(
+ {
+ "step_id": f"{step_id}-{mode}-{attempt}",
+ "previous_step_id": prev_step_id,
+ "previous_step_summary_path": prev_summary_path,
+ "is_last_step": str(is_last_step).lower(),
+ "step_dir": str(step_dir),
+ "fix_dir": str(fix_dir),
+ "conflict_dir": conflict_dir,
+ "ascend_path": str(ascend_path),
+ "triton_path": ctx.triton_ascend_path,
+ "reference_dir": _REF,
+ "mode": mode,
+ "error_logs": json.dumps(ctx.fix_errors, ensure_ascii=False),
+ "target_commit": ctx.target_commit,
+ "step_index": f"{ctx.current_step + 1}/{ctx.total_steps}",
+ "ascend_npu_ir_fix": str(ascend_npu_ir_fix).lower(),
+ "ascend_npu_ir_compat_ref": ascend_npu_ir_compat_ref,
+ }
+ )
+
+ # ── Fix validation gate ────────────────────────────────────────
+ is_valid, reason = validate_fix(
+ ascend_path, pre_files, result.modified_files
+ )
+ if not is_valid:
+ log.warning(f"Fix validation FAILED: {reason}")
+ log.warning("Reverting invalid changes...")
+ # Write rejection feedback so AI can adjust on next attempt
+ rejection_file = fix_dir / "fix_rejection.txt"
+ rejection_file.write_text(
+ f"VALIDATION REJECTED: {reason}\n"
+ f"Allowed paths: {', '.join(_ALLOWED_FIX_PREFIXES)}\n"
+ f"Modified files: {result.modified_files}\n",
+ encoding="utf-8",
+ )
+ _revert_illegal_changes(ascend_path)
+ return ctx
+
+ log.ai_result(
+ bool(result.modified_files),
+ result.modified_files,
+ (result.step_summary or "")[:500],
+ )
+ return ctx
+ except Exception as e:
+ log.error(f"AI fix failed: {e}")
+ return ctx
+
+
+# Paths AI is allowed to modify when fixing test failures.
+# Build fixes are restricted to third_party/ascend/ only.
+_ALLOWED_FIX_PREFIXES = [
+ "third_party/ascend/",
+ "python/triton/extension",
+ "python/triton/runtime/libentry.py",
+]
+
+
+def validate_fix(
+ ascend_path: Path,
+ pre_fix_files: set[str],
+ modified_files: list[str],
+) -> tuple[bool, str]:
+ """Validate that AI fixes only touch allowed paths.
+
+ Returns (is_valid, reason).
+ """
+ if not modified_files:
+ return False, "No files were modified"
+
+ for f in modified_files:
+ if not any(f.startswith(p) for p in _ALLOWED_FIX_PREFIXES):
+ return False, (
+ f"File '{f}' is outside allowed paths. "
+ f"Allowed: {', '.join(_ALLOWED_FIX_PREFIXES)}"
+ )
+
+ return True, "all changes within allowed paths"
+
+
+def _list_tracked_files(repo: Path) -> set[str]:
+ """Return the set of all tracked files in the repo."""
+ try:
+ output = run_git(repo, "ls-files")
+ return set(output.strip().splitlines())
+ except Exception:
+ return set()
+
+
+def _revert_illegal_changes(repo: Path) -> None:
+ """Revert all uncommitted changes and remove untracked files."""
+ try:
+ run_git(repo, "checkout", "--", ".")
+ run_git(repo, "clean", "-fd")
+ except Exception as e:
+ log.error(f"Failed to revert changes: {e}")
+
+
+def _detect_ascend_npu_ir_errors(ascend_path: Path, step_id: str) -> bool:
+ """Check if build errors are from AscendNPU-IR compilation failures."""
+ build_log = WORKSPACE_DIR / STEPS_DIR / step_id / "build.log"
+ if not build_log.exists():
+ return False
+ try:
+ content = build_log.read_text(encoding="utf-8", errors="replace").lower()
+ indicators = [
+ "AscendNPU-IR".lower(), "ascendnpu-ir",
+ "llvm::", "mlir::", "fatal error", "undefined reference",
+ ]
+ return any(ind in content for ind in indicators)
+ except Exception:
+ return False
diff --git a/src/TA_main2main_workflow/pipeline/ir_patch.py b/src/TA_main2main_workflow/pipeline/ir_patch.py
new file mode 100644
index 0000000..8b48dc2
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/ir_patch.py
@@ -0,0 +1,1069 @@
+"""Pipeline step: IR patch generation for LLVM version changes.
+
+Handles the full IR patch pipeline when the LLVM hash changes:
+ 1. Build baseline LLVM (pre-merge, from base branch's llvm-hash)
+ 2. Per-step: apply existing patch → build LLVM → build TA → test
+ 3. On failure: AI adjust patch → rebuild LLVM → rebuild TA → retest
+ 4. AI supplement missing IR patches → retest loop
+ 5. Fallback: full OP analysis pipeline
+
+Key entry points:
+ - ``build_baseline_llvm(ctx, config)`` — Build baseline LLVM once
+ - ``per_step_ir_patch(ctx, config, step)`` — Per-step IR patch flow
+"""
+
+from __future__ import annotations
+
+import json
+import os
+import shutil
+import subprocess
+import time
+from pathlib import Path
+
+from TA_main2main_workflow.agent.opencode_adapter import run_opencode_adapter
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.git import run_git, run_git_no_check, stream_cmd
+from TA_main2main_workflow.pipeline.build import build_and_fix_loop
+from TA_main2main_workflow.pipeline.test import (
+ run_tests, detect_oom_in_tests, rerun_tests_reduced_concurrency,
+ test_and_fix_loop, _run_pretest_and_fix,
+)
+from TA_main2main_workflow.pipeline.fix import ai_fix
+from TA_main2main_workflow.utils import (
+ WORKSPACE_DIR, STEPS_DIR, BUILD_RESULT_FILE, TEST_RESULT_FILE,
+ IR_ANALYSIS_DIR, IR_OPS_REPORT_FILE, IR_CHANGES_REPORT_FILE,
+ IR_DIAGNOSIS_FILE, IR_MAX_ITERATIONS, LLVM_CHANGE_ANALYSIS_DIR,
+ _ASCEND_BASELINE_LLVM_HASH,
+)
+
+log = get_logger(__name__)
+_REF = str(Path(__file__).parent.parent / "reference")
+
+# Maximum retries for LLVM patch apply/rebuild loop
+_MAX_LLVM_RETRIES = 10
+# Maximum retries for applying existing patch
+_MAX_PATCH_APPLY_RETRIES = 3
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Baseline LLVM build (pre-merge, called once)
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def build_baseline_llvm(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Build baseline LLVM from base branch's llvm-hash + existing patch.
+
+ Called once at the start of single-step mode. This ensures LLVM is
+ built and ready before any merge steps begin. After building, stashes
+ the patch changes so the LLVM tree is clean for subsequent steps.
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ llvm_project = config.llvm_project
+ llvm_install = config.llvm_install
+
+ hash_file = ascend_path / "cmake" / "llvm-hash.txt"
+ if not hash_file.exists():
+ log.info("No llvm-hash.txt — skipping baseline LLVM build")
+ return ctx
+
+ required_hash = hash_file.read_text(encoding="utf-8").strip()
+ if not required_hash:
+ return ctx
+
+ log.header("Build Baseline LLVM (pre-merge)")
+
+ # ── Print LLVM workspace info ──
+ log.key_value("LLVM project", str(llvm_project))
+ log.key_value("LLVM install prefix", str(llvm_install))
+ log.key_value("Target LLVM hash", required_hash[:12])
+
+ # Allow skipping baseline LLVM build (LLVM already built for current TA)
+ if config.skip_baseline_llvm:
+ log.status(True, "SKIP_BASELINE_LLVM set — assuming baseline LLVM already built")
+ return ctx.copy_with(build_passed=True)
+
+ # Ensure llvm-project exists
+ if not llvm_project.exists():
+ run_git(WORKSPACE_DIR, "clone", config.llvm_repo_url, str(llvm_project))
+
+ # Ensure clean workspace
+ _ensure_llvm_workspace_clean(llvm_project, "baseline LLVM build")
+
+ # Checkout the required hash
+ log.info(f"Checking out LLVM hash: {required_hash[:12]}")
+ _ensure_commit_available(llvm_project, required_hash)
+ run_git(llvm_project, "checkout", "-f", required_hash)
+ run_git(llvm_project, "clean", "-fd")
+
+ # Apply existing Ascend patch
+ patch_dir = ascend_path / "third_party/ascend/patch"
+ patch_files = sorted(patch_dir.glob("*.patch")) if patch_dir.exists() else []
+ if patch_files:
+ patch_file = patch_files[0]
+ log.info(f"Applying existing patch: {patch_file.name}")
+ try:
+ run_git(llvm_project, "apply", str(patch_file))
+ except Exception as e:
+ log.warning(f"Patch apply failed: {e} — will build without patch")
+ else:
+ log.info("No existing Ascend patch found — building clean LLVM")
+
+ # Build LLVM
+ try:
+ prefix = _do_llvm_build(llvm_project, llvm_install, required_hash)
+ log.status(True, f"Baseline LLVM built at {prefix}")
+ except Exception as e:
+ log.error(f"Baseline LLVM build failed: {e}")
+ return ctx.copy_with(build_passed=False)
+
+ # Stash/drop patch changes to leave clean tree
+ _ensure_llvm_workspace_clean(llvm_project, "post-baseline build")
+
+ # Set SKIP_LLVM_REBUILD so downstream builds don't wipe LLVM
+ os.environ["SKIP_LLVM_REBUILD"] = "true"
+
+ return ctx.copy_with(build_passed=True)
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Per-step IR patch pipeline
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def per_step_ir_patch(ctx: WorkflowContext, config: TAConfig,
+ step: dict) -> WorkflowContext:
+ """Full per-step IR patch pipeline for LLVM version changes.
+
+ Strategy: apply existing patch first → build → test → AI supplement.
+ Falls back to full OP analysis if the existing-patch-first approach
+ exhausts retries.
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ step_id = step["id"]
+
+ if config.skip_ir_patch:
+ log.header(f"IR Patch Pipeline — {step_id} — SKIPPED (SKIP_IR_PATCH=true)")
+ return ctx.copy_with(build_passed=True, test_passed=True)
+
+ step_dir = WORKSPACE_DIR / STEPS_DIR / step_id
+ step_dir.mkdir(parents=True, exist_ok=True)
+
+ target_llvm_hash = _get_current_llvm_hash(ascend_path)
+ log.header(f"IR Patch Pipeline — {step_id}")
+ log.key_value("target LLVM hash", target_llvm_hash[:12])
+
+ # ── Phase 1: Apply existing patch + build LLVM ──────────────────
+ if not _do_apply_existing_patch(ctx, config, step, target_llvm_hash):
+ log.warning("Existing patch apply failed after retries — falling back")
+ return _per_step_ir_patch_fallback(ctx, config, step, target_llvm_hash)
+
+ # ── Phase 2: Build Triton-Ascend with AI fix loop ──────────────
+ log.section(f"Build TA with AI Fix — {step_id}")
+ build_ctx = _do_ta_build_with_fix(ctx, config, step)
+ if not build_ctx.build_passed:
+ log.error(f"TA build failed for {step_id}")
+ return build_ctx
+
+ # ── Phase 3: Test + IR supplement loop ─────────────────────────
+ log.section(f"Test + IR Supplement — {step_id}")
+ ctx = _do_test_and_fix_with_ir_retry(build_ctx, config, step, target_llvm_hash)
+
+ return ctx
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Phase 1: Apply existing patch
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _do_apply_existing_patch(
+ ctx: WorkflowContext, config: TAConfig, step: dict, target_llvm_hash: str,
+) -> bool:
+ """Apply existing Ascend LLVM patch to llvm-project, with AI fix retry.
+
+ Returns True if patch applies and LLVM builds successfully.
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ llvm_project = config.llvm_project
+ llvm_install = config.llvm_install
+ step_id = step["id"]
+
+ patch_dir = ascend_path / "third_party/ascend/patch"
+ patch_files = sorted(patch_dir.glob("*.patch")) if patch_dir.exists() else []
+
+ if not patch_files:
+ log.info("No existing patch — building LLVM directly")
+ if config.skip_llvm_rebuild:
+ log.status(True, "SKIP_LLVM_REBUILD set — assuming LLVM already built")
+ return True
+ try:
+ _do_llvm_build(llvm_project, llvm_install, target_llvm_hash)
+ return True
+ except Exception:
+ return False
+
+ for attempt in range(1, _MAX_PATCH_APPLY_RETRIES + 1):
+ log.step(attempt, _MAX_PATCH_APPLY_RETRIES, "Apply existing patch")
+
+ # ── Clean LLVM workspace ──
+ _ensure_llvm_workspace_clean(llvm_project, f"patch apply attempt {attempt}")
+ _ensure_commit_available(llvm_project, target_llvm_hash)
+ run_git(llvm_project, "checkout", "-f", target_llvm_hash)
+ run_git(llvm_project, "clean", "-fd")
+
+ # ── Apply patch ──
+ patch_file = patch_files[0]
+ log.info(f"Applying: {patch_file.name}")
+ try:
+ run_git(llvm_project, "apply", str(patch_file))
+ except Exception as e:
+ log.error(f"Patch apply failed: {e}")
+ if attempt > 1:
+ continue
+ # First failure: try AI fix
+ log.info("Attempting AI patch adjustment...")
+ try:
+ _ai_adjust_patch_for_failure(ctx, config, step, str(e))
+ except Exception:
+ pass
+ continue
+ log.info(f"Applying existing patch Successfully (attempt {attempt})")
+ # ── Build LLVM ──
+ if config.skip_llvm_rebuild:
+ log.status(True, f"SKIP_LLVM_REBUILD set — assuming LLVM already built (attempt {attempt})")
+ return True
+ try:
+ _do_llvm_build(llvm_project, llvm_install, target_llvm_hash)
+ log.status(True, f"LLVM build with existing patch OK (attempt {attempt})")
+ return True
+ except Exception as e:
+ log.error(f"LLVM build failed: {e}")
+ if attempt < _MAX_PATCH_APPLY_RETRIES:
+ log.info("Attempting AI patch fix...")
+ try:
+ _ai_adjust_patch_for_failure(ctx, config, step, str(e))
+ except Exception:
+ pass
+
+ return False
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Phase 2: Build TA with AI fix
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _do_ta_build_with_fix(
+ ctx: WorkflowContext, config: TAConfig, step: dict,
+) -> WorkflowContext:
+ """Build Triton-Ascend with AI fix loop for compile errors."""
+ return build_and_fix_loop(ctx, config)
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Phase 3: Test + IR supplement loop
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _do_test_and_fix_with_ir_retry(
+ ctx: WorkflowContext, config: TAConfig, step: dict, target_llvm_hash: str,
+) -> WorkflowContext:
+ """Test with IR supplement loop.
+
+ Runs pytest, classifies failures. IR issues → supplement patch →
+ rebuild LLVM → rebuild TA → retest. Code issues → AI fix loop.
+ Max 3 IR supplement iterations.
+ """
+ step_id = step["id"]
+ ascend_path = Path(ctx.triton_ascend_path)
+ llvm_project = config.llvm_project
+ llvm_install = config.llvm_install
+ ir_max = config.ir_max_iterations
+
+ for ir_iter in range(ir_max + 1):
+ if ir_iter > 0:
+ log.header(f"IR Supplement Iteration {ir_iter}/{ir_max}")
+
+ # ── Pre-test: smoke check before every test run ──────────────
+ if not config.skip_e2e_test:
+ ctx = _run_pretest_and_fix(ctx, config, ascend_path)
+ if not ctx.test_passed:
+ log.error("Pre-test failed after all retries — aborting")
+ return ctx.copy_with(test_passed=False, pytest_passed=False)
+
+ # ── Run tests ──
+ ctx = run_tests(ctx, config)
+ if ctx.test_passed:
+ log.status(True, f"All tests passed (IR iter {ir_iter})")
+ return ctx.copy_with(test_passed=True, pytest_passed=True)
+
+ # ── OOM detection ──
+ if detect_oom_in_tests(ctx):
+ log.warning("OOM detected — reducing concurrency")
+ oom_ctx = rerun_tests_reduced_concurrency(ascend_path, config)
+ if oom_ctx is not None and oom_ctx.test_passed:
+ return ctx.copy_with(test_passed=True, pytest_passed=True)
+
+ if ir_iter >= ir_max:
+ log.error(f"IR supplement loop exhausted ({ir_max} iterations)")
+ break
+
+ # ── Classify failures: IR vs code ──
+ is_ir_issue = _classify_test_failures(ctx, config, step)
+ if is_ir_issue:
+ log.info(f"IR issues detected — supplementing patch (iter {ir_iter + 1}/{ir_max})")
+ _ir_supplement_patch(ctx, config, step, target_llvm_hash,
+ supplement_iter=ir_iter + 1)
+ # Rebuild LLVM with updated patch
+ if config.skip_llvm_rebuild:
+ log.status(True, "SKIP_LLVM_REBUILD set — skipping LLVM rebuild after supplement")
+ else:
+ try:
+ # Re-apply updated patch to clean workspace before building
+ _clean_checkout_apply_patch(
+ llvm_project, ascend_path, target_llvm_hash,
+ reason=f"IR supplement iter {ir_iter + 1}",
+ )
+ _do_llvm_build(llvm_project, llvm_install, target_llvm_hash)
+ except Exception as e:
+ log.error(f"LLVM rebuild after supplement failed: {e}")
+ continue
+ # Rebuild TA
+ ctx = _do_ta_build_with_fix(ctx, config, step)
+ if not ctx.build_passed:
+ return ctx
+ else:
+ log.info("Code issues detected — entering AI fix loop")
+ ctx = _do_ai_fix_loop(ctx, config, step)
+ if ctx.test_passed:
+ return ctx
+
+ return ctx.copy_with(test_passed=False, pytest_passed=False)
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Fallback: Full OP analysis pipeline
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _per_step_ir_patch_fallback(
+ ctx: WorkflowContext, config: TAConfig, step: dict, target_llvm_hash: str,
+) -> WorkflowContext:
+ """Fallback: Full OP analysis pipeline when existing-patch-first fails.
+
+ Does: IR OP analysis → IR change analysis → IR generate patches →
+ apply patches + rebuild LLVM → build TA → test.
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ llvm_project = config.llvm_project
+ llvm_install = config.llvm_install
+ step_id = step["id"]
+
+ log.header(f"IR Patch Fallback — Full OP Analysis — {step_id}")
+
+ # 1. IR OP analysis (AI scans Ascend code for MLIR OP usage)
+ log.section("IR OP Analysis")
+ try:
+ ops_report = _run_ir_op_analysis(ctx, config)
+ ctx = ctx.copy_with(ir_ops_report=ops_report, ir_analysis_done=True)
+ except Exception as e:
+ log.warning(f"OP analysis failed: {e}")
+
+ # 2. IR change analysis (AI compares OP definitions between LLVM versions)
+ log.section("IR Change Analysis")
+ try:
+ changes_report = _run_ir_change_analysis(ctx, config, target_llvm_hash)
+ ctx = ctx.copy_with(ir_changes_report=changes_report)
+ except Exception as e:
+ log.warning(f"Change analysis failed: {e}")
+
+ # 3. Generate IR patches
+ log.section("IR Patch Generation")
+ try:
+ _run_ir_generate_patches(ctx, config, step, target_llvm_hash)
+ except Exception as e:
+ log.warning(f"Patch generation failed: {e}")
+
+ # 4. Apply generated patch to clean workspace, then build LLVM
+ if config.skip_llvm_rebuild:
+ log.status(True, "SKIP_LLVM_REBUILD set — skipping LLVM build in fallback")
+ else:
+ try:
+ _clean_checkout_apply_patch(
+ llvm_project, ascend_path, target_llvm_hash,
+ reason="fallback IR pipeline",
+ )
+ _do_llvm_build(llvm_project, llvm_install, target_llvm_hash)
+ except Exception as e:
+ log.error(f"LLVM build failed: {e}")
+ return ctx.copy_with(build_passed=False)
+
+ # 5. Build TA
+ ctx = _do_ta_build_with_fix(ctx, config, step)
+ if not ctx.build_passed:
+ return ctx
+
+ # 6. Test
+ ctx = run_tests(ctx, config)
+ return ctx
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# IR analysis sub-steps
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _ir_ai_base(ctx: WorkflowContext, config: TAConfig, step_id: str,
+ ir_dir: Path, mode: str, error_logs: str = "[]",
+ extra: dict | None = None) -> dict:
+ """Build the common AI context dict for IR patch calls."""
+ ascend_path = Path(ctx.triton_ascend_path)
+ base = {
+ "step_id": f"{step_id}-{mode}",
+ "previous_step_id": "",
+ "previous_step_summary_path": "",
+ "is_last_step": "true",
+ "step_dir": str(WORKSPACE_DIR / STEPS_DIR / step_id),
+ "fix_dir": str(ir_dir),
+ "conflict_dir": str(WORKSPACE_DIR / "conflicts"),
+ "ascend_path": str(ascend_path),
+ "triton_path": ctx.triton_ascend_path,
+ "reference_dir": _REF,
+ "mode": mode,
+ "error_logs": error_logs,
+ "target_commit": ctx.target_commit,
+ "step_index": f"{ctx.current_step + 1}/{ctx.total_steps}",
+ "llvm_project_path": str(config.llvm_project),
+ }
+ if extra:
+ base.update(extra)
+ return base
+
+
+def _run_ir_op_analysis(ctx: WorkflowContext, config: TAConfig) -> dict:
+ """AI scans Ascend backend code for MLIR OP usage. Returns ops_report."""
+ step = ctx.steps[ctx.current_step] if ctx.steps else {"id": "step-0"}
+ step_id = step["id"]
+ ir_dir = WORKSPACE_DIR / STEPS_DIR / step_id / IR_ANALYSIS_DIR
+ ir_dir.mkdir(parents=True, exist_ok=True)
+
+ result = run_opencode_adapter(_ir_ai_base(ctx, config, step_id, ir_dir,
+ "ir_op_analysis"))
+
+ ops_file = ir_dir / IR_OPS_REPORT_FILE
+ if hasattr(result, 'step_summary') and result.step_summary:
+ try:
+ ops_data = json.loads(result.step_summary)
+ ops_file.write_text(json.dumps(ops_data, indent=2, ensure_ascii=False) + "\n", encoding="utf-8")
+ return ops_data
+ except json.JSONDecodeError:
+ pass
+ return {}
+
+
+def _run_ir_change_analysis(ctx: WorkflowContext, config: TAConfig,
+ target_llvm_hash: str) -> dict:
+ """AI compares OP .td definitions between baseline and target LLVM."""
+ step = ctx.steps[ctx.current_step] if ctx.steps else {"id": "step-0"}
+ step_id = step["id"]
+ ir_dir = WORKSPACE_DIR / STEPS_DIR / step_id / IR_ANALYSIS_DIR
+ ir_dir.mkdir(parents=True, exist_ok=True)
+
+ result = run_opencode_adapter(_ir_ai_base(ctx, config, step_id, ir_dir,
+ "ir_change_analysis",
+ error_logs=json.dumps({
+ "baseline_llvm_hash": _ASCEND_BASELINE_LLVM_HASH,
+ "target_llvm_hash": target_llvm_hash,
+ }, ensure_ascii=False),
+ extra={
+ "baseline_llvm_hash": _ASCEND_BASELINE_LLVM_HASH,
+ "target_llvm_hash": target_llvm_hash,
+ },
+ ))
+
+ changes_file = ir_dir / IR_CHANGES_REPORT_FILE
+ if hasattr(result, 'step_summary') and result.step_summary:
+ try:
+ changes_data = json.loads(result.step_summary)
+ changes_file.write_text(json.dumps(changes_data, indent=2, ensure_ascii=False) + "\n", encoding="utf-8")
+ return changes_data
+ except json.JSONDecodeError:
+ pass
+ return {}
+
+
+def _run_ir_generate_patches(ctx: WorkflowContext, config: TAConfig,
+ step: dict, target_llvm_hash: str = "") -> None:
+ """AI modifies the existing patch file in-place for new LLVM version."""
+ step_id = step["id"]
+ ir_dir = WORKSPACE_DIR / STEPS_DIR / step_id / IR_ANALYSIS_DIR
+ ir_dir.mkdir(parents=True, exist_ok=True)
+
+ ascend_path = Path(ctx.triton_ascend_path)
+ patch_dir = ascend_path / "third_party/ascend/patch"
+ patch_files = sorted(patch_dir.glob("*.patch")) if patch_dir.exists() else []
+ ascend_patch_file = str(patch_files[0]) if patch_files else ""
+
+ run_opencode_adapter(_ir_ai_base(ctx, config, step_id, ir_dir,
+ "ir_patch_gen",
+ error_logs=json.dumps({
+ "ops_report": ctx.ir_ops_report,
+ "changes_report": ctx.ir_changes_report,
+ }, ensure_ascii=False, default=str),
+ extra={
+ "baseline_llvm_hash": _ASCEND_BASELINE_LLVM_HASH,
+ "target_llvm_hash": target_llvm_hash,
+ "ascend_patch_file": ascend_patch_file,
+ },
+ ))
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# AI fix helpers for IR pipeline
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _ai_adjust_patch_for_failure(
+ ctx: WorkflowContext, config: TAConfig, step: dict, error_info: str,
+) -> None:
+ """AI adjusts the LLVM patch after build failure."""
+ step_id = step["id"]
+ ir_dir = WORKSPACE_DIR / STEPS_DIR / step_id / IR_ANALYSIS_DIR
+ ir_dir.mkdir(parents=True, exist_ok=True)
+
+ ascend_path = Path(ctx.triton_ascend_path)
+ patch_dir = ascend_path / "third_party/ascend/patch"
+ patch_files = sorted(patch_dir.glob("*.patch")) if patch_dir.exists() else []
+ ascend_patch_file = str(patch_files[0]) if patch_files else ""
+ target_llvm_hash = _get_current_llvm_hash(ascend_path)
+
+ run_opencode_adapter(_ir_ai_base(ctx, config, step_id, ir_dir,
+ "ir_patch_fix",
+ error_logs=json.dumps({"error": error_info[:5000]}, ensure_ascii=False),
+ extra={
+ "baseline_llvm_hash": _ASCEND_BASELINE_LLVM_HASH,
+ "target_llvm_hash": target_llvm_hash,
+ "ascend_patch_file": ascend_patch_file,
+ "adjust_mode": "patch_apply_failure",
+ "patch_error_type": "apply_or_build",
+ "patch_error_msg": error_info[:2000],
+ },
+ ))
+
+
+def _build_focused_change_report(
+ ctx: WorkflowContext, config: TAConfig, step: dict,
+ target_llvm_hash: str,
+) -> Path | None:
+ """Analyze OP definition diffs for affected OPs identified by ir_diagnose.
+
+ Uses ``git grep`` to find the .td files defining each affected OP,
+ then runs ``git diff baseline..target`` to collect the precise changes.
+
+ Writes the focused report to ``focused_changes.json`` in the IR
+ analysis directory so the supplement AI can read it.
+
+ Returns the path to the focused changes report, or None if no
+ affected OPs could be analyzed.
+ """
+ step_id = step["id"]
+ ir_dir = WORKSPACE_DIR / STEPS_DIR / step_id / IR_ANALYSIS_DIR
+ ir_dir.mkdir(parents=True, exist_ok=True)
+
+ baseline_hash = _ASCEND_BASELINE_LLVM_HASH
+ llvm_project = config.llvm_project
+
+ # ── Read affected OPs from diagnosis ───────────────────────────
+ diagnosis = _read_diagnosis(step_id, ir_dir)
+ affected_ops: list[dict] = []
+ if isinstance(diagnosis, dict):
+ failures = diagnosis.get("failures", [])
+ for f in failures:
+ if f.get("classification") == "ir_compatibility":
+ op_name = f.get("affected_op", "").strip()
+ if op_name and op_name not in {o["name"] for o in affected_ops}:
+ affected_ops.append({
+ "name": op_name,
+ "error_summary": f.get("error_summary", ""),
+ "rationale": f.get("rationale", ""),
+ })
+ if not affected_ops:
+ log.info("No ir_compatibility OPs in diagnosis — nothing to analyze")
+ return None
+
+ log.info(f"Analyzing changes for {len(affected_ops)} affected OP(s)...")
+
+ # ── For each affected OP, find .td definition and diff ──────────
+ analyzed: list[dict] = []
+ seen_td_files: set[str] = set()
+
+ for op in affected_ops:
+ op_name = op["name"]
+ td_relative = _find_td_file(llvm_project, target_llvm_hash, op_name)
+ if not td_relative:
+ log.info(f" {op_name}: .td file not found — skipping")
+ continue
+
+ entry: dict = {
+ "op_name": op_name,
+ "td_file": td_relative,
+ "error_summary": op["error_summary"],
+ "rationale": op["rationale"],
+ }
+
+ # Only diff each .td file once (multiple OPs in same file)
+ if td_relative not in seen_td_files:
+ seen_td_files.add(td_relative)
+ diff = _diff_td_file(llvm_project, baseline_hash, target_llvm_hash, td_relative)
+ entry["td_diff"] = diff[:8000] if diff else "(no diff)"
+ if diff and len(diff) > 8000:
+ entry["td_diff_truncated"] = True
+ else:
+ entry["td_diff"] = "(see above — already included)"
+
+ analyzed.append(entry)
+ log.info(f" {op_name}: {td_relative}")
+
+ if not analyzed:
+ return None
+
+ # ── Write focused report ───────────────────────────────────────
+ report = {
+ "source_llvm_hash": baseline_hash,
+ "target_llvm_hash": target_llvm_hash,
+ "diagnosis_source": "ir_diagnose",
+ "affected_ops": analyzed,
+ "summary": {
+ "total_affected_ops": len(analyzed),
+ "unique_td_files": len(seen_td_files),
+ },
+ }
+ report_path = ir_dir / "focused_changes.json"
+ report_path.write_text(
+ json.dumps(report, indent=2, ensure_ascii=False) + "\n",
+ encoding="utf-8",
+ )
+ log.status(True, f"Focused change analysis: {report_path}")
+ return report_path
+
+
+def _read_diagnosis(step_id: str, ir_dir: Path) -> dict:
+ """Read IR diagnosis JSON, trying both possible locations."""
+ for p in _diagnosis_candidates(step_id, ir_dir):
+ if p.exists():
+ try:
+ return json.loads(p.read_text(encoding="utf-8"))
+ except (json.JSONDecodeError, OSError):
+ pass
+ return {}
+
+
+def _find_diagnosis_file(step_id: str, ir_dir: Path) -> Path | None:
+ """Find the IR diagnosis file, trying both possible locations."""
+ for p in _diagnosis_candidates(step_id, ir_dir):
+ if p.exists():
+ return p
+ return None
+
+
+def _diagnosis_candidates(step_id: str, ir_dir: Path) -> list[Path]:
+ """Candidate paths for IR diagnosis, in priority order."""
+ step_dir = WORKSPACE_DIR / STEPS_DIR / step_id
+ return [
+ step_dir / "ir_diagnosis.json", # where AI writes per prompt
+ ir_dir / IR_DIAGNOSIS_FILE, # where classify writes parsed
+ ]
+
+
+def _find_td_file(llvm_project: Path, target_hash: str, op_name: str) -> str | None:
+ """Find the .td file that defines *op_name* at *target_hash*.
+
+ Uses ``git grep`` to search for ``def OpName`` in .td files.
+ """
+ # Strip dialect prefix for the def search (e.g., "triton::LoadOp" → "LoadOp")
+ short_name = op_name.split("::")[-1]
+ try:
+ result = subprocess.run(
+ ["git", "grep", "-l", f"def {short_name}", target_hash, "--", "*.td"],
+ cwd=str(llvm_project),
+ capture_output=True, text=True, timeout=30,
+ )
+ except (subprocess.TimeoutExpired, OSError):
+ return None
+ if result.returncode != 0 or not result.stdout.strip():
+ return None
+ # Return first match (most OPs are defined once)
+ return result.stdout.strip().split("\n")[0]
+
+
+def _diff_td_file(llvm_project: Path, baseline_hash: str, target_hash: str,
+ td_relative: str) -> str:
+ """Get the diff of a .td file between baseline and target LLVM."""
+ try:
+ result = subprocess.run(
+ ["git", "diff", f"{baseline_hash}..{target_hash}", "--", td_relative],
+ cwd=str(llvm_project),
+ capture_output=True, text=True, timeout=60,
+ )
+ except (subprocess.TimeoutExpired, OSError):
+ return ""
+ return result.stdout if result.returncode == 0 else ""
+
+
+def _ir_supplement_patch(
+ ctx: WorkflowContext, config: TAConfig, step: dict, target_llvm_hash: str,
+ supplement_iter: int = 1,
+) -> None:
+ """AI supplements the existing IR patch with missing OP IR changes.
+
+ AI is given:
+ - Test failure logs showing IR errors (collected from test-logs/)
+ - The existing patch file content (first 5000 bytes as context)
+ - The target LLVM commit and llvm-project path for OP definition lookup
+ - IR diagnosis from previous step (chained via previous_step_summary_path)
+ """
+ step_id = step["id"]
+ ir_dir = WORKSPACE_DIR / STEPS_DIR / step_id / IR_ANALYSIS_DIR
+ ir_dir.mkdir(parents=True, exist_ok=True)
+
+ ascend_path = Path(ctx.triton_ascend_path)
+ patch_dir = ascend_path / "third_party/ascend/patch"
+ patch_files = sorted(patch_dir.glob("*.patch")) if patch_dir.exists() else []
+ ascend_patch_file = str(patch_files[0]) if patch_files else ""
+
+ # Collect actual test failure logs
+ error_log_paths = _collect_test_error_logs()
+ if not error_log_paths:
+ log.warning("No test failure logs found — cannot diagnose IR issues")
+
+ # Include IR diagnosis if available (chain from classify step)
+ # Try both possible locations: where AI writes per prompt, where classify writes
+ diagnosis_path = _find_diagnosis_file(step_id, ir_dir)
+ if diagnosis_path:
+ error_log_paths.append(str(diagnosis_path))
+ log.info(f"Including IR diagnosis: {diagnosis_path}")
+ else:
+ diagnosis_path = None
+
+ # ── Focused OP change analysis ──
+ # Extract affected OPs from diagnosis, diff their .td definitions
+ # between baseline and target LLVM, write to focused_changes.json
+ focused_report_path = _build_focused_change_report(
+ ctx, config, step, target_llvm_hash,
+ )
+ if focused_report_path:
+ error_log_paths.append(str(focused_report_path))
+ log.info(f"Including focused change analysis: {focused_report_path}")
+
+ # Build patch content snippet for AI context
+ patch_content_snippet = ""
+ if ascend_patch_file:
+ try:
+ full = Path(ascend_patch_file).read_text(encoding="utf-8", errors="replace")
+ patch_content_snippet = full[:5000]
+ if len(full) > 5000:
+ patch_content_snippet += f"\n\n... ({len(full) - 5000} more bytes)"
+ except Exception:
+ pass
+
+ log.key_value("Existing patch", str(ascend_patch_file) if ascend_patch_file else "(none)")
+ log.key_value("Target LLVM", target_llvm_hash[:12])
+ log.key_value("Test error logs", str(len(error_log_paths)))
+ if focused_report_path:
+ log.key_value("Focused changes", str(focused_report_path))
+ if diagnosis_path:
+ log.key_value("Diagnosis", str(diagnosis_path))
+
+ run_opencode_adapter(_ir_ai_base(ctx, config, step_id, ir_dir,
+ "ir_generate_patch",
+ error_logs=json.dumps(error_log_paths, ensure_ascii=False),
+ extra={
+ "previous_step_id": "ir-diagnose",
+ "previous_step_summary_path": str(diagnosis_path or ""),
+ "focused_changes_path": str(focused_report_path or ""),
+ "target_llvm_hash": target_llvm_hash,
+ "baseline_llvm_hash": _ASCEND_BASELINE_LLVM_HASH,
+ "ascend_patch_file": ascend_patch_file,
+ "patch_content_snippet": patch_content_snippet,
+ "adjust_mode": "supplement",
+ "supplement_iteration": str(supplement_iter),
+ "ascend_npu_ir_compat_ref": str(
+ Path(__file__).parent.parent / "reference"
+ / "AscendNPU-IR_LLVM_VERSION_COMPAT.md"),
+ },
+ ))
+
+
+def _collect_test_error_logs() -> list[str]:
+ """Collect actual test failure log files for AI analysis.
+
+ Gathers all per-suite JUnit XMLs, result JSONs, and custom test
+ output logs so the AI has the full failure picture across all
+ sequentially-run test suites.
+ """
+ error_log_paths: list[str] = []
+ test_log_dir = WORKSPACE_DIR / "test-logs"
+ if test_log_dir.exists():
+ # JUnit XML per suite (pytest-junit-primary.xml, pytest-junit-extra-*.xml)
+ for f in sorted(test_log_dir.glob("pytest-junit-*.xml")):
+ error_log_paths.append(str(f))
+ # Result JSON per suite (test-result-primary.json, etc.)
+ for f in sorted(test_log_dir.glob("test-result-*.json")):
+ error_log_paths.append(str(f))
+ # Custom test command output
+ for f in sorted(test_log_dir.glob("*.log")):
+ error_log_paths.append(str(f))
+ # Legacy single-result file (backward compat)
+ test_result = WORKSPACE_DIR / TEST_RESULT_FILE
+ if test_result.exists():
+ error_log_paths.append(str(test_result))
+ return error_log_paths
+
+
+def _classify_test_failures(
+ ctx: WorkflowContext, config: TAConfig, step: dict,
+) -> bool:
+ """AI classifies test failures: IR compatibility vs code issues.
+
+ Collects actual test log files, invokes AI diagnosis, and writes
+ the result to ``IR_DIAGNOSIS_FILE`` so downstream supplement
+ steps can chain from it.
+
+ Returns True if IR issues are present (needs supplement),
+ False if purely code issues (needs AI fix).
+ """
+ step_id = step["id"]
+ ir_dir = WORKSPACE_DIR / STEPS_DIR / step_id / IR_ANALYSIS_DIR
+ ir_dir.mkdir(parents=True, exist_ok=True)
+
+ # Collect actual test failure logs (not just fix_errors paths)
+ error_log_paths = _collect_test_error_logs()
+ if not error_log_paths:
+ log.warning("No test failure logs found — assuming IR issues")
+ return True
+
+ log.info(f"Collected {len(error_log_paths)} log file(s) for AI diagnosis")
+ for p in error_log_paths[:5]:
+ log.info(f" - {p}")
+ if len(error_log_paths) > 5:
+ log.info(f" ... and {len(error_log_paths) - 5} more")
+
+ try:
+ result = run_opencode_adapter(_ir_ai_base(ctx, config, step_id, ir_dir,
+ "ir_diagnose",
+ error_logs=json.dumps(error_log_paths, ensure_ascii=False),
+ ))
+ except Exception as e:
+ log.error(f"IR diagnosis failed: {e}")
+ return True # Default to IR issue on failure
+
+ # Write diagnosis result for chaining
+ diagnosis_path = ir_dir / IR_DIAGNOSIS_FILE
+ summary = result.step_summary or ""
+ try:
+ diagnosis_data = json.loads(summary)
+ except json.JSONDecodeError:
+ diagnosis_data = {"summary": summary, "has_ir_issues": "ir_issue" in summary.lower()}
+ diagnosis_path.write_text(
+ json.dumps(diagnosis_data, indent=2, ensure_ascii=False) + "\n",
+ encoding="utf-8",
+ )
+ log.info(f"IR diagnosis written: {diagnosis_path}")
+
+ has_ir = diagnosis_data.get("has_ir_issues", False) or "ir_issue" in summary.lower()
+ return bool(has_ir)
+
+
+def _do_ai_fix_loop(
+ ctx: WorkflowContext, config: TAConfig, step: dict,
+) -> WorkflowContext:
+ """Standard AI fix loop for code issues (not IR-related)."""
+ return test_and_fix_loop(ctx, config)
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# LLVM workspace management
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _ensure_llvm_workspace_clean(llvm_project: Path, reason: str = "") -> None:
+ """Clean the LLVM working tree: stash changes, checkout HEAD, clean.
+
+ Idempotent — safe to call multiple times.
+ """
+ if not llvm_project.exists():
+ return
+
+ log.info(f"Cleaning LLVM workspace{f' ({reason})' if reason else ''}...")
+ try:
+ # Abort any in-progress merge
+ if (llvm_project / ".git" / "MERGE_HEAD").exists():
+ run_git(llvm_project, "merge", "--abort")
+ except Exception:
+ pass
+
+ try:
+ # Stash any local changes
+ run_git(llvm_project, "stash", "--include-untracked")
+ run_git(llvm_project, "stash", "drop")
+ except Exception:
+ # If stash fails (no changes), just reset
+ try:
+ run_git(llvm_project, "checkout", "--", ".")
+ run_git(llvm_project, "clean", "-fd")
+ except Exception:
+ pass
+
+
+def _ensure_commit_available(llvm_project: Path, commit_hash: str) -> None:
+ """Ensure *commit_hash* is available in the local LLVM repo.
+
+ Fetches from origin with retries if the commit is missing.
+ """
+ max_attempts = 6
+ for attempt in range(1, max_attempts + 1):
+ result = run_git_no_check(llvm_project, "cat-file", "-t", commit_hash)
+ if result.returncode == 0:
+ return
+ if attempt < max_attempts:
+ log.info(f"Commit {commit_hash[:12]} not found locally — fetching "
+ f"(attempt {attempt}/{max_attempts})...")
+ try:
+ run_git(llvm_project, "fetch", "origin", commit_hash)
+ except Exception:
+ time.sleep(2)
+ raise RuntimeError(
+ f"Failed to fetch LLVM commit {commit_hash[:12]} after {max_attempts} attempts"
+ )
+
+
+def _clean_checkout_apply_patch(
+ llvm_project: Path, ascend_path: Path, target_llvm_hash: str, reason: str = "",
+) -> bool:
+ """Clean workspace, checkout target hash, apply the current Ascend patch.
+
+ Returns True if a patch file was found and applied, False if no patch
+ exists (clean LLVM, no Ascend modifications).
+ """
+ _ensure_llvm_workspace_clean(llvm_project, reason or "prepare for patch apply")
+ _ensure_commit_available(llvm_project, target_llvm_hash)
+ run_git(llvm_project, "checkout", "-f", target_llvm_hash)
+ run_git(llvm_project, "clean", "-fd")
+
+ patch_dir = ascend_path / "third_party/ascend/patch"
+ patch_files = sorted(patch_dir.glob("*.patch")) if patch_dir.exists() else []
+ if not patch_files:
+ log.info("No Ascend patch found — building clean LLVM")
+ return False
+
+ patch_file = patch_files[0]
+ log.info(f"Applying patch: {patch_file.name}")
+ run_git(llvm_project, "apply", str(patch_file))
+ log.status(True, f"Patch applied: {patch_file.name}")
+ return True
+
+
+def _get_current_llvm_hash(ascend_path: Path) -> str:
+ """Read the current LLVM hash from triton-ascend's cmake/llvm-hash.txt."""
+ hash_file = ascend_path / "cmake" / "llvm-hash.txt"
+ if hash_file.exists():
+ return hash_file.read_text(encoding="utf-8").strip()
+ return ""
+
+
+
+def _do_llvm_build(llvm_project: Path, llvm_install: Path,
+ required_hash: str) -> str:
+ """Build and install LLVM from current working tree state.
+
+ Cleans build directory, runs cmake + ninja install, copies FileCheck.
+
+ Returns the LLVM install prefix path.
+
+ Raises RuntimeError on build failure.
+ """
+ build_dir = WORKSPACE_DIR / "llvm-build"
+ if build_dir.exists():
+ shutil.rmtree(build_dir)
+ build_dir.mkdir(parents=True, exist_ok=True)
+
+ log_dir = WORKSPACE_DIR / "llvm-logs"
+ log_dir.mkdir(parents=True, exist_ok=True)
+
+ log.info(f"Building LLVM (hash: {required_hash[:12]})...")
+ log.key_value("LLVM project", str(llvm_project))
+ log.key_value("LLVM install prefix", str(llvm_install))
+ log.key_value("Build log", str(log_dir / "llvm-build.log"))
+
+ # ── cmake configure ──────────────────────────────────────────────
+ cmake_cmd = [
+ "cmake", str(llvm_project / "llvm"),
+ "-G", "Ninja",
+ "-DCMAKE_BUILD_TYPE=Release",
+ "-DLLVM_ENABLE_ASSERTIONS=ON",
+ "-DLLVM_ENABLE_PROJECTS=mlir;llvm;lld",
+ "-DLLVM_TARGETS_TO_BUILD=host;NVPTX;AMDGPU",
+ f"-DCMAKE_INSTALL_PREFIX={llvm_install}",
+ "-DCMAKE_C_COMPILER=clang",
+ "-DCMAKE_CXX_COMPILER=clang++",
+ ]
+ llvm_build_log = log_dir / "llvm-build.log"
+ with open(llvm_build_log, "w", encoding="utf-8") as fh:
+ fh.write(f"=== cmake ===\n{' '.join(cmake_cmd)}\n\n")
+ fh.flush()
+ rc = stream_cmd(cmake_cmd, build_dir, fh, timeout=300,
+ label="Configuring LLVM with cmake")
+ if rc != 0:
+ raise RuntimeError(
+ f"LLVM cmake configure failed (exit {rc}). See: {llvm_build_log}"
+ )
+ log.status(True, "cmake configure OK")
+
+ # ── ninja build + install ───────────────────────────────────────
+ log.info("ninja install (this may take a while)...")
+ with open(llvm_build_log, "a", encoding="utf-8") as fh:
+ fh.write(f"\n=== ninja install ===\n")
+ fh.flush()
+ rc = stream_cmd(["ninja", "install"], build_dir, fh, timeout=7200,
+ label="ninja install")
+ if rc != 0:
+ raise RuntimeError(
+ f"LLVM ninja build failed (exit {rc}). See: {llvm_build_log}"
+ )
+ log.status(True, "ninja install OK")
+
+ # Copy FileCheck
+ fc_src = build_dir / "bin" / "FileCheck"
+ fc_dst = llvm_install / "bin" / "FileCheck"
+ if fc_src.exists():
+ fc_dst.parent.mkdir(parents=True, exist_ok=True)
+ shutil.copy2(fc_src, fc_dst)
+
+ # Write hash cache
+ llvm_install.mkdir(parents=True, exist_ok=True)
+ (llvm_install / ".llvm_hash").write_text(required_hash, encoding="utf-8")
+
+ log.status(True, f"LLVM build complete — {llvm_install}")
+ return str(llvm_install)
+
+
+def _detect_ascend_npu_ir_errors(ctx: WorkflowContext) -> bool:
+ """Check if build errors are from AscendNPU-IR compilation failures."""
+ step = ctx.steps[ctx.current_step] if ctx.steps else {"id": "step-0"}
+ step_dir = WORKSPACE_DIR / STEPS_DIR / step["id"]
+ build_log = step_dir / "build.log"
+ if not build_log.exists():
+ return False
+ try:
+ content = build_log.read_text(encoding="utf-8", errors="replace").lower()
+ indicators = [
+ "AscendNPU-IR".lower(), "ascendnpu-ir",
+ "llvm::", "mlir::",
+ "fatal error", "undefined reference",
+ ]
+ return any(ind in content for ind in indicators)
+ except Exception:
+ return False
diff --git a/src/TA_main2main_workflow/pipeline/merge.py b/src/TA_main2main_workflow/pipeline/merge.py
new file mode 100644
index 0000000..71f658e
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/merge.py
@@ -0,0 +1,149 @@
+"""Pipeline step: Execute git merge of upstream commits into triton-ascend.
+
+For the first step, creates a work branch. Subsequent steps merge
+incrementally on the same work branch.
+"""
+
+from __future__ import annotations
+
+import json
+import time
+from pathlib import Path
+
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.git import run_git, run_git_no_check
+from TA_main2main_workflow.utils import (
+ WORKSPACE_DIR, STEPS_DIR, get_base_branch_ref,
+)
+
+log = get_logger(__name__)
+
+_MERGE_RESULT = "merge_result.json"
+
+
+def merge_upstream_commit(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Merge this step's upstream commits.
+
+ On the first step (current_step == 0): creates a new work branch.
+ On subsequent steps: merges incrementally on the existing work branch.
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ step = (
+ ctx.steps[ctx.current_step]
+ if ctx.steps
+ else {"id": "step-0", "end_commit": ctx.target_commit}
+ )
+ step_id = step.get("id", "step-0")
+ step_dir = WORKSPACE_DIR / STEPS_DIR / step_id
+ step_dir.mkdir(parents=True, exist_ok=True)
+ result_file = step_dir / _MERGE_RESULT
+
+ # Resume: skip if merge_result.json already exists for this step
+ if config.resume and result_file.exists():
+ log.info(f"Resume: {_MERGE_RESULT} exists, skipping merge")
+ mr = json.loads(result_file.read_text(encoding="utf-8"))
+ return ctx.copy_with(
+ merge_has_conflicts=mr.get("has_conflicts", False),
+ conflict_files=mr.get("conflict_files", []),
+ )
+
+ # ── Step 0: Create work branch ──────────────────────────────────
+ if ctx.current_step == 0:
+ _create_work_branch(ascend_path, config)
+ work_branch = run_git(ascend_path, "branch", "--show-current").strip()
+ ctx = ctx.copy_with(work_branch=work_branch)
+ log.info(f"Work branch: {work_branch}")
+ else:
+ # Ensure we're on the work branch
+ current_branch = run_git(ascend_path, "branch", "--show-current").strip()
+ if ctx.work_branch and current_branch != ctx.work_branch:
+ log.warning(f"Expected '{ctx.work_branch}' but on '{current_branch}' — switching")
+ run_git(ascend_path, "checkout", ctx.work_branch)
+
+ # ── Do the merge ────────────────────────────────────────────────
+ log.info(f"Merging {step['end_commit'][:12]} ...")
+ merge_proc = run_git_no_check(
+ ascend_path, "merge", "--no-ff", "--no-edit", step["end_commit"]
+ )
+
+ conflict_files = run_git(
+ ascend_path, "diff", "--name-only", "--diff-filter=U"
+ ).strip()
+ conflict_files = (
+ [f for f in conflict_files.splitlines() if f] if conflict_files else []
+ )
+ has_conflicts = len(conflict_files) > 0
+
+ result = {
+ "target_commit": step["end_commit"],
+ "merge_exit_code": merge_proc.returncode,
+ "has_conflicts": has_conflicts,
+ "conflict_files": conflict_files,
+ }
+ result_file.write_text(
+ json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
+ )
+
+ if has_conflicts:
+ log.conflict_list(conflict_files)
+ elif merge_proc.returncode != 0:
+ log.warning(f"Merge exited with code {merge_proc.returncode} "
+ f"but no conflict markers found — continuing")
+ else:
+ log.key_value("merge exit code", str(merge_proc.returncode))
+ log.key_value("conflict files", "0")
+
+ return ctx.copy_with(
+ merge_has_conflicts=has_conflicts, conflict_files=conflict_files
+ )
+
+
+def _create_work_branch(repo: Path, config: TAConfig) -> None:
+ """Create a work branch from the configured base ref.
+
+ Cleans the working tree before branching: aborts stale merges,
+ resets to a pristine state, then creates the work branch.
+ """
+ base_ref = get_base_branch_ref(config.work_branch_base)
+ timestamp = time.strftime("%Y%m%d-%H%M%S")
+ branch_name = f"sync/main2main-{timestamp}"
+
+ # ── 1. Abort any stale merge ────────────────────────────────────
+ if (repo / ".git" / "MERGE_HEAD").exists():
+ log.warning("Stale merge in progress — aborting")
+ try:
+ run_git(repo, "merge", "--abort")
+ except Exception:
+ run_git(repo, "reset", "--hard", "HEAD")
+ # Clean up leftover merge files
+ for fname in ("MERGE_MODE", "MERGE_MSG", "CHERRY_PICK_HEAD"):
+ p = repo / ".git" / fname
+ if p.exists():
+ p.unlink()
+
+ # ── 2. Reset to pristine state ─────────────────────────────────
+ try:
+ run_git(repo, "checkout", "--detach")
+ except Exception:
+ pass
+ run_git(repo, "reset", "--hard", "HEAD")
+ run_git(repo, "clean", "-fd")
+
+ # ── 3. Fetch the work branch base remote ────────────────────────
+ try:
+ run_git(repo, "fetch", config.work_branch_base)
+ except Exception:
+ log.warning(f"Could not fetch remote '{config.work_branch_base}' — using origin")
+
+ # ── 4. Checkout base ref and create work branch ────────────────
+ try:
+ run_git(repo, "checkout", "-B", branch_name, base_ref)
+ except Exception:
+ # Fallback: use origin/base_branch
+ fallback = f"origin/{config.base_branch}"
+ log.warning(f"Could not use {base_ref}, falling back to {fallback}")
+ run_git(repo, "checkout", "-B", branch_name, fallback)
+
+ log.info(f"Created work branch: {branch_name}")
diff --git a/src/TA_main2main_workflow/pipeline/plan.py b/src/TA_main2main_workflow/pipeline/plan.py
new file mode 100644
index 0000000..f37bed6
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/plan.py
@@ -0,0 +1,334 @@
+"""Pipeline step: Plan steps — split upstream commits by line budget.
+
+Groups upstream commits into ordered steps based on changed lines in
+key source directories. LLVM-hash-changing commits get solo steps.
+"""
+
+from __future__ import annotations
+
+import json
+from pathlib import Path
+from typing import Any
+
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils import (
+ WORKSPACE_DIR, STEPS_FILE, STEPS_DIR, LLVM_HASH_FILE, SOURCE_DIRS,
+)
+from TA_main2main_workflow.utils.git import run_git
+from TA_main2main_workflow.utils.logging import get_logger
+
+log = get_logger(__name__)
+
+
+def run_plan(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Plan merge steps (with resume support).
+
+ If ``config.resume`` is set and ``steps.json`` exists, loads cached
+ results. Otherwise runs full planning via :func:`plan_steps`.
+ """
+ steps_file = WORKSPACE_DIR / STEPS_FILE
+ if config.resume and steps_file.exists():
+ log.info("Resume: steps.json exists, skipping plan")
+ plan = json.loads(steps_file.read_text(encoding="utf-8"))
+ return ctx.copy_with(steps=plan["steps"], total_steps=len(plan["steps"]))
+
+ return plan_steps(ctx, config)
+
+
+def plan_steps(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Split upstream commits into steps and populate ctx.steps.
+
+ If progressive_merge is disabled or there is only 1 commit,
+ creates a single step covering all commits.
+ """
+ triton_path = Path(ctx.triton_ascend_path)
+ base = ctx.merge_base
+ target = ctx.target_commit
+ line_budget = config.line_budget
+
+ commits = ctx.upstream_commits
+ log.info(
+ f"[plan] Scanning {len(commits)} upstream commits ({base[:8]}..{target[:8]})"
+ )
+ log.info(f"[plan] Line budget: {line_budget} (no commit-count limit)")
+
+ if config.progressive_merge:
+ lines_per_commit, llvm_commits, source_touching = _scan_commits(triton_path, commits)
+ steps = _plan_steps_inner(
+ commits, lines_per_commit, base, line_budget, llvm_commits
+ )
+ _enrich_steps(triton_path, steps)
+
+ # Print per-step detail (matching pre-refactor output)
+ for s in steps:
+ reason_tag = ""
+ if s.get("reason") == "llvm_version":
+ reason_tag = " [LLVM VERSION]"
+ elif s.get("reason") == "oversized":
+ reason_tag = " [OVERSIZED]"
+ budget_label = "OVERSIZED" if s["source_changed_lines"] > line_budget else "OK"
+ log.info(
+ f" {s['id']}: {s['commit_count']} commits, "
+ f"{s['source_changed_lines']} lines ({budget_label})"
+ f"{reason_tag}"
+ )
+
+ total_commits = sum(s["commit_count"] for s in steps)
+ plan = {
+ "base_commit": base,
+ "target_commit": target,
+ "line_budget": line_budget,
+ "total_source_commits": source_touching,
+ "total_commits": total_commits,
+ "total_steps": len(steps),
+ "steps": steps,
+ }
+ _write_plan(plan)
+
+ log.info(
+ f"[plan] Generated {len(steps)} step(s) totaling "
+ f"{source_touching} source-touching commits"
+ )
+ return ctx.copy_with(steps=steps, total_steps=len(steps))
+ else:
+ return ctx.copy_with(
+ total_steps=1,
+ steps=[
+ {
+ "index": 1,
+ "id": "step-1",
+ "commits": commits,
+ "commit_count": len(commits),
+ "start_commit": base,
+ "end_commit": target,
+ "source_changed_lines": ctx.changed_lines_total,
+ "reason": "single_step",
+ }
+ ],
+ )
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# LLVM hash change detection helpers
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def llvm_hash_changed_after_merge(ctx: WorkflowContext) -> bool:
+ """Check if the merge changed cmake/llvm-hash.txt in triton-ascend.
+
+ Compares the current hash with the pre-step ascend HEAD.
+ Should be called AFTER a merge step completes.
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ hash_file = ascend_path / LLVM_HASH_FILE
+ if not hash_file.exists():
+ return False
+ current_hash = hash_file.read_text(encoding="utf-8").strip()
+
+ # Compare with pre-step state
+ if ctx.step_start_ascend_head:
+ try:
+ old_content = run_git(
+ ascend_path, "show",
+ f"{ctx.step_start_ascend_head}:{LLVM_HASH_FILE}"
+ ).strip()
+ return old_content != current_hash
+ except Exception:
+ pass
+ return False
+
+
+def get_current_llvm_hash(ascend_path: Path) -> str:
+ """Read the current LLVM hash from triton-ascend's cmake/llvm-hash.txt."""
+ hash_file = ascend_path / LLVM_HASH_FILE
+ if hash_file.exists():
+ return hash_file.read_text(encoding="utf-8").strip()
+ return ""
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Internal helpers
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _source_lines_for_commit(repo: Path, sha: str) -> int:
+ """Return total lines changed in SOURCE_DIRS only, per directory sum."""
+ total = 0
+ for d in SOURCE_DIRS:
+ try:
+ output = run_git(
+ repo, "diff-tree", "--no-commit-id", "-r", "--numstat",
+ sha, "--", f":(top){d}",
+ )
+ except Exception:
+ continue
+ for line in output.strip().splitlines():
+ if not line.strip():
+ continue
+ parts = line.split("\t")
+ if len(parts) >= 2:
+ try:
+ total += int(parts[0] or 0) + int(parts[1] or 0)
+ except ValueError:
+ pass
+ return total
+
+
+def _commit_changed_llvm_hash(repo: Path, sha: str) -> bool:
+ try:
+ output = run_git(
+ repo, "diff-tree", "--no-commit-id", "--name-only", "-r", sha)
+ return LLVM_HASH_FILE in output
+ except Exception:
+ return False
+
+
+def _scan_commits(
+ repo: Path, commits: list[dict[str, str]]
+) -> tuple[dict[str, int], set[str], int]:
+ lines_per_commit: dict[str, int] = {}
+ llvm_commits: set[str] = set()
+ source_touching = 0
+ for i, c in enumerate(commits):
+ lines = _source_lines_for_commit(repo, c["sha"])
+ lines_per_commit[c["sha"]] = lines
+ if lines > 0:
+ source_touching += 1
+ if _commit_changed_llvm_hash(repo, c["sha"]):
+ llvm_commits.add(c["sha"])
+ log.info(f"[plan] LLVM version change detected: {c['sha'][:8]} {c['subject'][:80]}")
+ if (i + 1) % 50 == 0:
+ log.info(f"[plan] ... scanned {i + 1}/{len(commits)} commits")
+ # Print zero-lines summary (matching pre-refactor)
+ if source_touching < len(commits):
+ log.info(f"[plan] {len(commits) - source_touching} commits touch zero "
+ f"source lines — included in steps with 0 line contribution")
+ if llvm_commits:
+ log.info(f"[plan] {len(llvm_commits)} commit(s) changed LLVM hash "
+ f"— each will be a solo merge step")
+ return lines_per_commit, llvm_commits, source_touching
+
+
+def _make_step(
+ index: int,
+ commits: list[dict[str, str]],
+ start: str,
+ lines: int,
+ budget: int,
+ reason: str = "line_budget",
+) -> dict[str, Any]:
+ return {
+ "index": index,
+ "id": f"step-{index}",
+ "commits": commits,
+ "commit_count": len(commits),
+ "start_commit": start,
+ "end_commit": commits[-1]["sha"],
+ "source_changed_lines": lines,
+ "line_budget": budget,
+ "reason": reason,
+ }
+
+
+def _plan_steps_inner(
+ commits: list[dict[str, str]],
+ lines_per_commit: dict[str, int],
+ base: str,
+ budget: int,
+ llvm_commits: set[str],
+) -> list[dict[str, Any]]:
+ steps: list[dict[str, Any]] = []
+ step_commits: list[dict[str, str]] = []
+ step_lines = 0
+ start = base
+
+ for commit in commits:
+ sha = commit["sha"]
+ lines = lines_per_commit.get(sha, 0)
+
+ # LLVM change → solo step
+ if sha in llvm_commits:
+ if step_commits:
+ steps.append(
+ _make_step(len(steps) + 1, step_commits, start, step_lines, budget)
+ )
+ start = steps[-1]["end_commit"]
+ step_commits, step_lines = [], 0
+ steps.append(
+ _make_step(
+ len(steps) + 1, [commit], start, lines, budget, reason="llvm_version"
+ )
+ )
+ start = steps[-1]["end_commit"]
+ continue
+
+ # Oversized → solo step
+ if lines > budget:
+ if step_commits:
+ steps.append(
+ _make_step(len(steps) + 1, step_commits, start, step_lines, budget)
+ )
+ start = steps[-1]["end_commit"]
+ step_commits, step_lines = [], 0
+ steps.append(
+ _make_step(len(steps) + 1, [commit], start, lines, budget, reason="oversized")
+ )
+ start = steps[-1]["end_commit"]
+ continue
+
+ # Would exceed budget → flush
+ if step_lines + lines > budget:
+ steps.append(
+ _make_step(len(steps) + 1, step_commits, start, step_lines, budget)
+ )
+ start = steps[-1]["end_commit"]
+ step_commits, step_lines = [], 0
+
+ step_commits.append(commit)
+ step_lines += lines
+
+ if step_commits:
+ steps.append(
+ _make_step(len(steps) + 1, step_commits, start, step_lines, budget)
+ )
+
+ return steps
+
+
+def _enrich_steps(repo: Path, steps: list[dict[str, Any]]) -> None:
+ for step in steps:
+ # Build pathspec args for SOURCE_DIRS filtering
+ pathspecs = [f":(top){d}" for d in SOURCE_DIRS]
+ step["upstream_patch"] = run_git(
+ repo, "diff",
+ f"{step['start_commit']}..{step['end_commit']}",
+ "--", *pathspecs,
+ )
+ step["changed_files"] = run_git(
+ repo, "diff", "--name-only",
+ f"{step['start_commit']}..{step['end_commit']}",
+ "--", *pathspecs,
+ )
+ step["files_changed"] = sorted(
+ f for f in step["changed_files"].strip().splitlines() if f
+ )
+
+
+def _write_plan(plan: dict[str, Any]) -> None:
+ steps_dir = WORKSPACE_DIR / STEPS_DIR
+ steps_dir.mkdir(parents=True, exist_ok=True)
+ (WORKSPACE_DIR / STEPS_FILE).write_text(
+ json.dumps(plan, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
+ )
+ for step in plan["steps"]:
+ step_dir = steps_dir / step["id"]
+ step_dir.mkdir(parents=True, exist_ok=True)
+ (step_dir / "upstream.patch").write_text(
+ step["upstream_patch"], encoding="utf-8"
+ )
+ (step_dir / "changed_files.txt").write_text(
+ step["changed_files"], encoding="utf-8"
+ )
+ lines = [f"{c['sha'][:8]} {c['subject']}" for c in step["commits"]]
+ (step_dir / "commits.txt").write_text("\n".join(lines) + "\n", encoding="utf-8")
diff --git a/src/TA_main2main_workflow/pipeline/pre_ci.py b/src/TA_main2main_workflow/pipeline/pre_ci.py
new file mode 100644
index 0000000..3a0fd12
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/pre_ci.py
@@ -0,0 +1,101 @@
+"""Pre-CI verification checks.
+
+Runs before committing: scans for leftover merge conflict markers,
+validates Python syntax, and removes temporary build/test artifacts.
+"""
+
+from __future__ import annotations
+
+import ast
+import shutil
+from pathlib import Path
+
+from TA_main2main_workflow.utils.logging import get_logger
+
+log = get_logger(__name__)
+
+_CONFLICT_MARKERS = (b"<<<<<<<", b"=======", b">>>>>>>")
+_TEMP_PATTERNS = [
+ "result_profiling",
+ "__pycache__",
+ ".pytest_cache",
+ "*.pyc",
+ "*.pyo",
+ "*.orig",
+ "*.rej",
+ "*.log",
+ ".DS_Store",
+ "*.lock",
+]
+
+
+def run_pre_ci_check(repo: str | Path, step_id: str = "") -> bool:
+ """Run pre-CI checks on *repo* and return True if all pass."""
+ repo = Path(repo)
+ ok = _check_conflict_markers(repo)
+ if not ok:
+ log.error(f"Pre-CI [{step_id}]: conflict markers found!")
+ py_ok = _check_python_syntax(repo)
+ if not py_ok:
+ log.error(f"Pre-CI [{step_id}]: Python syntax errors found!")
+ return ok and py_ok
+
+
+def cleanup_temp_files(repo: str | Path) -> None:
+ """Remove temporary/build artifacts from *repo*."""
+ repo = Path(repo)
+ for pattern in _TEMP_PATTERNS:
+ if "*" in pattern:
+ ext = pattern.lstrip("*")
+ for f in repo.rglob(ext):
+ if f.is_file():
+ try:
+ f.unlink()
+ except OSError:
+ pass
+ else:
+ for d in repo.rglob(pattern):
+ if d.is_dir():
+ try:
+ shutil.rmtree(d)
+ except OSError:
+ pass
+ log.info("Temp files cleaned")
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Internal
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _check_conflict_markers(repo: Path) -> bool:
+ """Scan tracked source files for unresolved merge conflict markers."""
+ dirty = False
+ for ext in (".py", ".cpp", ".c", ".h", ".hpp", ".td", ".mlir", ".txt", ".md"):
+ for f in repo.rglob(f"*{ext}"):
+ if ".git" in f.parts:
+ continue
+ try:
+ content = f.read_bytes()
+ if any(m in content for m in _CONFLICT_MARKERS):
+ log.warning(f"Conflict marker in: {f}")
+ dirty = True
+ except OSError:
+ pass
+ return not dirty
+
+
+def _check_python_syntax(repo: Path) -> bool:
+ """Validate Python syntax for all .py files in the repo."""
+ ok = True
+ for py_file in repo.rglob("*.py"):
+ if ".git" in py_file.parts or "__pycache__" in py_file.parts:
+ continue
+ try:
+ ast.parse(py_file.read_text(encoding="utf-8"), filename=str(py_file))
+ except SyntaxError as e:
+ log.warning(f"Syntax error in {py_file}: {e}")
+ ok = False
+ except Exception:
+ pass
+ return ok
diff --git a/src/TA_main2main_workflow/pipeline/prepare.py b/src/TA_main2main_workflow/pipeline/prepare.py
new file mode 100644
index 0000000..dc1a0f5
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/prepare.py
@@ -0,0 +1,163 @@
+"""Pipeline step 0: Prepare workspace — clone repos, configure remotes, fetch.
+
+This is the first step of the workflow. It ensures the local environment
+is ready before any detection or merge work begins:
+
+1. Clone triton-ascend if no local path is given (skip if already exists)
+2. Verify ``origin`` points to the correct remote URL; fix if not
+3. Ensure ``triton-upstream`` remote exists, pointing to upstream Triton
+4. Fetch both remotes (with built-in retry)
+5. Checkout the configured base branch, fast-forward to origin
+
+Output context fields set:
+ - ``origin_remote``, ``upstream_remote`` — remote names
+ - ``triton_ascend_path`` — absolute path to the triton-ascend repo
+ - ``target_commit`` — the upstream commit to sync to (HEAD of
+ triton-upstream/main when not explicitly given)
+ - ``ascend_head`` — the HEAD of the configured base branch
+ - ``original_branch`` — the base branch name
+"""
+
+from __future__ import annotations
+
+from pathlib import Path
+
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.git import run_git, run_git_no_check
+from TA_main2main_workflow.utils import WORKSPACE_DIR
+
+log = get_logger(__name__)
+
+ORIGIN_REMOTE = "origin"
+UPSTREAM_REMOTE = "triton-upstream"
+
+
+def prepare(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Set up workspace: clone, remotes, fetch, checkout.
+
+ This is idempotent — safe to call on an already-prepared workspace.
+ """
+ WORKSPACE_DIR.mkdir(parents=True, exist_ok=True)
+
+ # ── 1. Ensure triton-ascend exists ──────────────────────────────────
+ ascend_path = _ensure_repo(config, WORKSPACE_DIR)
+
+ # ── 2. Ensure origin points to the correct remote URL ───────────────
+ _fix_origin(ascend_path, config.triton_ascend_url)
+
+ # ── 3. Ensure triton-upstream remote ────────────────────────────────
+ _ensure_remote(ascend_path, UPSTREAM_REMOTE, config.triton_upstream_url)
+
+ # ── 4. Fetch both remotes ───────────────────────────────────────────
+ log.info(f"Fetching {ORIGIN_REMOTE} ...")
+ run_git(ascend_path, "fetch", ORIGIN_REMOTE)
+ log.info(f"Fetching {UPSTREAM_REMOTE} ...")
+ run_git(ascend_path, "fetch", UPSTREAM_REMOTE)
+
+ # ── 5. Checkout base branch ─────────────────────────────────────────
+ base_branch = config.base_branch
+ base_ref = f"{ORIGIN_REMOTE}/{base_branch}"
+
+ # Abort any stale merge before checkout
+ if (ascend_path / ".git" / "MERGE_HEAD").exists():
+ try:
+ run_git(ascend_path, "merge", "--abort")
+ except Exception:
+ run_git(ascend_path, "reset", "--hard", "HEAD")
+
+ # Force checkout to origin's version of the base branch
+ run_git(ascend_path, "checkout", "-B", base_branch, base_ref)
+
+ # ── 6. Resolve ascend HEAD ──────────────────────────────────────────
+ try:
+ ascend_head = run_git(ascend_path, "rev-parse", base_ref).strip()
+ except Exception:
+ raise RuntimeError(
+ f"Cannot resolve '{base_ref}'. "
+ f"Fetch it first:\n"
+ f" cd {ascend_path} && git fetch {ORIGIN_REMOTE} {base_branch}"
+ )
+
+ # ── 7. Resolve target commit (default: triton-upstream/main HEAD) ──
+ target_commit = config.target_commit
+ if not target_commit:
+ upstream_ref = f"{UPSTREAM_REMOTE}/main"
+ try:
+ target_commit = run_git(ascend_path, "rev-parse", upstream_ref).strip()
+ except Exception:
+ raise RuntimeError(
+ f"Cannot resolve upstream HEAD from '{upstream_ref}'. "
+ f"Specify --target-commit or ensure '{upstream_ref}' exists. "
+ f"Try: cd {ascend_path} && git fetch {UPSTREAM_REMOTE}"
+ )
+
+ log.section("Workspace ready")
+ log.key_value("triton-ascend", str(ascend_path))
+ log.key_value("base branch", base_branch)
+ log.key_value("ascend HEAD", ascend_head[:12])
+ log.key_value("target commit", target_commit[:12])
+
+ return ctx.copy_with(
+ triton_ascend_path=str(ascend_path),
+ target_commit=target_commit,
+ ascend_head=ascend_head,
+ original_branch=base_branch,
+ origin_remote=ORIGIN_REMOTE,
+ upstream_remote=UPSTREAM_REMOTE,
+ )
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Internal helpers
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _ensure_repo(config: TAConfig, workspace: Path) -> Path:
+ """Return path to triton-ascend repo, cloning if necessary."""
+ if config.triton_ascend_path:
+ path = Path(config.triton_ascend_path)
+ if not path.exists():
+ raise FileNotFoundError(f"triton-ascend path does not exist: {path}")
+ log.info(f"Using existing repo: {path}")
+ return path
+
+ target = workspace / "triton-ascend"
+ if target.exists():
+ log.info(f"Repo exists, skip clone: {target}")
+ else:
+ log.info(f"Cloning {config.triton_ascend_url} → {target}")
+ run_git(workspace, "clone", config.triton_ascend_url, str(target))
+ return target
+
+
+def _ensure_remote(repo: Path, name: str, url: str) -> None:
+ """Add a git remote if it doesn't already exist."""
+ result = run_git_no_check(repo, "remote")
+ if name not in result.stdout:
+ run_git(repo, "remote", "add", name, url)
+
+
+def _fix_origin(repo: Path, expected_url: str) -> None:
+ """Ensure ``origin`` points to *expected_url*. Update it if not."""
+ current = _get_remote_url(repo, ORIGIN_REMOTE)
+ if current is None:
+ log.warning(f"No '{ORIGIN_REMOTE}' remote found — adding it")
+ run_git(repo, "remote", "add", ORIGIN_REMOTE, expected_url)
+ return
+
+ if current.rstrip("/") == expected_url.rstrip("/"):
+ log.info(f"origin URL OK: {current}")
+ return
+
+ log.warning(f"origin URL mismatch — updating to {expected_url}")
+ run_git(repo, "remote", "set-url", ORIGIN_REMOTE, expected_url)
+
+
+def _get_remote_url(repo: Path, name: str) -> str | None:
+ """Return the fetch URL of remote *name*, or None if it doesn't exist."""
+ result = run_git_no_check(repo, "remote", "get-url", name)
+ if result.returncode == 0:
+ return result.stdout.strip()
+ return None
diff --git a/src/TA_main2main_workflow/pipeline/push_pr.py b/src/TA_main2main_workflow/pipeline/push_pr.py
new file mode 100644
index 0000000..e535baa
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/push_pr.py
@@ -0,0 +1,395 @@
+"""Pipeline step: Push work branch and create GitHub PR.
+
+Pushes to the user's fork through the CI proxy, then creates a PR
+from the fork branch to the upstream repo via ``gh`` CLI or REST API.
+Matches pre-refactor behaviour exactly.
+"""
+
+from __future__ import annotations
+
+import json
+import os
+import subprocess
+import time
+from datetime import datetime
+from pathlib import Path
+
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.git import run_git, run_git_no_check
+from TA_main2main_workflow.utils.submodule import push_submodule
+from TA_main2main_workflow.utils import WORKSPACE_DIR, FINAL_TARGET_PATCH_FILE, FINAL_SUMMARY_FILE
+
+log = get_logger(__name__)
+
+_MAX_PUSH_RETRIES = 5
+_MAX_PR_RETRIES = 5
+_RETRY_DELAY_BASE = 10 # seconds, multiplied by attempt number
+
+
+def push_and_create_pr(
+ ascend_path: str | Path,
+ github_repo: str,
+ summary_path: str | Path | None = None,
+ target_commit: str = "",
+ work_branch: str = "",
+) -> str:
+ """Push work branch to fork (via proxy) and create a GitHub PR.
+
+ Returns the PR URL on success.
+
+ Raises RuntimeError if push or PR creation fails after all retries.
+ """
+ ascend_path = Path(ascend_path)
+ summary_path = Path(summary_path) if summary_path else None
+
+ # ── 0. Push submodule first ────────────────────────────────────
+ push_submodule(ascend_path)
+
+ # ── 1. Determine branch ────────────────────────────────────────
+ branch = work_branch or run_git(ascend_path, "branch", "--show-current").strip()
+ log.info(f"Pushing branch: {branch}")
+
+ # ── 2. Fork owner (pushing to private fork, not upstream) ──────
+ fork_owner = os.environ.get("TA_FORK_OWNER") or "TecJesh"
+
+ # ── 3. Generate summary if missing ─────────────────────────────
+ summary_file = summary_path or (WORKSPACE_DIR / FINAL_SUMMARY_FILE)
+ if not summary_file.exists():
+ summary_file.write_text(
+ f"# Triton-Ascend Upstream Sync\n\n"
+ f"Branch: `{branch}`\n"
+ f"Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n",
+ encoding="utf-8",
+ )
+
+ # ── 4. Ensure gh auth ──────────────────────────────────────────
+ _ensure_gh_auth(ascend_path)
+
+ # ── 5. Push to fork through proxy ──────────────────────────────
+ token = os.environ.get("GH_TOKEN") or os.environ.get("GITHUB_TOKEN") or ""
+ _push_to_fork(ascend_path, branch, fork_owner, token)
+
+ # ── 6. Create PR from fork → upstream ──────────────────────────
+ pr_url = _create_pr(
+ ascend_path, github_repo, branch, fork_owner, token,
+ summary_file, target_commit,
+ )
+
+ return pr_url
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Internal
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _ensure_gh_auth(repo: Path) -> None:
+ """Ensure gh CLI is authenticated.
+
+ When GH_TOKEN is set, logs gh into github.com directly (necessary
+ when the git remote points to a proxy host). Also embeds the token
+ in the origin URL as a fallback for git push through proxy.
+ """
+ token = os.getenv("GH_TOKEN", "") or os.getenv("GITHUB_TOKEN", "")
+ if not token:
+ log.warning("No GH_TOKEN set — push/PR may fail")
+
+ if not token:
+ # Fall back to interactive auth check
+ try:
+ subprocess.run(
+ ["gh", "auth", "status"],
+ check=True, capture_output=True, text=True,
+ )
+ log.info("gh CLI already authenticated")
+ except subprocess.CalledProcessError:
+ log.warning("gh not authenticated and GH_TOKEN not set")
+ try:
+ subprocess.run(
+ ["gh", "auth", "setup-git"],
+ check=True, capture_output=True, text=True,
+ )
+ except Exception:
+ pass
+ return
+
+ log.info("Using GH_TOKEN from environment")
+
+ # Step 1: Explicitly login gh CLI against github.com.
+ # This is essential when the git remote points to a proxy host —
+ # gh needs to know about github.com independently of git remotes.
+ try:
+ subprocess.run(
+ ["gh", "auth", "login", "--with-token", "--hostname", "github.com"],
+ input=token.encode(), capture_output=True, timeout=30,
+ )
+ log.info("gh auth login OK")
+ except Exception as e:
+ log.warning(f"gh auth login failed: {e}")
+
+ # Step 2: Verify
+ try:
+ result = subprocess.run(
+ ["gh", "auth", "status", "--hostname", "github.com"],
+ capture_output=True, text=True, timeout=30,
+ )
+ log.info(f"gh auth status: {result.stdout.strip()}")
+ except Exception:
+ pass
+
+ # Step 3: Configure git credential helper (best-effort)
+ try:
+ subprocess.run(
+ ["gh", "auth", "setup-git", "--hostname", "github.com"],
+ capture_output=True, text=True, timeout=30,
+ )
+ log.info("gh auth setup-git OK")
+ except Exception:
+ pass
+
+ # Step 4: Embed token in origin URL (fallback for push through proxy)
+ try:
+ origin_url = run_git(repo, "remote", "get-url", "origin").strip()
+ if origin_url.startswith("https://"):
+ clean_url = origin_url.replace("https://", "", 1)
+ if "@" in clean_url:
+ clean_url = clean_url.split("@", 1)[1]
+ new_url = f"https://x-access-token:{token}@{clean_url}"
+ run_git(repo, "remote", "set-url", "origin", new_url)
+ safe = f"https://x-access-token:***@{clean_url}"
+ log.info(f"origin URL rewritten with token: {safe}")
+ except Exception as e:
+ log.warning(f"Could not rewrite origin URL: {e}")
+
+
+def _push_to_fork(repo: Path, branch: str, fork_owner: str, token: str) -> None:
+ """Push work branch to the user's fork through the CI proxy.
+
+ Creates a temporary remote ``ta-fork-push`` that goes through
+ ``gh-proxy.test.osinfra.cn`` to the user's fork, pushes, and
+ removes the temporary remote. Retries up to 5 times.
+ """
+ if not token or not fork_owner:
+ log.warning("No GH_TOKEN or TA_FORK_OWNER — falling back to direct push")
+ for attempt in range(1, _MAX_PUSH_RETRIES + 1):
+ log.info(f"Push attempt {attempt}/{_MAX_PUSH_RETRIES}...")
+ try:
+ run_git(repo, "push", "--force-with-lease", "origin", branch)
+ log.status(True, f"Pushed {branch}")
+ return
+ except Exception as e:
+ log.warning(f"Push failed (attempt {attempt}): {e}")
+ if attempt < _MAX_PUSH_RETRIES:
+ time.sleep(_RETRY_DELAY_BASE * attempt)
+ raise RuntimeError(f"Push failed after {_MAX_PUSH_RETRIES} attempts")
+ return
+
+ fork_remote = "ta-fork-push"
+ fork_url = (
+ f"https://x-access-token:{token}@"
+ f"gh-proxy.test.osinfra.cn/"
+ f"https://github.com/{fork_owner}/triton-ascend.git"
+ )
+
+ log.info(f"Pushing to fork {fork_owner}/triton-ascend via proxy...")
+ log.debug(f"fork remote: {fork_remote}")
+ log.debug(f"fork URL (masked): https://x-access-token:***@gh-proxy.test.osinfra.cn/https://github.com/{fork_owner}/triton-ascend.git")
+
+ last_error = ""
+ for attempt in range(1, _MAX_PUSH_RETRIES + 1):
+ log.info(f"Push attempt {attempt}/{_MAX_PUSH_RETRIES}...")
+ try:
+ # Remove stale temp remote
+ run_git_no_check(repo, "remote", "remove", fork_remote)
+ run_git(repo, "remote", "add", fork_remote, fork_url)
+
+ push_result = subprocess.run(
+ ["git",
+ "-c", "http.https://github.com/.extraheader=",
+ "push", "--force-with-lease", fork_remote, branch],
+ cwd=str(repo), capture_output=True, text=True,
+ )
+ # Clean up temp remote
+ run_git_no_check(repo, "remote", "remove", fork_remote)
+
+ if push_result.returncode == 0:
+ if push_result.stdout.strip():
+ log.info(f"push stdout: {push_result.stdout.strip()}")
+ log.status(True, f"Pushed {branch} to fork")
+ return
+
+ last_error = push_result.stderr.strip() or "(no stderr)"
+ log.warning(f"Push failed (attempt {attempt}): {last_error}")
+ except Exception as e:
+ last_error = str(e)
+ log.warning(f"Push failed (attempt {attempt}): {e}")
+ try:
+ run_git_no_check(repo, "remote", "remove", fork_remote)
+ except Exception:
+ pass
+
+ if attempt < _MAX_PUSH_RETRIES:
+ time.sleep(_RETRY_DELAY_BASE * attempt)
+
+ raise RuntimeError(f"Push failed after {_MAX_PUSH_RETRIES} attempts: {last_error}")
+
+
+def _create_pr(
+ repo: Path, github_repo: str, branch: str, fork_owner: str, token: str,
+ summary_file: Path, target_commit: str,
+) -> str:
+ """Create PR via gh CLI (with fork-aware origin swap).
+
+ Temporarily sets origin to the fork URL so ``gh`` can detect the
+ GitHub host, creates the PR from ``fork_owner:branch`` to the
+ upstream repo, then restores the saved origin.
+ """
+ pr_body = summary_file.read_text(encoding="utf-8") if summary_file.exists() else ""
+
+ title = _build_pr_title(target_commit)
+ head = f"{fork_owner}:{branch}" if fork_owner else branch
+ base_branch = os.getenv("TA_PR_BASE_BRANCH", "upstream-sync")
+
+ log.info(f"Creating PR: head={head}, base={base_branch}, repo={github_repo}")
+
+ # Save origin, swap to fork URL so gh CLI recognizes github.com
+ saved_origin = run_git(repo, "config", "--get", "remote.origin.url").strip()
+ if token and fork_owner:
+ pr_origin = (
+ f"https://x-access-token:{token}@"
+ f"github.com/{fork_owner}/triton-ascend.git"
+ )
+ else:
+ pr_origin = saved_origin
+
+ run_git(repo, "remote", "set-url", "origin", pr_origin)
+
+ last_error = ""
+ for attempt in range(1, _MAX_PR_RETRIES + 1):
+ try:
+ pr_url = _create_pr_via_gh(github_repo, title, pr_body, head, base_branch)
+ log.status(True, f"PR created: {pr_url}")
+ return pr_url
+ except Exception as e:
+ last_error = str(e)
+ log.warning(f"PR create attempt {attempt}/{_MAX_PR_RETRIES} FAILED: {last_error}")
+ if attempt < _MAX_PR_RETRIES:
+ time.sleep(_RETRY_DELAY_BASE * attempt)
+ finally:
+ # Always restore origin
+ try:
+ run_git(repo, "remote", "set-url", "origin", saved_origin)
+ except Exception:
+ pass
+
+ # Restore origin one more time in case of exception path
+ try:
+ run_git(repo, "remote", "set-url", "origin", saved_origin)
+ except Exception:
+ pass
+
+ # Fallback: GitHub REST API
+ log.info("Falling back to GitHub REST API...")
+ try:
+ return _create_pr_via_api(github_repo, head, title, pr_body, base_branch, token)
+ except Exception as e:
+ raise RuntimeError(f"PR creation failed after all attempts: {last_error}; API fallback: {e}")
+
+
+def _build_pr_title(target_commit: str = "") -> str:
+ """Build PR title in conventional commit format.
+
+ Example: [Sync](feat) Merge upstream triton commits abc12345
+
+ Env vars:
+ PR_AUTHOR — user tag (default: "Sync")
+ PR_TYPE — conventional commit type (default: "feat")
+ """
+ author = os.getenv("PR_AUTHOR", "Sync").strip()
+ pr_type = os.getenv("PR_TYPE", "feat").strip()
+ if target_commit:
+ return f"[{author}]({pr_type}) Merge upstream triton commits {target_commit[:8]}"
+ ts = datetime.now().strftime("%Y%m%d-%H%M%S")
+ return f"[{author}]({pr_type}) Merge upstream triton commits {ts}"
+
+
+def _create_pr_via_gh(
+ github_repo: str, title: str, body: str,
+ head_ref: str, base_branch: str,
+) -> str:
+ """Create a GitHub PR via the gh CLI.
+
+ Uses GH_TOKEN env var directly (overrides any auto GITHUB_TOKEN from
+ actions/checkout) so the PR can reference branches on the user's fork.
+ """
+ gh_token = os.environ.get("GH_TOKEN") or ""
+ cmd = [
+ "gh", "pr", "create",
+ "--title", title,
+ "--body", body,
+ "--head", head_ref,
+ "--base", base_branch,
+ "--repo", github_repo,
+ ]
+ log.info(f"Running: GH_HOST=github.com {' '.join(cmd)}")
+ result = subprocess.run(
+ cmd,
+ capture_output=True, text=True, timeout=60,
+ env={**os.environ,
+ "GITHUB_TOKEN": gh_token,
+ "GH_TOKEN": gh_token,
+ "GH_HOST": "github.com"},
+ )
+ if result.returncode != 0:
+ raise RuntimeError(
+ f"gh pr create failed (exit {result.returncode}): "
+ f"{result.stderr.strip()}"
+ )
+ pr_url = result.stdout.strip()
+ if not pr_url:
+ raise RuntimeError("gh pr create returned empty output")
+ return pr_url
+
+
+def _create_pr_via_api(
+ github_repo: str, head: str, title: str, body: str,
+ base: str, token: str,
+) -> str:
+ """Create a GitHub PR via the REST API (fallback).
+
+ Uses the REST API directly to avoid host-detection issues when git
+ remotes are rewritten by url.insteadOf proxy.
+ """
+ if not token:
+ raise RuntimeError("No GH_TOKEN or GITHUB_TOKEN set")
+
+ data = {
+ "title": title,
+ "head": head,
+ "base": base,
+ "body": body or "",
+ }
+
+ url = f"https://api.github.com/repos/{github_repo}/pulls"
+ cmd = [
+ "curl", "-s", "-X", "POST", url,
+ "-H", f"Authorization: Bearer {token}",
+ "-H", "Accept: application/vnd.github+json",
+ "-H", "X-GitHub-Api-Version: 2022-11-28",
+ "-H", "Content-Type: application/json",
+ "-d", json.dumps(data),
+ ]
+ result = subprocess.run(cmd, capture_output=True, text=True, timeout=30)
+ if result.returncode != 0:
+ raise RuntimeError(f"API PR creation failed: {result.stderr}")
+
+ try:
+ resp = json.loads(result.stdout)
+ if "html_url" in resp:
+ return resp["html_url"]
+ if "message" in resp:
+ raise RuntimeError(f"GitHub API error: {resp['message']}")
+ except json.JSONDecodeError:
+ pass
+
+ raise RuntimeError(f"Unexpected API response: {result.stdout[:500]}")
diff --git a/src/TA_main2main_workflow/pipeline/resolve.py b/src/TA_main2main_workflow/pipeline/resolve.py
new file mode 100644
index 0000000..2bfea41
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/resolve.py
@@ -0,0 +1,109 @@
+"""Pipeline step: AI resolve merge conflicts."""
+
+from __future__ import annotations
+
+import json
+from pathlib import Path
+
+from TA_main2main_workflow.agent.opencode_adapter import run_opencode_adapter
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.git import run_git
+from TA_main2main_workflow.pipeline.pre_ci import cleanup_temp_files, run_pre_ci_check
+from TA_main2main_workflow.utils import STEPS_DIR, WORKSPACE_DIR
+
+log = get_logger(__name__)
+_REF = str(Path(__file__).parent.parent / "reference")
+
+
+def resolve_conflicts(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """AI-driven merge conflict resolution with retry loop."""
+ if config.skip_ai_analysis:
+ log.warning("SKIP_AI_ANALYSIS=true — cannot resolve conflicts")
+ return ctx
+
+ ascend_path = Path(ctx.triton_ascend_path)
+ step = ctx.steps[ctx.current_step] if ctx.current_step < len(ctx.steps) else None
+ step_id = step["id"] if step else "step-0"
+ step_dir = WORKSPACE_DIR / STEPS_DIR / step_id
+ step_dir.mkdir(parents=True, exist_ok=True)
+
+ # Compute context for AI: previous step info for continuity
+ prev_step_id = ""
+ prev_summary_path = ""
+ if ctx.current_step > 0 and ctx.current_step <= len(ctx.steps):
+ prev = ctx.steps[ctx.current_step - 1]
+ prev_step_id = prev["id"]
+ prev_summary_path = str(
+ WORKSPACE_DIR / STEPS_DIR / prev_step_id / "step_summary.md"
+ )
+ is_last_step = (ctx.current_step >= ctx.total_steps - 1)
+
+ log.header("AI Conflict Resolution")
+ for attempt in range(1, config.max_retries + 1):
+ log.step(attempt, config.max_retries, "AI conflict resolution")
+ cf = [
+ f
+ for f in run_git(ascend_path, "diff", "--name-only", "--diff-filter=U")
+ .strip()
+ .splitlines()
+ if f
+ ]
+ if not cf:
+ log.status(True, "Already resolved!")
+ break
+ try:
+ run_opencode_adapter(
+ {
+ "step_id": f"{step_id}-conflict-{attempt}",
+ "previous_step_id": prev_step_id,
+ "previous_step_summary_path": prev_summary_path,
+ "is_last_step": str(is_last_step).lower(),
+ "step_dir": str(step_dir),
+ "conflict_dir": str(step_dir),
+ "ascend_path": str(ascend_path),
+ "triton_path": ctx.triton_ascend_path,
+ "reference_dir": _REF,
+ "mode": "conflict",
+ "error_logs": json.dumps(cf, ensure_ascii=False),
+ "target_commit": ctx.target_commit,
+ "step_index": f"{ctx.current_step + 1}/{ctx.total_steps}",
+ }
+ )
+ except Exception as e:
+ log.error(f"AI call failed: {e}")
+ if attempt < config.max_retries:
+ continue
+ break
+ if not run_git(ascend_path, "diff", "--name-only", "--diff-filter=U").strip():
+ log.status(True, f"Resolved (attempt {attempt})")
+ break
+ cf_remain = [
+ f
+ for f in run_git(ascend_path, "diff", "--name-only", "--diff-filter=U")
+ .strip()
+ .splitlines()
+ if f
+ ]
+ log.status(False, f"{len(cf_remain)} conflict(s) remain")
+ else:
+ log.error(f"Failed after {config.max_retries} attempts")
+ return ctx
+
+ cleanup_temp_files(ascend_path)
+ try:
+ run_git(ascend_path, "add", "-A")
+ run_git(ascend_path, "commit", "--no-edit", "-s")
+ log.status(True, "Committed resolution")
+ except Exception:
+ pass
+ pre_ci_ok = run_pre_ci_check(str(ascend_path), step_id="conflict-resolution")
+ if pre_ci_ok:
+ log.status(True, "Pre-CI check passed after conflict resolution")
+ else:
+ log.warning("Pre-CI check found issues after conflict resolution")
+ return ctx.copy_with(
+ merge_has_conflicts=False,
+ conflict_files_resolved=ctx.conflict_files_resolved + len(ctx.conflict_files),
+ )
diff --git a/src/TA_main2main_workflow/pipeline/test.py b/src/TA_main2main_workflow/pipeline/test.py
new file mode 100644
index 0000000..8b72683
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/test.py
@@ -0,0 +1,457 @@
+"""Pipeline step: Run pytest unit tests with retry/fix loop.
+
+Handles OOM detection and automatic concurrency reduction on retries.
+"""
+
+from __future__ import annotations
+
+import json
+import os
+import shutil
+import subprocess
+import time
+import xml.etree.ElementTree as ET
+from pathlib import Path
+
+from TA_main2main_workflow.utils.config import TAConfig
+from TA_main2main_workflow.utils.context import WorkflowContext
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils.tracker import timed
+from TA_main2main_workflow.utils import TEST_RESULT_FILE, WORKSPACE_DIR, STEPS_DIR
+from TA_main2main_workflow.pipeline.build import build_triton, commit_fixes
+from TA_main2main_workflow.pipeline.fix import ai_fix
+
+log = get_logger(__name__)
+
+# Single-file pre-test to smoke-check before running the full suite.
+# Must pass before any other tests are attempted.
+_PRETEST_FILE = "third_party/ascend/unittest/pytest_ut/test_add.py"
+
+
+def test_and_fix_loop(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Test + AI fix loop with OOM detection and reduced concurrency retry.
+
+ Used in single-step mode. After each test failure:
+ 1. Check for OOM → rerun with halved concurrency (up to 5 retries)
+ 2. If still failing: AI fix → rebuild TA → retest
+ """
+ if config.skip_e2e_test:
+ log.info("SKIP_E2E_TEST=true — treating tests as passed")
+ return ctx.copy_with(test_passed=True, pytest_passed=True)
+
+ ascend_path = Path(ctx.triton_ascend_path)
+
+ attempt = 0
+ while attempt <= config.max_retries:
+ ctx = ctx.copy_with(retry_count=attempt)
+
+ # ── Pre-test: smoke check before every test run ──────────────
+ ctx = _run_pretest_and_fix(ctx, config, ascend_path)
+ if not ctx.test_passed:
+ return ctx.copy_with(test_passed=False, pytest_passed=False)
+
+ if attempt > 0:
+ # ── OOM detection: check BEFORE AI fix ──────────────────
+ if detect_oom_in_tests(ctx):
+ log.warning("OOM detected in test output — reducing concurrency")
+ oom_ctx = rerun_tests_reduced_concurrency(ascend_path, config)
+ if oom_ctx is not None and oom_ctx.test_passed:
+ return ctx.copy_with(
+ test_passed=True, pytest_passed=True,
+ test_fix_count=ctx.test_fix_count,
+ )
+
+ log.header(f"Test Fix Attempt {attempt}/{config.max_retries}")
+ ctx = ai_fix(ctx, config, attempt=attempt, mode="fix")
+
+ # Rebuild TA after AI fix (old behavior: rebuild before retest)
+ with timed("test-fix-rebuild"):
+ ctx = build_triton(ctx, config, clean=False)
+ if not ctx.build_passed:
+ log.error("Rebuild after AI test fix failed")
+ continue
+
+ # Run tests
+ with timed("test"):
+ ctx = run_tests(ctx, config)
+
+ if ctx.test_passed:
+ # Commit AI test fixes with AI-authored message
+ if attempt > 0:
+ commit_fixes(ctx, config)
+ return ctx.copy_with(
+ test_passed=True, pytest_passed=True,
+ test_fix_count=ctx.test_fix_count + (1 if attempt > 0 else 0),
+ )
+
+ log.info(f"Tests failed (attempt {attempt + 1}) — retrying")
+ attempt += 1
+
+ return ctx.copy_with(test_passed=False, pytest_passed=False)
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Internal
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+def _run_pretest_and_fix(
+ ctx: WorkflowContext, config: TAConfig, ascend_path: Path,
+) -> WorkflowContext:
+ """Run a single-file pre-test with its own fix loop.
+
+ Must pass before the full test suite runs. Uses the same
+ OOM detection → AI fix → rebuild pattern as the main loop.
+ Called before every test retry to smoke-check the build.
+ """
+ pretest_path = ascend_path / _PRETEST_FILE
+ if not pretest_path.exists():
+ log.warning(f"Pre-test file not found: {_PRETEST_FILE} — skipping")
+ return ctx.copy_with(test_passed=True)
+
+ log.section("Pre-Test (smoke check)")
+ for pretest_attempt in range(config.max_retries + 1):
+ if pretest_attempt > 0:
+ if detect_oom_in_tests(ctx):
+ log.warning("OOM in pre-test — reducing concurrency")
+ oom_ctx = rerun_tests_reduced_concurrency(ascend_path, config)
+ if oom_ctx is not None and oom_ctx.test_passed:
+ log.status(True, "Pre-test passed (OOM rerun)")
+ return ctx.copy_with(test_passed=True)
+
+ log.header(f"Pre-Test Fix {pretest_attempt}/{config.max_retries}")
+ ctx = ai_fix(ctx, config, attempt=pretest_attempt, mode="fix")
+ with timed("pretest-fix-rebuild"):
+ ctx = build_triton(ctx, config, clean=False)
+ if not ctx.build_passed:
+ continue
+
+ with timed("pretest"):
+ ctx = _run_pytest(ctx, config, [_PRETEST_FILE],
+ test_procs=1, label="pretest")
+ if ctx.test_passed:
+ if pretest_attempt > 0:
+ commit_fixes(ctx, config)
+ log.status(True, "Pre-test passed")
+ return ctx.copy_with(test_passed=True)
+
+ log.info(f"Pre-test failed (attempt {pretest_attempt + 1})")
+
+ log.error(f"Pre-test failed after {config.max_retries} retries")
+ return ctx.copy_with(test_passed=False)
+
+
+def run_tests(ctx: WorkflowContext, config: TAConfig,
+ python_exe: str = "", test_procs: int = 0) -> WorkflowContext:
+ """Execute tests sequentially.
+
+ 1. Default pytest UT (primary test dir) — always runs first
+ 2. Extra test dirs — each runs individually, one after another
+ 3. Custom test command (``TA_TEST_COMMAND``) — runs last if set
+
+ Test results and error logs accumulate across all runs so AI fix
+ steps can see the full failure picture.
+ """
+ test_dirs = list(config.test_dirs)
+ if not test_dirs and not config.test_command:
+ log.warning("No test directories or test command configured")
+ return ctx.copy_with(test_passed=True, test_log_dir=str(WORKSPACE_DIR / "test-logs"))
+
+ primary = test_dirs[0] if test_dirs else None
+ extras = test_dirs[1:] if len(test_dirs) > 1 else []
+
+ all_passed = True
+ all_errors: list[str] = list(ctx.fix_errors)
+
+ # ── Step 1: Default pytest UT ─────────────────────────────────────
+ if primary:
+ log.section("Default pytest UT")
+ ctx = _run_pytest(ctx, config, [primary], python_exe=python_exe,
+ test_procs=test_procs, label="primary")
+ all_passed = all_passed and ctx.test_passed
+ all_errors.extend(ctx.fix_errors)
+
+ # ── Step 2: Extra test dirs — one by one ──────────────────────────
+ for i, extra_dir in enumerate(extras):
+ log.section(f"Extra Tests ({i + 1}/{len(extras)}): {extra_dir}")
+ extra_ctx = _run_pytest(ctx, config, [extra_dir], python_exe=python_exe,
+ test_procs=test_procs, label=f"extra-{i + 1}")
+ if not extra_ctx.test_passed:
+ all_passed = False
+ all_errors.extend(extra_ctx.fix_errors)
+
+ # ── Step 3: Custom test command ───────────────────────────────────
+ if config.test_command:
+ log.section("Custom Test Command (TA_TEST_COMMAND)")
+ custom_ctx = _run_custom_test(ctx, config)
+ if not custom_ctx.test_passed:
+ all_passed = False
+ all_errors.extend(custom_ctx.fix_errors)
+
+ # Merge: pass only if ALL suites pass; accumulate all error paths
+ ctx = ctx.copy_with(
+ test_passed=all_passed,
+ fix_errors=all_errors,
+ test_log_dir=str(WORKSPACE_DIR / "test-logs"),
+ )
+
+ if all_passed:
+ log.status(True, "All test suites passed")
+ else:
+ log.error(f"Tests FAILED — {len(all_errors)} error log(s)")
+ return ctx
+
+
+def _run_custom_test(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Execute a user-supplied test command and capture results.
+
+ The command is run via ``bash -c`` in the ascend repo root. stdout
+ and stderr are captured to ``test-logs/test-output.log``. If the
+ command produces a JUnit XML file (``--junitxml=...``), it is parsed
+ for detailed pass/fail/counts. Otherwise only the exit code is used.
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ test_log_dir = WORKSPACE_DIR / "test-logs"
+ test_log_dir.mkdir(parents=True, exist_ok=True)
+
+ cmd = config.test_command
+ output_log = test_log_dir / "test-output.log"
+
+ log.section("Run Tests (custom command)")
+ log.key_value("command", cmd)
+ log.key_value("output log", str(output_log))
+
+ _start = time.time()
+ with open(output_log, "w", encoding="utf-8") as fh:
+ fh.write(f"=== TA_TEST_COMMAND ===\n{cmd}\n\n")
+ fh.flush()
+ proc = subprocess.Popen(
+ ["bash", "-c", cmd],
+ cwd=str(ascend_path),
+ stdout=fh, stderr=subprocess.STDOUT,
+ )
+ try:
+ rc = proc.wait(timeout=7200)
+ except subprocess.TimeoutExpired:
+ proc.kill()
+ rc = -1
+ log.warning("Test command timed out after 7200s")
+
+ elapsed = time.time() - _start
+ log.info(f"Test command finished in {elapsed:.0f}s, exit={rc}")
+
+ # Try to parse JUnit XML if present (extract --junitxml=... from command)
+ pf = pe = tp = 0
+ junit_xml: Path | None = None
+ import re as _re
+ m = _re.search(r"--junitxml[= ](\S+)", cmd)
+ if m:
+ junit_xml = Path(m.group(1))
+ if not junit_xml.is_absolute():
+ junit_xml = ascend_path / junit_xml
+ if junit_xml.exists():
+ try:
+ tree = ET.parse(str(junit_xml))
+ root = tree.getroot()
+ suites = [root] if root.tag != "testsuites" else root.findall("testsuite")
+ for s in suites:
+ tp += int(s.get("tests", 0))
+ pf += int(s.get("failures", 0))
+ pe += int(s.get("errors", 0))
+ except Exception:
+ log.warning(f"Could not parse JUnit XML: {junit_xml}")
+
+ passed = rc == 0 and pf == 0 and pe == 0
+
+ # Write per-suite result file (unique name so it doesn't clobber pytest results)
+ result_file = test_log_dir / "test-result-custom.json"
+ summary = {
+ "label": "custom",
+ "exit_code": rc,
+ "passed": passed,
+ "test_log": str(junit_xml or output_log),
+ "test_command": cmd,
+ "passed_count": tp,
+ "failed_count": pf,
+ "error_count": pe,
+ }
+ result_file.write_text(
+ json.dumps(summary, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
+ )
+
+ if not passed:
+ log.error(f"Tests FAILED (exit={rc}, {pf} failed, {pe} errors)")
+ return ctx.copy_with(
+ test_passed=False,
+ fix_errors=[str(junit_xml or output_log), str(result_file)],
+ test_log_dir=str(test_log_dir),
+ )
+ log.status(True, f"All tests passed ({tp} passed)" if tp else f"Tests passed (exit 0)")
+ return ctx.copy_with(test_passed=True, test_log_dir=str(test_log_dir))
+
+
+# ---------------------------------------------------------------------------
+# Default pytest runner (used when TA_TEST_COMMAND is not set)
+# ---------------------------------------------------------------------------
+
+
+def _run_pytest(ctx: WorkflowContext, config: TAConfig,
+ test_dirs: list[str],
+ python_exe: str = "", test_procs: int = 0,
+ label: str = "pytest") -> WorkflowContext:
+ """Execute pytest for the given *test_dirs* in a single invocation.
+
+ Each call with a unique *label* writes to a separate JUnit XML file
+ (``pytest-junit-{label}.xml``) so results from different test suites
+ don't clobber each other.
+ """
+ ascend_path = Path(ctx.triton_ascend_path)
+ test_log_dir = WORKSPACE_DIR / "test-logs"
+ test_log_dir.mkdir(parents=True, exist_ok=True)
+
+ python_exe = python_exe or config.python_exe or os.getenv("PYTHON", "python3.10")
+ procs = test_procs or config.test_procs
+
+ # Resolve test directories, skipping missing ones
+ test_paths: list[Path] = []
+ for d in test_dirs:
+ p = (ascend_path / d).resolve()
+ if p.exists():
+ test_paths.append(p)
+ else:
+ log.warning(f"Test directory not found, skipping: {p}")
+
+ if not test_paths:
+ log.warning(f"[{label}] No test directories found — treating as passed")
+ return ctx.copy_with(test_passed=True)
+
+ junit_xml = test_log_dir / f"pytest-junit-{label}.xml"
+ pytest_bin = shutil.which("pytest")
+ cmd = (
+ [pytest_bin] if pytest_bin
+ else [python_exe, "-m", "pytest"]
+ )
+ cmd += [str(p) for p in test_paths]
+ cmd += ["-n", str(procs), f"--junitxml={junit_xml}"]
+
+ log.key_value(f"[{label}] test dirs", ", ".join(str(p.relative_to(ascend_path)) for p in test_paths))
+ log.info(f"[{label}] cmd: {' '.join(cmd)}")
+ _start = time.time()
+ try:
+ result = subprocess.run(cmd, cwd=ascend_path, timeout=3000)
+ rc = result.returncode
+ except subprocess.TimeoutExpired:
+ rc = -1
+ log.warning(f"[{label}] pytest timed out after 3000s")
+
+ elapsed = time.time() - _start
+ log.info(f"[{label}] pytest finished in {elapsed:.0f}s, returncode={rc}")
+
+ pf = pe = tp = 0
+ if junit_xml.exists():
+ try:
+ tree = ET.parse(junit_xml)
+ root = tree.getroot()
+ suites = [root] if root.tag != "testsuites" else root.findall("testsuite")
+ for s in suites:
+ tp += int(s.get("tests", 0))
+ pf += int(s.get("failures", 0))
+ pe += int(s.get("errors", 0))
+ except Exception:
+ pass
+
+ passed = pf == 0 and pe == 0
+
+ # Write per-suite result file (unique name so they don't clobber)
+ result_file = test_log_dir / f"test-result-{label}.json"
+ summary = {
+ "label": label,
+ "exit_code": 0 if passed else 1,
+ "passed": passed,
+ "test_log": str(junit_xml),
+ "test_dirs": [str(p) for p in test_paths],
+ "passed_count": tp,
+ "failed_count": pf,
+ "error_count": pe,
+ }
+ result_file.write_text(
+ json.dumps(summary, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
+ )
+
+ if not passed:
+ log.error(f"[{label}] Tests FAILED ({pf} failed, {pe} errors)")
+ return ctx.copy_with(
+ test_passed=False,
+ fix_errors=[str(junit_xml), str(result_file)],
+ test_log_dir=str(test_log_dir),
+ )
+ log.status(True, f"[{label}] All tests passed ({tp} passed)")
+ return ctx.copy_with(test_passed=True, test_log_dir=str(test_log_dir))
+
+
+def detect_oom_in_tests(ctx: WorkflowContext) -> bool:
+ """Check test output for Out-Of-Memory indicators."""
+ test_log_dir = Path(ctx.test_log_dir) if ctx.test_log_dir else None
+ if not test_log_dir or not test_log_dir.exists():
+ return False
+
+ oom_keywords = [
+ "OutOfMemoryError", "out of memory", "MemoryError",
+ "Cannot allocate memory", "OOM", "Killed",
+ "Exit code 137", "exit code 137",
+ "CUDA error", "cuMemAlloc", "NPU error",
+ ]
+ # Scan ALL JUnit XML files (each test suite writes its own)
+ for junit_xml in sorted(test_log_dir.glob("pytest-junit-*.xml")):
+ try:
+ content = junit_xml.read_text(encoding="utf-8", errors="replace").lower()
+ for kw in oom_keywords:
+ if kw.lower() in content:
+ log.warning(f"OOM indicator '{kw}' found in: {junit_xml.name}")
+ return True
+ except Exception:
+ pass
+ # Also scan custom test output log
+ output_log = test_log_dir / "test-output.log"
+ if output_log.exists():
+ try:
+ content = output_log.read_text(encoding="utf-8", errors="replace").lower()
+ for kw in oom_keywords:
+ if kw.lower() in content:
+ log.warning(f"OOM indicator found in test output: '{kw}'")
+ return True
+ except Exception:
+ pass
+ return False
+
+
+def rerun_tests_reduced_concurrency(
+ ascend_path: Path, config: TAConfig, max_reruns: int = 5
+) -> WorkflowContext | None:
+ """Rerun pytest with progressively halved concurrency.
+
+ Returns a new WorkflowContext with test results, or None if all fail.
+
+ Stops early when OOM indicators disappear from test output — remaining
+ failures are real code issues, not memory-related (matches pre-refactor).
+ """
+ original_procs = config.test_procs
+ for r in range(max_reruns):
+ procs = max(1, original_procs // (2 ** (r + 1)))
+ if procs >= original_procs:
+ break
+ log.info(f"Rerunning tests with {procs} workers (attempt {r + 1}/{max_reruns})...")
+
+ ascend_path_str = str(ascend_path)
+ ctx = WorkflowContext(triton_ascend_path=ascend_path_str)
+ ctx = run_tests(ctx, config, test_procs=procs)
+ if ctx.test_passed:
+ log.status(True, f"Tests passed with {procs} workers")
+ return ctx
+ # Stop if OOM is gone — remaining failures are code issues
+ if not detect_oom_in_tests(ctx):
+ log.info("OOM resolved — remaining failures are not memory-related, stopping rerun")
+ return ctx
+
+ log.error(f"Tests still failing after {max_reruns} concurrency reductions")
+ return None
diff --git a/src/TA_main2main_workflow/reference/AscendNPU-IR_LLVM_VERSION_COMPAT.md b/src/TA_main2main_workflow/reference/AscendNPU-IR_LLVM_VERSION_COMPAT.md
new file mode 100644
index 0000000..cfa8128
--- /dev/null
+++ b/src/TA_main2main_workflow/reference/AscendNPU-IR_LLVM_VERSION_COMPAT.md
@@ -0,0 +1,521 @@
+# AscendNPU-IR LLVM 版本兼容性适配指南
+
+本文档整理了 AscendNPU-IR(bishengir)为适配 LLVM 20/21/22 版本变更所做的全部兼容性适配,供后续版本升级参考。
+
+---
+
+## 1. 编译系统:版本标志定义
+
+**文件:** `CMakeLists.txt:17-36`
+
+```cmake
+option(LLVM_MAJOR_VERSION_20_COMPATIBLE "NPUIR build with LLVM 20" OFF)
+if(LLVM_MAJOR_VERSION_20_COMPATIBLE)
+ add_definitions(-D__LLVM_MAJOR_VERSION_20_COMPATIBLE__)
+endif()
+
+option(LLVM_MAJOR_VERSION_21_COMPATIBLE "NPUIR build with LLVM 21 or later" OFF)
+if(LLVM_MAJOR_VERSION_21_COMPATIBLE)
+ add_definitions(-D__LLVM_MAJOR_VERSION_21_COMPATIBLE__)
+endif()
+
+option(LLVM_MAJOR_VERSION_22_COMPATIBLE "NPUIR build with LLVM 22" OFF)
+if(LLVM_MAJOR_VERSION_22_COMPATIBLE)
+ add_definitions(-D__LLVM_MAJOR_VERSION_21_COMPATIBLE__) # LLVM 22 也兼容 21 的变更
+ add_definitions(-D__LLVM_MAJOR_VERSION_22_COMPATIBLE__)
+endif()
+```
+
+**规则:** LLVM 22 同时启用 21 和 22 的标志,因为 21 的兼容性变更在 22 中仍需保留。
+
+**子模块 TableGen 宏传递:**
+
+`bishengir/include/bishengir/Dialect/HACC/IR/CMakeLists.txt:9-11`:
+```cmake
+if(LLVM_MAJOR_VERSION_21_COMPATIBLE)
+ list(APPEND tblgen_feat_list -D__LLVM_MAJOR_VERSION_21_COMPATIBLE__)
+endif()
+```
+
+`bishengir/include/bishengir/Dialect/HFusion/IR/CMakeLists.txt:6-11`:
+```cmake
+if(LLVM_MAJOR_VERSION_21_COMPATIBLE)
+ list(APPEND tblgen_feat_list -D__LLVM_MAJOR_VERSION_21_COMPATIBLE__)
+endif()
+if(LLVM_MAJOR_VERSION_22_COMPATIBLE)
+ list(APPEND tblgen_feat_list -D__LLVM_MAJOR_VERSION_22_COMPATIBLE__)
+endif()
+```
+
+---
+
+## 2. 兼容模式分类
+
+### 2.1 `bufferization::ToMemrefOp` → `bufferization::ToBufferOp`(LLVM 22)
+
+**变更说明:** LLVM 22 将 `bufferization::ToMemrefOp` 重命名为 `bufferization::ToBufferOp`。
+
+**影响文件(共 6 处):**
+
+| 文件 | 行号 |
+|------|------|
+| `lib/ExecutionEngine/ConvertHIVMToUpstream.cpp` | 719-723 |
+| `lib/ExecutionEngine/CreateHostMain.cpp` | 228-234 |
+| `lib/Dialect/HIVM/IR/HIVMImpl.h` | 78-86 |
+| `lib/Dialect/HIVM/Utils/Utils.cpp` | 103-107, 709-713, 914-922 |
+| `lib/Dialect/HIVM/Transforms/InsertLoadStoreForMixCV/Utils.cpp` | 136-142 |
+
+**兼容模式:** `#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__` 走旧 API,`#else` 走新 API。
+
+```cpp
+// 模式 A: 创建 op(类型作为 op 名称)
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+ using bufferCastOp = bufferization::ToMemrefOp;
+#else
+ using bufferCastOp = bufferization::ToBufferOp;
+#endif
+
+// 模式 B: isa 类型匹配(isa → isa)
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+ } else if (auto toMemrefOp = v.getDefiningOp()) {
+#else
+ } else if (auto toBufferOp = v.getDefiningOp()) {
+#endif
+
+// 模式 C: isa 列表
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+ isa<..., bufferization::ToMemrefOp, bufferization::ToTensorOp>(userOp)
+#else
+ isa<..., bufferization::ToBufferOp, bufferization::ToTensorOp>(userOp)
+#endif
+```
+
+**`ToTensorOp` 构造函数变更(LLVM 22,与上述相关但独立):**
+
+LLVM 22 的 `bufferization::ToTensorOp::create` 需要显式传入 tensor 类型,旧版从 memref 自动推导:
+
+```cpp
+// 旧版: 自动推导返回类型
+rewriter.create(loc, alloc, true, true);
+// 新版: 需要显式传入 tensor type
+auto tensorType = RankedTensorType::get(targetShape, elementType);
+rewriter.create(loc, tensorType, alloc, true, true);
+```
+
+影响文件:
+- `lib/Dialect/HIVM/Utils/Utils.cpp:883-891`
+- `lib/Dialect/HIVM/Transforms/InsertLoadStoreForMixCV/Utils.cpp:136-142`
+
+---
+
+### 2.2 `getStridesAndOffset` 从自由函数变成成员函数(LLVM 21)
+
+**变更说明:** LLVM 21 之前 `getStridesAndOffset(memrefType)` 是全局自由函数;LLVM 21+ 变为 `memrefType.getStridesAndOffset()` 成员函数。
+
+**影响文件:**
+
+| 文件 | 行号 |
+|------|------|
+| `lib/Dialect/HIVM/IR/HIVMImpl.cpp` | 378-382 |
+| `lib/Dialect/HIVM/IR/HIVMTraits.cpp` | 47-52 |
+| `lib/Dialect/HIVM/Utils/Utils.cpp` | 1166-1170, 1296-1300 |
+| `lib/Dialect/HIVM/IR/BiShengIRAggregatedOpInterface/DecomposeOperation.cpp` | 1173-1179 |
+
+**兼容模式:**
+
+```cpp
+// 模式 A: 结构化绑定接收
+#ifndef __LLVM_MAJOR_VERSION_21_COMPATIBLE__
+ auto [strides, offset] = getStridesAndOffset(memrefType);
+#else
+ auto [strides, offset] = memrefType.getStridesAndOffset();
+#endif
+
+// 模式 B: 函数失败检查
+#ifndef __LLVM_MAJOR_VERSION_21_COMPATIBLE__
+ if (failed(getStridesAndOffset(srcType, srcStrides, srcOffset)))
+#else
+ if (failed(srcType.getStridesAndOffset(srcStrides, srcOffset)))
+#endif
+```
+
+---
+
+### 2.3 `RegionBuilderFn` 签名增加 `emitError` 回调(LLVM 22)
+
+**变更说明:** LLVM 22 中 `RegionBuilderFn` 的类型签名增加了第四个参数 `function_ref`,用于在 region builder 内部报错。
+
+**影响文件:**
+- `lib/Dialect/HFusion/IR/HFusionOps.cpp`(10+ 处)
+- `tools/bishengir-hfusion-ods-gen/bishengir-hfusion-ods-yaml-gen.cpp`
+
+**兼容模式:**
+
+```cpp
+// 类型别名定义
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+using RegionBuilderFn = llvm::function_ref)>;
+#else
+using RegionBuilderFn = llvm::function_ref,
+ function_ref)>;
+#endif
+
+// 调用点: 旧版 3 参数,新版 4 参数
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+ regionBuilder(b, *body, attrs);
+#else
+ regionBuilder(b, *body, attrs, [&]() {
+ return mlir::emitError(opBuilder.getUnknownLoc());
+ });
+#endif
+
+// 函数签名定义
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+std::function)>
+#else
+std::function,
+ function_ref)>
+#endif
+ReduceWithIndexOp::getRegionBuilder() { ... }
+
+// 函数体内接收额外参数(所有 getRegionBuilder 方法)
+ return [](ImplicitLocOpBuilder &b, Block &block,
+ ArrayRef attrs
+#ifdef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+ , function_ref emitError
+#endif
+ ) { ... };
+```
+
+涉及的自定义 Op:`ReduceWithIndexOp`, `ArangeOp`, `GatherOp`, `GatherMaskOp`, `Conv1DOp`, `Conv2DOp`, `Conv3DOp`
+
+---
+
+### 2.4 `MeshDialect` 头文件移除(LLVM 22)
+
+**变更说明:** LLVM 22 中 `mlir/Dialect/Mesh/IR/MeshDialect.h` 被移除。
+
+**影响文件:** `include/bishengir/Dialect/HFusion/IR/HFusion.h:21-23`
+
+**兼容模式:**
+
+```cpp
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+#include "mlir/Dialect/Mesh/IR/MeshDialect.h"
+#endif
+```
+
+---
+
+### 2.5 `CopyOpInterface` 从 MLIR 上游移除 → 本地 vendored(LLVM 22)
+
+**变更说明:** LLVM 22 移除了 `CopyOpInterface`(PR #157711)。AscendNPU-IR 在本地定义了一个等价的 interface。
+
+**文件结构:**
+- `include/bishengir/Interfaces/CopyOpInterface.td` — vendored 定义
+- `include/bishengir/Interfaces/CopyOpInterface.h` — 条件包含
+
+**影响文件:**
+- `include/bishengir/Interfaces/CopyOpInterface.h:21-28`
+- `lib/Dialect/HIVM/IR/HIVMInterfaces.cpp:42-45`
+
+**兼容模式:**
+
+```cpp
+// CopyOpInterface.h: LLVM 22 时使用本地 vendored 版本
+#ifdef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+#include "mlir/IR/OpDefinition.h"
+#include "bishengir/Interfaces/CopyOpInterface.h.inc"
+#endif
+
+// HIVMInterfaces.cpp: LLVM 22 时编译本地实现
+#ifdef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+#include "bishengir/Interfaces/CopyOpInterface.cpp.inc"
+#endif
+```
+
+---
+
+### 2.6 `linalg::ElemwiseBinaryOp` / `linalg::ElemwiseUnaryOp` 移除(LLVM 22)
+
+**变更说明:** LLVM 22 移除了 `linalg::ElemwiseBinaryOp` 和 `linalg::ElemwiseUnaryOp`,改用 `isElementwiseOp(op)` 通用函数替代。
+
+**影响文件:** `lib/Dialect/Utils/Util.cpp:1161-1216`
+
+**兼容模式:**
+
+```cpp
+// 模式 A: isa 检查替换为 isElementwiseOp
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+ return isa_and_present(op);
+#else
+ return isa_and_present(op) || isElementwiseOp(op);
+#endif
+
+// 模式 B: 特定 UnaryOp 检查
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+ return isa_and_present(op);
+#else
+ if (isa_and_present(op)) { return true; }
+ if (isElementwiseOp(op)) {
+ auto genericOp = dyn_cast(op);
+ return genericOp.getNumDpsInputs() == 1;
+ }
+ return false;
+#endif
+
+// 模式 C: 合法 Op 列表中移除
+bool isLegalOp(Operation *op) {
+ if (isa(op)) { return true; }
+}
+```
+
+---
+
+### 2.7 `TargetSystemSpecAttr` / `DeviceIDTargetDeviceSpecPair` → `DataLayoutEntryInterface`(LLVM 22)
+
+**变更说明:** LLVM 22 修改了 `TargetSystemSpecAttr` 的内部类型,从 `DeviceIDTargetDeviceSpecPair` 变成通用的 `DataLayoutEntryAttr` / `DataLayoutEntryInterface`。
+
+**影响文件:** `lib/Dialect/HACC/Utils/Utils.cpp:189-195`
+
+**兼容模式:**
+
+```cpp
+void setNPUTargetSpec(ModuleOp op, HACCTargetDeviceSpecInterface spec) {
+ MLIRContext *ctx = op->getContext();
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+ SmallVector entries;
+ entries.push_back({StringAttr::get(ctx, kNPUStr), spec});
+#else
+ SmallVector entries;
+ entries.push_back(
+ DataLayoutEntryAttr::get(ctx, StringAttr::get(ctx, kNPUStr), spec));
+#endif
+ op->setAttr(TargetSystemSpecAttr::name,
+ TargetSystemSpecAttr::get(ctx, entries));
+}
+```
+
+---
+
+### 2.8 `DISubprogramAttr::get` 参数变更(LLVM 20 / 22)
+
+**变更说明:** LLVM 20+ 中 `LLVM::DISubprogramAttr::get` 的参数增加了 `LLVM::DISubprogramFlags` 枚举替代原先的 `unsigned`,且中间参数的顺序/类型有调整。
+
+**影响文件:** `lib/Dialect/HACC/Utils/Utils.cpp:284-287`
+
+**兼容模式:**
+
+```cpp
+#if defined(__LLVM_MAJOR_VERSION_20_COMPATIBLE__) || defined(__LLVM_MAJOR_VERSION_22_COMPATIBLE__)
+ auto newAttr = LLVM::DISubprogramAttr::get(
+ llvmFunc->getContext(), DistinctAttr(), LLVM::DICompileUnitAttr(),
+ originalAttr.getScope(), originalAttr.getName(),
+ originalAttr.getLinkageName(), originalAttr.getFile(), unsigned(),
+ unsigned(), LLVM::DISubprogramFlags::Optimized,
+ originalAttr.getType(), {}, {});
+#else
+ auto newAttr = LLVM::DISubprogramAttr::get(
+ llvmFunc->getContext(), DistinctAttr(), LLVM::DICompileUnitAttr(),
+ originalAttr.getScope(), originalAttr.getName(),
+ originalAttr.getLinkageName(), originalAttr.getFile(), unsigned(),
+ unsigned(), originalAttr.getType(), {}, {});
+#endif
+```
+
+关键差异:
+- LLVM 20/22:`get(..., unsigned(), unsigned(), Flags, type, {}, {})` 多了 `Flags` 参数
+- LLVM 19:`get(..., unsigned(), unsigned(), type, {}, {})` 参数较少
+
+---
+
+### 2.9 `Record::getDirectSuperClasses` 返回类型变更(LLVM 20/21)
+
+**变更说明:** TableGen 的 `Record::getDirectSuperClasses` 方法签名逐版变化:
+- LLVM 19: `getDirectSuperClasses(SmallVector)`,元素为非 const
+- LLVM 20/21: `getDirectSuperClasses(SmallVector)`,输出参数
+- LLVM 22: `getDirectSuperClasses() -> ArrayRef>`,返回值
+
+**影响文件:** `tools/bishengir-target-spec-tblgen/TargetSpecGen.cpp:105-119`
+
+**兼容模式:**
+
+```cpp
+#if defined(__LLVM_MAJOR_VERSION_22_COMPATIBLE__)
+ // LLVM 22: 返回 ArrayRef>
+ auto superClasses = derivedClassRecord->getDirectSuperClasses();
+ const Record *superClass = superClasses.front().first;
+#elif defined(__LLVM_MAJOR_VERSION_21_COMPATIBLE__) || defined(__LLVM_MAJOR_VERSION_20_COMPATIBLE__)
+ // LLVM 20/21: 输出参数 SmallVector
+ SmallVector superClasses;
+ derivedClassRecord->getDirectSuperClasses(superClasses);
+ const Record *superClass = superClasses.front();
+#else
+ // LLVM 19: 输出参数 SmallVector
+ SmallVector superClasses;
+ derivedClassRecord->getDirectSuperClasses(superClasses);
+ Record *superClass = superClasses.front();
+#endif
+```
+
+同一文件中,函数签名也需要适配(共 3 处:`emitStrToSymFnForDeviceTarget`、`emitSymToStrFnForDeviceTarget` 等):
+
+```cpp
+// LLVM 20+:Record 指针是 const
+#if defined(__LLVM_MAJOR_VERSION_20_COMPATIBLE__) || defined(__LLVM_MAJOR_VERSION_21_COMPATIBLE__)
+static void emitStrToSymFn(const std::vector &records, ...);
+#else
+static void emitStrToSymFn(const std::vector &records, ...);
+#endif
+```
+
+以及在 `emitSymToStrFn` 函数体内,`switch` 语句格式略有不同:
+
+```cpp
+// LLVM 20+:
+OS << " switch (val) {\n";
+// LLVM 19:
+OS << formatv(" switch (val) {{\n", enumName);
+```
+
+---
+
+### 2.10 `arith::ConstantIntOp` 参数顺序变更(LLVM 22)
+
+**变更说明:** LLVM 22 中 `arith::ConstantIntOp::create` 的参数顺序从 `(loc, value, type)` 变成 `(loc, type, value)`。
+
+**影响文件:** `lib/Dialect/HFusion/IR/HFusionOps.cpp:2747-2751`
+
+**兼容模式:**
+
+```cpp
+#ifndef __LLVM_MAJOR_VERSION_22_COMPATIBLE__
+ return b.create(loc, 0, ty); // 旧: 值在前
+#else
+ return b.create(loc, ty, static_cast(0)); // 新: 类型在前
+#endif
+```
+
+---
+
+### 2.11 `concatAffineMaps` 增加 MLIRContext 参数(LLVM 21)
+
+**变更说明:** LLVM 21 中 `concatAffineMaps` 从 2 参数变为 3 参数,新增 `MLIRContext *` 参数。
+
+**影响文件:** `include/bishengir/Dialect/HIVM/IR/HIVMInterfaces.td:539`
+
+**兼容模式(TableGen):**
+
+```tablegen
+// HIVMInterfaces.td
+#ifndef __LLVM_MAJOR_VERSION_21_COMPATIBLE__
+ concatAffineMaps(maps) // LLVM 19/20: 2 参数
+#else
+ concatAffineMaps(maps, $_op.getContext()) // LLVM 21+: 3 参数
+#endif
+```
+
+---
+
+### 2.12 TableGen `$_op` 自引用限制(LLVM 21)
+
+**变更说明:** LLVM 21 的 TableGen 在 interface 方法的默认实现中,`$_op` 自引用不再能直接调用成员函数,必须先 cast 到 `ConcreteOp`。
+
+**影响文件:** `include/bishengir/Dialect/HIVM/Interfaces/OpPipeInterface.td:63,82,101`
+
+**兼容模式:**
+
+```tablegen
+// LLVM 19/20: 直接使用 $_op
+$_op.getPipe()
+
+// LLVM 21+: 先 cast 到 ConcreteOp
+ConcreteOp op = $_op;
+return op.getPipe();
+```
+
+---
+
+### 2.13 `RecordKeeper` const 修饰符变更(LLVM 20/21 vs 22)
+
+**变更说明:** TableGen 工具中 `RecordKeeper` 引用在 LLVM 20/21 是 `const`,LLVM 19 和 LLVM 22+ 是非 const。
+
+**影响文件:** `tools/bishengir-target-spec-tblgen/bishengir-target-spec-tblgen.cpp:46`
+
+**兼容模式:**
+
+```cpp
+#if defined(__LLVM_MAJOR_VERSION_20_COMPATIBLE__) || defined(__LLVM_MAJOR_VERSION_21_COMPATIBLE__)
+static bool bishengirTargetSpecGenMain(raw_ostream &os, const RecordKeeper &records)
+#else
+static bool bishengirTargetSpecGenMain(raw_ostream &os, RecordKeeper &records)
+#endif
+```
+
+---
+
+### 2.14 `StringSwitch` 格式化字符串差异(LLVM 20/21)
+
+**变更说明:** `emitSymToStrFnForDeviceTarget` 函数中 switch 语句的格式字符串有微小差异。
+
+**影响文件:** `tools/bishengir-target-spec-tblgen/TargetSpecGen.cpp:273`
+
+```cpp
+#if defined(__LLVM_MAJOR_VERSION_20_COMPATIBLE__) || defined(__LLVM_MAJOR_VERSION_21_COMPATIBLE__)
+ OS << " switch (val) {\n"; // 直接字符串
+#else
+ OS << formatv(" switch (val) {{\n", enumName); // 带格式
+#endif
+```
+
+---
+
+## 3. 快速参考
+
+| 版本 | 宏 | 关键变更 |
+|------|-----|---------|
+| LLVM 20 | `__LLVM_MAJOR_VERSION_20_COMPATIBLE__` | `DISubprogramAttr::get` 签名、`getDirectSuperClasses` const 化 |
+| LLVM 21 | `__LLVM_MAJOR_VERSION_21_COMPATIBLE__` | `getStridesAndOffset` 成员函数化、上一条的 const 化 |
+| LLVM 22 | `__LLVM_MAJOR_VERSION_22_COMPATIBLE__` + 上一条 | 见下表 |
+
+### LLVM 21 变更速查表
+
+| # | 变更项 | 适配方式 |
+|---|--------|---------|
+| 1 | `getStridesAndOffset` 从自由函数→成员函数 | 自由函数 for <21, 成员函数 for ≥21 |
+| 2 | `concatAffineMaps` 增加 `MLIRContext*` 参数 | TableGen `#ifndef` 分支 |
+| 3 | TableGen `$_op` 自引用需 cast `ConcreteOp` | TableGen `#ifndef` 分支,LLVM 21+ 先 cast |
+| 4 | `Record` 和 `RecordKeeper` const 化 | `#if` 三版本分支 |
+
+### LLVM 22 变更速查表
+
+| # | 变更项 | 适配方式 |
+|---|--------|---------|
+| 1 | `bufferization::ToMemrefOp` → `ToBufferOp` | `#ifndef` 用旧名,`#else` 用新名 |
+| 2 | `bufferization::ToTensorOp` 构造函数多一个 `Type` 参数 | `#ifndef` 旧签名,`#else` 新签名 |
+| 3 | `RegionBuilderFn` 增加 `emitError` 参数 | `#ifndef` 3参数,`#else` 4参数 |
+| 4 | `MeshDialect.h` 移除 | `#ifndef` 才 include |
+| 5 | `CopyOpInterface` 移除 | 本地 vendored,`#ifdef` 才 include/编译 |
+| 6 | `linalg::ElemwiseBinaryOp/UnaryOp` 移除 | 改用 `isElementwiseOp()` |
+| 7 | `DeviceIDTargetDeviceSpecPair` → `DataLayoutEntryInterface` | `#ifndef` 旧类型,`#else` 新类型 |
+| 8 | `getDirectSuperClasses` 返回值变化 | `#if` 三版本分支 |
+| 9 | `arith::ConstantIntOp` 参数顺序 | `#ifndef` 旧序,`#else` 新序 |
+
+---
+
+## 4. 新版本适配流程
+
+升级到 LLVM N+1 时,建议按以下步骤操作:
+
+1. **建 CMake Option:** 在根 `CMakeLists.txt` 中添加 `LLVM_MAJOR_VERSION_N_COMPATIBLE` 选项
+2. **传递到子模块:** 在各 `CMakeLists.txt` 的 `tblgen_feat_list` 中追加宏
+3. **分类处理变更:** 参照上述模式,在每个变更点用 `#ifndef __LLVM_MAJOR_VERSION_N_COMPATIBLE__` 保留旧版路径
+4. **复用机制:** 如果 N+1 也需保留 N 的变更,在 N+1 的 CMake 分支中同时 `add_definitions` N 的宏
+5. **清理旧版:** 不再需要支持的旧版本,删除对应 CMake option 和 `#ifndef` 分支
diff --git a/src/TA_main2main_workflow/reference/error-pattern-examples.md b/src/TA_main2main_workflow/reference/error-pattern-examples.md
index 6f61d3e..d2efe97 100644
--- a/src/TA_main2main_workflow/reference/error-pattern-examples.md
+++ b/src/TA_main2main_workflow/reference/error-pattern-examples.md
@@ -100,6 +100,35 @@ or references.
---
+## MLIR Op Rename: ToBufferOp ↔ ToMemrefOp
+
+**Error:** `error: 'ToBufferOp' is not a member of 'mlir::bufferization'`
+or `error: unknown type name 'ToBufferOp'`
+
+**Cause:** LLVM renamed `bufferization::ToBufferOp` back to
+`bufferization::ToMemrefOp` (the direction depends on the LLVM version).
+
+**Fix — when compiler cannot find ToBufferOp (→ use ToMemrefOp):**
+```bash
+# Find all references
+grep -rn "ToBufferOp" third_party/ascend/ lib/Target/Ascend/ \
+ --include="*.cpp" --include="*.h"
+```
+Replace all `bufferization::ToBufferOp` with `bufferization::ToMemrefOp`.
+
+**Fix — when compiler cannot find ToMemrefOp (→ use ToBufferOp):**
+Replace all `bufferization::ToMemrefOp` with `bufferization::ToBufferOp`.
+
+Also check for `using` aliases and `isa<>` / `dyn_cast<>` templates:
+```cpp
+// Old
+isa(op)
+// New
+isa(op)
+```
+
+---
+
## Backend Registration Change
**Error:** Ascend backend not found, device initialization failure, or
diff --git a/src/TA_main2main_workflow/scripts/__init__.py b/src/TA_main2main_workflow/scripts/__init__.py
deleted file mode 100644
index e69de29..0000000
diff --git a/src/TA_main2main_workflow/scripts/build_test.py b/src/TA_main2main_workflow/scripts/build_test.py
deleted file mode 100644
index d8c1ff5..0000000
--- a/src/TA_main2main_workflow/scripts/build_test.py
+++ /dev/null
@@ -1,678 +0,0 @@
-#!/usr/bin/env python3
-"""Build Triton-Ascend and run tests.
-
-Build steps:
- 1. Check LLVM version and rebuild if needed (unless SKIP_LLVM_REBUILD=true)
- 2. Build C++ extensions (CMake / setup.py build)
- 3. Install Python package in development mode
- 4. Run pre-commit checks (optional)
- 5. Run pytest unit tests
-
-Environment variables:
- LLVM_PROJECT_PATH — path to llvm-project repo (default: ~/workspace/llvm-project)
- LLVM_INSTALL_PREFIX_SYNC — where to install LLVM (default: ~/workspace/llvm-install-sync)
- SKIP_LLVM_REBUILD — set to "true" to skip LLVM rebuild check
-
-Output:
- - workspace/build_result.json
- - workspace/test_result.json
- - workspace/build.log
- - workspace/llvm_build.log
-"""
-
-from __future__ import annotations
-
-import json
-import os
-import signal
-import subprocess
-import sys
-import threading
-import time
-from pathlib import Path
-
-from TA_main2main_workflow.utils import (
- WORKSPACE_DIR, BUILD_RESULT_FILE, BUILD_LOG_FILE, TEST_RESULT_FILE,
- get_base_branch_ref,
-)
-
-
-def _run_to_log(cmd: list[str], cwd: Path, log_path: Path,
- env: dict | None = None, timeout: int | None = None,
- progress_line: bool = False) -> subprocess.CompletedProcess:
- """Run a command, tee output to log file and console.
-
- Output is streamed to the log file (full) and console (last line with
- \\r or every line depending on progress_line).
-
- If *timeout* is set and the subprocess does not exit within that many
- seconds, the entire process group is killed (via os.killpg). Pass
- timeout=None (the default) to block indefinitely — suitable for
- commands whose runtime is unbounded (e.g. pytest).
-
- Returns a CompletedProcess with returncode and a pointer to the log.
- """
- log_path.parent.mkdir(parents=True, exist_ok=True)
- proc_env = os.environ.copy()
- if env:
- proc_env.update(env)
-
- print(f" Running: {' '.join(cmd)}")
-
- # start_new_session=True gives the process its own process group.
- # close_fds=True prevents pytest-xdist workers from inheriting the
- # parent's file descriptors (especially stdin), which would cause
- # the process to hang at ~98% completion.
- # On timeout we can kill the entire group (pytest-xdist workers too).
- proc = subprocess.Popen(
- cmd, cwd=cwd, env=proc_env,
- stdin=subprocess.DEVNULL,
- stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True,
- start_new_session=True,
- close_fds=True,
- )
- assert proc.stdout is not None
-
- # ── Read output in a background thread so the main thread can
- # enforce the timeout via proc.wait(). When the process is
- # killed on timeout the pipe closes, unblocking the reader. ──
- last_line: str = ""
- read_error: Exception | None = None
-
- def _reader() -> None:
- nonlocal last_line, read_error
- try:
- with log_path.open("w", encoding="utf-8") as fh:
- for line in proc.stdout:
- fh.write(line)
- if progress_line:
- stripped = line.rstrip()
- if stripped:
- last_line = stripped
- # \r returns to line start, \033[K clears residue
- print(f"\r {stripped[:120]}\033[K", end="", flush=True)
- else:
- print(line, end="", flush=True)
- except Exception as exc:
- read_error = exc
- finally:
- if progress_line and last_line:
- print() # final newline
-
- reader = threading.Thread(target=_reader, daemon=True)
- reader.start()
-
- timed_out = False
- try:
- proc.wait(timeout=timeout)
- except subprocess.TimeoutExpired:
- timed_out = True
- print(f"\n ✗ Timeout after {timeout}s — killing process group "
- f"(pgid {proc.pid})...")
- # Kill the entire process group — catches pytest-xdist workers
- # that inherited the session from the main process.
- try:
- os.killpg(proc.pid, signal.SIGKILL)
- except (ProcessLookupError, OSError):
- pass
- try:
- proc.wait(timeout=30)
- except subprocess.TimeoutExpired:
- print(f" ✗ Process group did not respond to SIGKILL")
-
- # ── Close stdout pipe to unblock the reader thread ──
- # pytest-xdist workers may inherit the write end of the pipe,
- # keeping it open after the main process exits. Closing our
- # read end forces EOF on the pipe, unblocking the reader.
- try:
- proc.stdout.close()
- except Exception:
- pass
-
- # Wait for the reader thread to finish flushing the last lines
- reader.join(timeout=10)
-
- if read_error:
- # Ignore ValueError from closed pipe — this is expected
- if not isinstance(read_error, ValueError) or "closed" not in str(read_error).lower():
- print(f" ⚠ Reader thread error: {read_error}")
-
- if timed_out:
- raise subprocess.TimeoutExpired(cmd, timeout)
-
- if proc.returncode != 0:
- print(f" ✗ Exit code: {proc.returncode} (full log: {log_path})")
-
- return subprocess.CompletedProcess(
- cmd, proc.returncode,
- stdout="", stderr=f"See {log_path}"
- )
-
-
-def apply_llvm_patches(patch_dir: Path, llvm_project: Path,
- target_hash: str = "",
- patch_file: Path | None = None) -> dict:
- """Apply generated LLVM patch to llvm-project after cleaning stale state.
-
- 1. Clean any stale modifications in llvm-project (git checkout -- .)
- 2. Checkout the target LLVM commit
- 3. Apply the patch with 'git apply'
-
- If *patch_file* is given it is used directly; otherwise
- ``patch_dir / "ir_compat.patch"`` is used.
-
- This is a deterministic operation — no AI involved.
- Returns a dict with 'applied', 'failed', 'all_ok'.
- """
- if patch_file is None:
- patch_file = patch_dir / "ir_compat.patch"
- if not patch_file.exists():
- print(f" [llvm-patch] {patch_file.name} not found — nothing to apply")
- return {"applied": [], "failed": [], "all_ok": True}
-
- print(f"\n{'=' * 60}")
- print(f" Apply IR compat patch to LLVM")
- print(f"{'=' * 60}")
-
- # ── Step 1: Clean stale modifications ──
- print(" [llvm-patch] Cleaning stale changes in llvm-project...")
- subprocess.run(
- ["git", "checkout", "--", "."],
- cwd=llvm_project, capture_output=True, text=True,
- )
- subprocess.run(
- ["git", "clean", "-fd"],
- cwd=llvm_project, capture_output=True, text=True,
- )
- print(" [llvm-patch] Working tree cleaned")
-
- # ── Step 2: Checkout target LLVM commit ──
- if target_hash:
- print(f" [llvm-patch] Checking out LLVM commit: {target_hash[:12]}")
- result = subprocess.run(
- ["git", "checkout", target_hash],
- cwd=llvm_project, capture_output=True, text=True,
- )
- if result.returncode != 0:
- print(f" [llvm-patch] FAILED to checkout {target_hash[:12]}: "
- f"{result.stderr.strip()[-200:]}")
- return {"applied": [], "failed": [{
- "patch": str(patch_file),
- "error": f"git checkout failed: {result.stderr.strip()}",
- }], "all_ok": False}
- print(f" [llvm-patch] Checked out: {target_hash[:12]}")
-
- # ── Step 3: Apply the patch ──
- print(f" [llvm-patch] Applying: {patch_file.name}")
- # dry-run first
- proc = subprocess.run(
- ["git", "apply", "--check", str(patch_file)],
- cwd=llvm_project, capture_output=True, text=True,
- )
- if proc.returncode != 0:
- print(f" [llvm-patch] Patch does NOT apply cleanly:")
- print(f" {proc.stderr.strip()[-400:]}")
- return {"applied": [], "failed": [{
- "patch": str(patch_file),
- "error": proc.stderr.strip(),
- }], "all_ok": False}
-
- result = subprocess.run(
- ["git", "apply", str(patch_file)],
- cwd=llvm_project, capture_output=True, text=True,
- )
- if result.returncode != 0:
- print(f" [llvm-patch] FAILED: {result.stderr.strip()[-200:]}")
- return {"applied": [], "failed": [{
- "patch": str(patch_file),
- "error": result.stderr.strip(),
- }], "all_ok": False}
-
- print(f" [llvm-patch] ✓ Patch applied successfully")
- return {"applied": [str(patch_file)], "failed": [], "all_ok": True}
-
-
-def build_llvm(llvm_project: Path, llvm_install: Path,
- required_hash: str = "") -> str:
- """Build and install LLVM from the current working tree state.
-
- Does NOT check out commits, fetch, stash, or compare hashes.
- Assumes the caller has already prepared the working tree (checked
- out the correct commit, applied any patches, etc.).
-
- Cleans the build directory, runs cmake + ninja install, copies
- FileCheck, and writes the hash cache.
-
- Returns the LLVM install prefix path.
- """
- import shutil
-
- print(f"\n{'=' * 60}")
- print(f" Building LLVM from current working tree")
- if required_hash:
- print(f" Target hash: {required_hash[:12]}")
- print(f"{'=' * 60}")
-
- # Clean build directory
- llvm_build_log = WORKSPACE_DIR / "llvm_build.log"
- build_dir = llvm_project / "build"
- if build_dir.exists():
- shutil.rmtree(build_dir)
- build_dir.mkdir()
-
- cmake_cmd = [
- "cmake", str(llvm_project / "llvm"),
- "-G", "Ninja",
- "-DCMAKE_BUILD_TYPE=Release",
- "-DLLVM_ENABLE_ASSERTIONS=ON",
- "-DLLVM_ENABLE_PROJECTS=mlir;llvm;lld",
- "-DLLVM_TARGETS_TO_BUILD=host;NVPTX;AMDGPU",
- f"-DCMAKE_INSTALL_PREFIX={llvm_install}",
- "-DCMAKE_C_COMPILER=clang",
- "-DCMAKE_CXX_COMPILER=clang++",
- ]
- print(f" [llvm] Configuring...")
- _run_to_log(cmake_cmd, build_dir, llvm_build_log, timeout=300, progress_line=True)
-
- print(f" [llvm] Building (this may take a while)...")
- _run_to_log(
- ["ninja", "install"],
- build_dir, llvm_build_log, timeout=7200, progress_line=True,
- )
-
- # Copy FileCheck — not installed by ninja install
- filecheck_src = build_dir / "bin" / "FileCheck"
- filecheck_dst = llvm_install / "bin" / "FileCheck"
- if filecheck_src.exists():
- filecheck_dst.parent.mkdir(parents=True, exist_ok=True)
- shutil.copy2(filecheck_src, filecheck_dst)
- print(f" [llvm] Copied FileCheck to {filecheck_dst}")
- else:
- print(f" [llvm] WARNING: FileCheck not found at {filecheck_src}")
-
- # Write the hash cache
- llvm_install.mkdir(parents=True, exist_ok=True)
- hash_cache = llvm_install / ".llvm_hash"
- if required_hash:
- hash_cache.write_text(required_hash, encoding="utf-8")
- print(f" [llvm] Build complete — install prefix: {llvm_install}")
-
- return str(llvm_install)
-
-
-def _check_and_rebuild_llvm(repo_path: Path, force_rebuild: bool = False) -> str:
- """Check if LLVM version changed and rebuild if needed.
-
- Reads cmake/llvm-hash.txt from triton-ascend, compares with the
- last-built hash stored at {LLVM_INSTALL_PREFIX_SYNC}/.llvm_hash.
- If they differ (or no previous build exists), checks out the
- required commit in the pre-cloned llvm-project and rebuilds.
-
- When force_rebuild is True, skips the hash comparison and always
- rebuilds. Used after applying IR compatibility patches to LLVM.
-
- Environment variables:
- LLVM_PROJECT_PATH — path to llvm-project (default: ~/llvm-project)
- LLVM_INSTALL_PREFIX_SYNC — where to install LLVM (default: ~/llvm-install-sync)
-
- Returns the LLVM install prefix path.
- """
- llvm_project = Path(os.path.expanduser(
- os.getenv("LLVM_PROJECT_PATH", "~/llvm-project")))
- llvm_install = Path(os.path.expanduser(
- os.getenv("LLVM_INSTALL_PREFIX_SYNC", "~/llvm-install-sync")))
-
- # Read the required LLVM hash from triton-ascend
- llvm_hash_file = repo_path / "cmake" / "llvm-hash.txt"
- if not llvm_hash_file.exists():
- print(f" [llvm] {llvm_hash_file} not found — skipping LLVM rebuild")
- return str(llvm_install)
-
- required_hash = llvm_hash_file.read_text(encoding="utf-8").strip()
- if not required_hash:
- print(" [llvm] llvm-hash.txt is empty — skipping LLVM rebuild")
- return str(llvm_install)
-
- # Check last-built hash (skip when forcing rebuild)
- hash_cache = llvm_install / ".llvm_hash"
- if not force_rebuild and hash_cache.exists():
- last_hash = hash_cache.read_text(encoding="utf-8").strip()
- if last_hash == required_hash:
- print(f" [llvm] LLVM hash unchanged ({required_hash[:12]}) — skip rebuild")
- return str(llvm_install)
-
- if force_rebuild:
- print(f"\n{'=' * 60}")
- print(f" LLVM force rebuild requested (IR patches applied)")
- else:
- print(f"\n{'=' * 60}")
- print(f" LLVM version changed!")
- print(f" Previous: {hash_cache.read_text(encoding='utf-8').strip()[:12] if hash_cache.exists() else '(none)'}")
- print(f" Required: {required_hash[:12]}")
- print(f" Rebuilding LLVM...")
- print(f"{'=' * 60}")
-
- # Ensure llvm-project exists
- if not llvm_project.exists():
- raise RuntimeError(
- f"LLVM project not found at {llvm_project}. "
- f"Clone it with: git clone https://github.com/llvm/llvm-project.git {llvm_project}"
- )
-
- # ── Ensure the target commit is available locally ──
- # First check if the commit object exists via git cat-file -t.
- # If it doesn't, fetch it from origin (with retries).
- cat_proc = subprocess.run(
- ["git", "cat-file", "-t", required_hash],
- cwd=llvm_project, capture_output=True, text=True,
- )
- if cat_proc.returncode != 0:
- print(f" [llvm] Commit {required_hash[:12]} NOT found locally — fetching from origin...")
- for attempt in range(1, 7):
- fetch_proc = subprocess.run(
- ["git", "fetch", "origin", required_hash, "--no-tags"],
- cwd=llvm_project, capture_output=True, text=True, timeout=2000,
- )
- if fetch_proc.returncode == 0:
- print(f" [llvm] Fetch succeeded (attempt {attempt})")
- break
- print(f" [llvm] Fetch attempt {attempt}/6 failed — retrying...")
- else:
- raise RuntimeError(
- f"Failed to fetch LLVM commit {required_hash[:12]} after 6 attempts")
-
- # ── Clean working tree and checkout ──
- # Stash any local modifications so checkout doesn't fail on dirty tree,
- # then drop the stash (we want a pristine upstream checkout, not local edits).
- subprocess.run(
- ["git", "stash", "push", "--include-untracked",
- "-m", f"auto-stash-before-checkout-{required_hash[:12]}"],
- cwd=llvm_project, capture_output=True, text=True,
- )
- _run_cmd(
- ["git", "checkout", required_hash],
- cwd=llvm_project,
- timeout=2000,
- )
- # Drop the stash we just created (discard any local working-tree changes)
- subprocess.run(
- ["git", "stash", "drop", "--quiet"],
- cwd=llvm_project, capture_output=True, text=True,
- )
-
- return build_llvm(llvm_project, llvm_install, required_hash)
-
-
-def _run_cmd(cmd: list[str], cwd: Path, timeout: int = 300) -> str:
- """Run a command, return stdout. Raise on failure."""
- proc = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True, timeout=timeout)
- if proc.returncode != 0:
- print(f" [llvm] Command failed: {' '.join(cmd)}")
- print(f" stderr: {proc.stderr.strip()[-500:]}")
- return proc.stdout.strip()
-
-
-def build_triton_ascend(
- repo_path: Path,
- llvm_prefix: str = "",
- conda_env: str = "",
- build_dir: str = "build",
- clean_build: bool = False,
- python_exe: str = "python3",
-) -> dict:
- """Build the Triton-Ascend C++ extensions and Python package.
-
- python_exe: Python executable to use for setup.py install
- (default 'python3', use 'python3.10' / 'python3.11' for dual tests).
- """
- print("\n=== Building Triton-Ascend ===")
-
- # ── Check and rebuild LLVM if needed ──
- skip_llvm = os.getenv("SKIP_LLVM_REBUILD", "false").lower() == "true"
- if skip_llvm:
- print(" SKIP_LLVM_REBUILD=true — skipping LLVM version check")
- else:
- resolved_llvm_prefix = _check_and_rebuild_llvm(repo_path)
- if resolved_llvm_prefix and not llvm_prefix:
- llvm_prefix = resolved_llvm_prefix
-
- build_log = WORKSPACE_DIR / BUILD_LOG_FILE
-
- env = {}
- if llvm_prefix:
- env["LLVM_BUILD_DIR"] = llvm_prefix
- env["LLVM_INSTALL_PREFIX"] = llvm_prefix
-
- steps: list[dict] = []
- all_passed = True
-
- if clean_build:
- build_dir_path = repo_path / build_dir
- if build_dir_path.exists():
- print(f" Cleaning build directory: {build_dir_path}")
- subprocess.run(["rm", "-rf", str(build_dir_path)], check=False)
- steps.append({"step": "clean", "passed": True})
-
- print(" Building C++ extensions...")
-
- # --- Build via setup.py (retained for reference) ---
- # build_cmd = [
- # sys.executable, "-m", "pip", "install", "-e", ".",
- # "--no-build-isolation",
- # ]
-
- build_env = env.copy()
- build_env.update({
- "LLVM_SYSPATH": llvm_prefix,
- "TRITON_BUILD_WITH_CCACHE": "true",
- "TRITON_BUILD_WITH_CLANG_LLD": "true",
- "TRITON_BUILD_PROTON": "OFF",
- "DEBUG": "1",
- "TRITON_WHEEL_NAME": "triton-ascend",
- "TRITON_APPEND_CMAKE_ARGS": "-DTRITON_BUILD_UT=OFF",
- })
- build_cmd = [python_exe, "setup.py", "install"]
- build_proc = _run_to_log(build_cmd, repo_path, build_log, env=build_env, timeout=1800, progress_line=True)
- build_passed = build_proc.returncode == 0
- steps.append({
- "step": "setup_py_install",
- "passed": build_passed,
- "exit_code": build_proc.returncode,
- "log": str(build_log),
- })
- if not build_passed:
- all_passed = False
- print(" Build FAILED!")
- else:
- # Clear triton cache after a successful build
- cache_dir = Path.home() / ".triton" / "cache"
- if cache_dir.exists():
- print(f" Clearing triton cache: {cache_dir}")
- subprocess.run(["rm", "-rf", str(cache_dir)], check=False)
- steps.append({"step": "clear_cache", "passed": True})
-
- result = {
- "all_passed": all_passed,
- "steps": steps,
- "build_log": str(build_log),
- }
- (WORKSPACE_DIR / BUILD_RESULT_FILE).write_text(
- json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
- )
- return result
-
-
-def run_tests(
- repo_path: Path,
- test_dir: str = "third_party/ascend/unittest/pytest_ut",
- num_procs: int = 16,
- conda_env: str = "",
- python_exe: str = "",
-) -> dict:
- """Run pytest unit tests and return structured results.
-
- python_exe: Python executable for pytest (default '' uses PYTHON env var
- or 'python3'). Set to 'python3.10' / 'python3.11' for dual tests.
- """
- print("\n=== Running Tests ===")
- test_log_dir = WORKSPACE_DIR / "test-logs"
- test_log_dir.mkdir(parents=True, exist_ok=True)
-
- test_dir_path = repo_path / test_dir
-
- env = {}
- if conda_env:
- env["CONDA_DEFAULT_ENV"] = conda_env
-
- python_exe = python_exe or os.getenv("PYTHON", "python3.10")
-
- # Resolve to absolute path — test_dir_path may be relative, and the
- # subprocess cwd is repo_path. A relative path relative to repo_path
- # would double-up (e.g. triton-ascend/triton-ascend/third_party/…).
- test_dir_abs = test_dir_path.resolve()
-
- if not test_dir_abs.exists():
- print(f" WARNING: test directory not found: {test_dir_abs}")
- print(f" Skipping tests — directory does not exist after merge.")
- passed = False
- summary = {
- "exit_code": -1,
- "passed": False,
- "error": f"Test directory not found: {test_dir_abs}",
- "test_dir": str(test_dir_abs),
- }
- else:
- # Print bishengir-compile path before running tests
- import shutil
- bishengir_compile_path = shutil.which("bishengir-compile")
- print(f" bishengir-compile: {bishengir_compile_path or 'NOT FOUND'}")
-
- # JUnit XML for structured result parsing (replaces raw log regex).
- junit_xml = test_log_dir / "pytest-junit.xml"
-
- # Prefer pytest console script; fall back to python -m pytest.
- # -s : no capture — stdout/stderr inherit from the terminal.
- # pytest-xdist skips its internal IO-thread capture layer,
- # avoiding the fork()+IO-thread deadlock that hangs at 97%.
- # --junitxml : structured XML report for AI to read test results.
- pytest_bin = shutil.which("pytest")
- if pytest_bin:
- pytest_cmd = [
- pytest_bin, str(test_dir_abs),
- "-n", str(num_procs),
- # "-sv",
- f"--junitxml={junit_xml}",
- ]
- else:
- pytest_cmd = [
- python_exe, "-m", "pytest",
- str(test_dir_abs),
- "-n", str(num_procs),
- # "-sv",
- f"--junitxml={junit_xml}",
- ]
-
- proc_env = os.environ.copy()
- if env:
- proc_env.update(env)
-
- print(f" cwd: {repo_path}")
- print(f" cmd: {' '.join(pytest_cmd)}")
- print(f" junitxml: {junit_xml}")
- print(f" (stdout inherits terminal — no pipe, no tee, no capture)")
-
- # Run pytest with a 1000s timeout.
- _start = time.time()
- _timed_out = False
- try:
- result = subprocess.run(
- pytest_cmd,
- cwd=repo_path, env=proc_env,
- timeout=1000,
- )
- _rc = result.returncode
- except subprocess.TimeoutExpired:
- _timed_out = True
- _rc = -1
- print(f" pytest timed out after 1000s", flush=True)
-
- _elapsed = time.time() - _start
- if not _timed_out:
- print(f" pytest finished in {_elapsed:.0f}s, returncode={_rc}")
-
- # Parse JUnit XML first — real test results take priority over
- # process exit status.
- _pf = _pe = 0
- _tp = 0
- if junit_xml.exists():
- try:
- import xml.etree.ElementTree as ET
- tree = ET.parse(junit_xml)
- root = tree.getroot()
- suites = [root]
- if root.tag == "testsuites":
- suites = root.findall("testsuite")
- for suite in suites:
- _tp += int(suite.get("tests", 0))
- _pf += int(suite.get("failures", 0))
- _pe += int(suite.get("errors", 0))
- except Exception:
- pass
-
- # passed = no test failures. Timeout in teardown (all tests
- # already finished) is NOT a test failure.
- passed = (_pf == 0 and _pe == 0)
-
- summary = {
- "exit_code": 0 if passed else 1,
- "passed": passed,
- "test_log": str(junit_xml),
- "test_dir": str(test_dir_path),
- "passed_count": _tp,
- "failed_count": _pf,
- "error_count": _pe,
- }
- if _timed_out:
- summary["timed_out"] = True
-
- precommit_config = repo_path / ".pre-commit-config.yaml"
- if precommit_config.exists():
- print("\n Running pre-commit checks...")
- precommit_log = test_log_dir / "precommit.log"
- precommit_passed = True
- try:
- pc_proc = subprocess.run(
- ["pre-commit", "run", "--from-ref", get_base_branch_ref(), "--to-ref", "HEAD"],
- cwd=repo_path,
- stdout=precommit_log.open("w"),
- stderr=subprocess.STDOUT,
- timeout=300,
- )
- precommit_passed = pc_proc.returncode == 0
- except (subprocess.TimeoutExpired, FileNotFoundError):
- precommit_passed = False
-
- # ── If pre-commit auto-fixed files, amend the latest commit ──
- if not precommit_passed:
- print(" Pre-commit found issues — checking for auto-fixes...")
- from TA_main2main_workflow.utils import run_git_no_check
- status_proc = run_git_no_check(repo_path, "status", "--porcelain")
- if status_proc.stdout.strip():
- print(" Pre-commit applied auto-fixes, amending commit...")
- run_git_no_check(repo_path, "add", "-u")
- run_git_no_check(repo_path, "commit", "--amend", "--no-edit")
- print(" Commit amended with pre-commit fixes.")
- else:
- print(" Pre-commit failed but no auto-fixes were applied "
- "(manual review may be needed).")
- else:
- print(" Pre-commit checks passed.")
-
- summary["precommit_passed"] = precommit_passed
-
- result_path = WORKSPACE_DIR / TEST_RESULT_FILE
- result_path.write_text(
- json.dumps(summary, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
- )
-
- return summary
diff --git a/src/TA_main2main_workflow/scripts/detect_commits.py b/src/TA_main2main_workflow/scripts/detect_commits.py
deleted file mode 100644
index 2d58aaa..0000000
--- a/src/TA_main2main_workflow/scripts/detect_commits.py
+++ /dev/null
@@ -1,147 +0,0 @@
-#!/usr/bin/env python3
-"""Detect the commit gap between triton-ascend and upstream triton.
-
-For a merge-based workflow (Triton-Ascend is a fork of Triton), we:
- 1. Find the merge-base between the current triton-ascend branch and the
- upstream triton target commit.
- 2. List commits on the upstream side since that merge-base.
- 3. Determine total changed files and lines for planning.
-
-Output: workspace/detect.json
-"""
-
-from __future__ import annotations
-
-import json
-import os
-from pathlib import Path
-
-from TA_main2main_workflow.utils import (
- WORKSPACE_DIR, DETECT_FILE, run_git, get_repo_head, get_merge_base,
- ENV_BASE_BRANCH, get_base_branch_ref,
-)
-
-
-def _list_upstream_commits(repo: Path, merge_base: str, target: str) -> list[dict]:
- """List commits between merge_base and target, ordered chronologically."""
- log_output = run_git(
- repo, "log", "--reverse", "--format=%H%x1f%s",
- f"{merge_base}..{target}"
- )
- commits: list[dict] = []
- for line in log_output.strip().splitlines():
- if not line.strip():
- continue
- parts = line.split("\x1f", 1)
- commits.append({
- "sha": parts[0].strip(),
- "subject": parts[1].strip() if len(parts) > 1 else "",
- })
- return commits
-
-
-def _count_changed_lines(repo: Path, merge_base: str, target: str) -> dict:
- """Count changed lines in key source directories."""
- dirs = ["python/triton/", "lib/", "include/", "third_party/nvidia/", "third_party/amd/"]
- result = {}
- total = 0
- for d in dirs:
- try:
- output = run_git(
- repo, "diff", "--numstat", merge_base, target, "--", f":(top){d}"
- )
- except Exception:
- result[d] = 0
- continue
- lines = 0
- for line in output.strip().splitlines():
- if not line.strip():
- continue
- parts = line.split("\t")
- if len(parts) >= 3:
- added = int(parts[0]) if parts[0] != "-" else 0
- deleted = int(parts[1]) if parts[1] != "-" else 0
- lines += added + deleted
- result[d] = lines
- total += lines
- result["total"] = total
- return result
-
-
-def _changed_files(repo: Path, merge_base: str, target: str) -> list[str]:
- """Return list of changed files between merge_base and target."""
- output = run_git(repo, "diff", "--name-only", merge_base, target)
- return sorted(f for f in output.strip().splitlines() if f)
-
-
-def detect(
- triton_ascend_path: Path,
- triton_path: Path,
- target_commit: str | None = None,
-) -> tuple[dict, bool]:
- """Detect upstream commits that need to be merged.
-
- Returns (result_dict, has_new_commits).
- """
- # ── Fetch latest from the configured base branch ──
- base_branch = os.getenv(ENV_BASE_BRANCH, "main")
- base_ref = get_base_branch_ref()
- try:
- run_git(triton_ascend_path, "fetch", "origin", base_branch)
- print(f"[detect] Fetched {base_ref} (private fork)")
- except Exception:
- print(f"[detect] Warning: could not fetch {base_ref}, using local refs")
-
- # ── Fetch latest from upstream-triton ──
- if not target_commit:
- try:
- run_git(triton_ascend_path, "fetch", "upstream-triton", "--prune")
- except Exception:
- print("[detect] Warning: could not fetch upstream-triton, using local refs")
-
- # ── Use the configured base branch as the ascend reference (not checkout HEAD) ──
- # The work branch will be created from the base branch, so the merge_base
- # must be computed against it — otherwise we'd include commits
- # that are already on the base branch.
- try:
- ascend_head = run_git(triton_ascend_path, "rev-parse", base_ref).strip()
- except Exception:
- ascend_head = get_repo_head(triton_ascend_path)
- print(f"[detect] Warning: {base_ref} not available, using checkout HEAD")
-
- target = target_commit if target_commit else get_repo_head(triton_path)
-
- # ── Debug: print key refs ──
- checkout_head = get_repo_head(triton_ascend_path)
- print(f"[detect] Checkout HEAD : {checkout_head[:12]}")
- print(f"[detect] {base_ref} : {ascend_head[:12]}")
- print(f"[detect] upstream target: {target[:12]}")
-
- merge_base = get_merge_base(triton_ascend_path, ascend_head, target)
- print(f"[detect] merge_base : {merge_base[:12]}")
-
- commits = _list_upstream_commits(triton_path, merge_base, target)
- has_new = len(commits) > 0 and merge_base != target
-
- if has_new:
- print(f"[detect] {len(commits)} new upstream commits to merge "
- f"({commits[0]['sha'][:8]}..{commits[-1]['sha'][:8]})")
- else:
- print("[detect] No new upstream commits — already up to date")
-
- result = {
- "ascend_head": ascend_head,
- "target_commit": target,
- "merge_base": merge_base,
- "upstream_commits_count": len(commits),
- "upstream_commits": commits,
- "changed_lines": _count_changed_lines(triton_path, merge_base, target),
- "changed_files": _changed_files(triton_path, merge_base, target),
- "changed_files_count": len(_changed_files(triton_path, merge_base, target)),
- }
-
- (WORKSPACE_DIR / DETECT_FILE).write_text(
- json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
- )
-
- return result, has_new
diff --git a/src/TA_main2main_workflow/scripts/merge_upstream.py b/src/TA_main2main_workflow/scripts/merge_upstream.py
deleted file mode 100644
index fbeb400..0000000
--- a/src/TA_main2main_workflow/scripts/merge_upstream.py
+++ /dev/null
@@ -1,341 +0,0 @@
-#!/usr/bin/env python3
-"""Perform git merge of upstream triton commits into triton-ascend work branch.
-
-Creates a work branch based on the latest main from triton-lang/triton-ascend
-(fetched fresh each run), then merges the target upstream commit.
-If merge conflicts occur, they are recorded for later AI resolution.
-
-Output:
- - workspace/merge_result.json
- - workspace/merge.log (raw git merge output)
-"""
-
-from __future__ import annotations
-
-import json
-import os
-import subprocess
-from datetime import datetime
-from pathlib import Path
-
-from TA_main2main_workflow.utils import (
- WORKSPACE_DIR, MERGE_RESULT_FILE, MERGE_LOG_FILE, CONFLICT_LOG_DIR,
- run_git, run_git_no_check, has_merge_conflicts, get_conflict_files,
- ENV_BASE_BRANCH, get_base_branch_ref,
-)
-
-
-def _check_tracked_changes(repo: Path) -> bool:
- """Return True if tracked files have uncommitted changes (modified or staged)."""
- unstaged = run_git_no_check(repo, "diff", "--quiet")
- staged = run_git_no_check(repo, "diff", "--cached", "--quiet")
- return unstaged.returncode != 0 or staged.returncode != 0
-
-
-def _auto_stash(repo: Path) -> str:
- """Stash all changes (including untracked). Returns the stash name."""
- ts = datetime.now().strftime("%Y%m%d-%H%M%S")
- name = f"ta-sync-{ts}"
- run_git(repo, "stash", "push", "-u", "-m", name)
- print(f"[merge] Auto-stashed changes as '{name}'")
- return name
-
-
-def _abort_stale_merge(repo: Path) -> None:
- """Abort any stale merge in progress."""
- merge_head = repo / ".git" / "MERGE_HEAD"
- if merge_head.exists():
- print("[merge] Found stale MERGE_HEAD, running git merge --abort")
- try:
- run_git(repo, "merge", "--abort")
- except subprocess.CalledProcessError:
- print("[merge] Warning: git merge --abort failed, trying git reset --hard HEAD")
- run_git(repo, "reset", "--hard", "HEAD")
- for f in [".git/MERGE_MODE", ".git/MERGE_MSG", ".git/CHERRY_PICK_HEAD"]:
- p = repo / f
- if p.exists():
- p.unlink()
-
-
-def _ensure_upstream_ascend_remote(repo: Path) -> str:
- """Ensure a remote for triton-lang/triton-ascend exists and return its name.
-
- Checks existing remotes for one that points to triton-lang/triton-ascend.
- If none found, adds a remote named 'upstream-ascend'.
- Returns the remote name to use for fetching.
- """
- ASCEND_UPSTREAM_URL = "https://github.com/triton-lang/triton-ascend.git"
-
- # Check if any existing remote already points to the ascend upstream
- remotes_proc = run_git_no_check(repo, "remote", "-v")
- for line in remotes_proc.stdout.strip().splitlines():
- if ASCEND_UPSTREAM_URL in line:
- remote_name = line.split()[0]
- print(f"[merge] Found existing remote '{remote_name}' → {ASCEND_UPSTREAM_URL}")
- return remote_name
-
- # Not found — add a new remote
- remote_name = "upstream-ascend"
- print(f"[merge] Adding remote '{remote_name}' → {ASCEND_UPSTREAM_URL}")
- run_git(repo, "remote", "add", remote_name, ASCEND_UPSTREAM_URL)
- return remote_name
-
-
-def _create_work_branch(repo: Path, suffix: str = "") -> str:
- """Create and checkout a work branch for the merge.
-
- Default: branch from triton-lang/triton-ascend/upstream-sync.
- Set TA_WORK_BRANCH_BASE=origin to branch from the local fork instead,
- and TA_BASE_BRANCH to change the branch name on either remote.
- """
- ts = datetime.now().strftime("%Y%m%d-%H%M%S")
- branch = f"auto/upstream-sync-{ts}{'-' + suffix if suffix else ''}"
-
- if _check_tracked_changes(repo):
- auto_stash = os.getenv("AUTO_STASH", "false").lower() == "true"
- if auto_stash:
- _auto_stash(repo)
- else:
- print("[merge] ERROR: Tracked files have uncommitted changes.")
- print("[merge] Hint: git stash push -u -m 'pre-sync-stash'")
- print("[merge] or set AUTO_STASH=true to auto-stash before sync")
- raise RuntimeError(
- "Working tree has uncommitted changes to tracked files. "
- "Commit or stash changes before running sync."
- )
-
- _abort_stale_merge(repo)
-
- # Base branch for work branches. Defaults to the upstream repo's
- # upstream-sync branch (triton-lang/triton-ascend). Set
- # TA_WORK_BRANCH_BASE=origin to branch from the local fork instead.
- _upstream_branch = os.getenv(ENV_BASE_BRANCH, "upstream-sync")
- branch_base = os.getenv("TA_WORK_BRANCH_BASE", "upstream-ascend")
-
- if branch_base == "upstream-ascend":
- upstream_remote = _ensure_upstream_ascend_remote(repo)
- print(f"[merge] Fetching latest {_upstream_branch} from "
- f"'{upstream_remote}'...")
- run_git(repo, "fetch", upstream_remote, _upstream_branch)
- base_ref = f"{upstream_remote}/{_upstream_branch}"
- else:
- base_ref = get_base_branch_ref()
- base_branch = os.getenv(ENV_BASE_BRANCH, "main")
- print(f"[merge] Fetching latest {base_ref} from origin...")
- try:
- run_git(repo, "fetch", "origin", base_branch)
- except Exception:
- print(f"[merge] Warning: could not fetch {base_ref}, using local ref")
-
- # ── Reset to the pristine base ref and clean the working tree ──
- # This guarantees the work branch starts from exactly the right
- # commit, with no leftover artifacts from previous runs.
- print(f"[merge] Resetting working tree to {base_ref}...")
- run_git(repo, "checkout", "--detach", base_ref)
- run_git(repo, "reset", "--hard", "HEAD")
- run_git(repo, "clean", "-fd")
-
- # Resolve base ref to a commit so we can log both the name and the SHA
- base_sha = run_git(repo, "rev-parse", base_ref).strip()
- print(f"[merge] Base branch: {base_ref} commit: {base_sha[:12]}")
-
- print(f"[merge] Creating work branch '{branch}' from {base_ref}")
- proc = run_git_no_check(repo, "checkout", "-B", branch, base_ref)
- if proc.returncode != 0:
- print(f"[merge] ERROR: git checkout -B {branch} {base_ref} failed")
- print(f"[merge] stderr: {proc.stderr.strip()}")
- raise RuntimeError(f"Failed to create work branch '{branch}': {proc.stderr.strip()}")
-
- print(f"[merge] Created work branch: {branch} (based on {base_ref})")
- return branch
-
-
-def _get_conflict_content(repo: Path, filepath: str) -> str:
- """Get the content of a conflicted file (with conflict markers)."""
- file_path = Path(repo) / filepath
- if file_path.exists():
- return file_path.read_text(encoding="utf-8", errors="replace")
- return ""
-
-
-def _save_conflict_info(repo: Path, conflict_files: list[str], log_dir: Path) -> list[dict]:
- """Save conflict file contents and return structured conflict info."""
- conflicts = []
- for f in conflict_files:
- content = _get_conflict_content(repo, f)
- conflict_file = log_dir / f"{f.replace('/', '_')}.conflict"
- conflict_file.parent.mkdir(parents=True, exist_ok=True)
- conflict_file.write_text(content, encoding="utf-8")
- conflicts.append({
- "file": f,
- "conflict_snapshot": str(conflict_file),
- "size_bytes": len(content),
- })
- return conflicts
-
-
-def run_merge(
- triton_ascend_path: Path,
- triton_path: Path,
- target_commit: str,
-) -> dict:
- """Merge upstream triton *target_commit* into triton-ascend.
-
- Returns a dict with merge status, branch name, conflict info.
- """
- ascend_path = Path(triton_ascend_path)
-
- original_branch = run_git(ascend_path, "branch", "--show-current").strip()
- if not original_branch:
- original_branch = run_git(ascend_path, "rev-parse", "HEAD").strip()
-
- work_branch = _create_work_branch(ascend_path)
-
- try:
- run_git(ascend_path, "fetch", "upstream-triton", "--prune")
- except subprocess.CalledProcessError:
- print("[merge] Warning: could not fetch upstream-triton, assuming target is reachable")
-
- if Path(triton_path) != ascend_path:
- try:
- run_git(ascend_path, "fetch", str(triton_path), target_commit)
- except subprocess.CalledProcessError:
- print("[merge] Warning: could not fetch target from triton path")
-
- print(f"[merge] Merging {target_commit[:12]} into {work_branch}")
- merge_proc = run_git_no_check(
- ascend_path, "merge", "--no-ff", "--no-edit", target_commit
- )
-
- # Use a timestamped log file so each merge step's output is preserved
- ts = datetime.now().strftime("%Y%m%d-%H%M%S")
- merge_log_path = WORKSPACE_DIR / f"merge-{ts}.log"
- merge_log_path.write_text(
- f"STDOUT:\n{merge_proc.stdout}\n\nSTDERR:\n{merge_proc.stderr}\n",
- encoding="utf-8",
- )
- # Also write/update the canonical merge log for quick access to the latest
- (WORKSPACE_DIR / MERGE_LOG_FILE).write_text(
- f"STDOUT:\n{merge_proc.stdout}\n\nSTDERR:\n{merge_proc.stderr}\n",
- encoding="utf-8",
- )
-
- has_conflicts = has_merge_conflicts(ascend_path)
- conflict_files = get_conflict_files(ascend_path) if has_conflicts else []
-
- conflict_dir = WORKSPACE_DIR / CONFLICT_LOG_DIR
- conflict_info = []
- if has_conflicts:
- conflict_dir.mkdir(parents=True, exist_ok=True)
- conflict_info = _save_conflict_info(ascend_path, conflict_files, conflict_dir)
-
- result = {
- "work_branch": work_branch,
- "original_branch": original_branch,
- "target_commit": target_commit,
- "merge_exit_code": merge_proc.returncode,
- "has_conflicts": has_conflicts,
- "conflict_files": conflict_files,
- "conflict_count": len(conflict_files),
- "conflicts": conflict_info,
- "merge_log": str(merge_log_path),
- "conflict_dir": str(conflict_dir) if has_conflicts else "",
- }
-
- result_path = WORKSPACE_DIR / f"merge_result-{ts}.json"
- result_path.write_text(
- json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
- )
- # Also write/update the canonical result for quick access to the latest
- (WORKSPACE_DIR / MERGE_RESULT_FILE).write_text(
- json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
- )
-
- return result
-
-
-def run_merge_incremental(
- triton_ascend_path: Path,
- triton_path: Path,
- target_commit: str,
- work_branch: str,
-) -> dict:
- """Merge *target_commit* into an already-existing work branch.
-
- Used for progressive step-by-step merging: the first step calls
- run_merge() to create the work branch, and subsequent steps call
- run_merge_incremental() to merge their end_commit on top.
-
- Does NOT create a new branch or stash changes — it assumes we're
- already on the work branch from a previous step.
- """
- ascend_path = Path(triton_ascend_path)
-
- # Verify we're on the expected work branch
- current_branch = run_git(ascend_path, "branch", "--show-current").strip()
- if current_branch != work_branch:
- print(f"[merge] Switching from '{current_branch}' to work branch '{work_branch}'")
- run_git(ascend_path, "checkout", work_branch)
-
- # Fetch the target commit if needed
- try:
- run_git(ascend_path, "fetch", "upstream-triton", "--prune")
- except subprocess.CalledProcessError:
- print("[merge] Warning: could not fetch upstream-triton, assuming target is reachable")
-
- if Path(triton_path) != ascend_path:
- try:
- run_git(ascend_path, "fetch", str(triton_path), target_commit)
- except subprocess.CalledProcessError:
- print("[merge] Warning: could not fetch target from triton path")
-
- print(f"[merge] Incremental merge {target_commit[:12]} into {work_branch}")
- merge_proc = run_git_no_check(
- ascend_path, "merge", "--no-ff", "--no-edit", target_commit
- )
-
- # Use a timestamped log file so each merge step's output is preserved
- ts = datetime.now().strftime("%Y%m%d-%H%M%S")
- merge_log_path = WORKSPACE_DIR / f"merge-{ts}.log"
- merge_log_path.write_text(
- f"STDOUT:\n{merge_proc.stdout}\n\nSTDERR:\n{merge_proc.stderr}\n",
- encoding="utf-8",
- )
- # Also write/update the canonical merge log for quick access to the latest
- (WORKSPACE_DIR / MERGE_LOG_FILE).write_text(
- f"STDOUT:\n{merge_proc.stdout}\n\nSTDERR:\n{merge_proc.stderr}\n",
- encoding="utf-8",
- )
-
- has_conflicts = has_merge_conflicts(ascend_path)
- conflict_files = get_conflict_files(ascend_path) if has_conflicts else []
-
- conflict_dir = WORKSPACE_DIR / CONFLICT_LOG_DIR
- conflict_info = []
- if has_conflicts:
- conflict_dir.mkdir(parents=True, exist_ok=True)
- conflict_info = _save_conflict_info(ascend_path, conflict_files, conflict_dir)
-
- result = {
- "work_branch": work_branch,
- "original_branch": current_branch,
- "target_commit": target_commit,
- "merge_exit_code": merge_proc.returncode,
- "has_conflicts": has_conflicts,
- "conflict_files": conflict_files,
- "conflict_count": len(conflict_files),
- "conflicts": conflict_info,
- "merge_log": str(merge_log_path),
- "conflict_dir": str(conflict_dir) if has_conflicts else "",
- }
-
- result_path = WORKSPACE_DIR / f"merge_result-{ts}.json"
- result_path.write_text(
- json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
- )
- # Also write/update the canonical result for quick access to the latest
- (WORKSPACE_DIR / MERGE_RESULT_FILE).write_text(
- json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
- )
-
- return result
diff --git a/src/TA_main2main_workflow/scripts/plan_steps.py b/src/TA_main2main_workflow/scripts/plan_steps.py
deleted file mode 100644
index 37ffc80..0000000
--- a/src/TA_main2main_workflow/scripts/plan_steps.py
+++ /dev/null
@@ -1,380 +0,0 @@
-#!/usr/bin/env python3
-"""Deterministic step planner for the TA main2main upstream sync pipeline.
-
-Splits a range of upstream Triton commits into ordered steps based on changed
-lines in key source directories. Every commit between base and target is
-included — no commits are skipped, including those that touch zero source
-lines (they are still tracked but contribute 0 to the line budget).
-
-Algorithm (in priority order):
- 1. LLVM version change → solo step:
- If a commit modifies cmake/llvm-hash.txt it MUST be merged alone,
- regardless of its source-line count. Pending commits are flushed first.
- 2. Oversized single commit:
- A commit whose source lines exceed LINE_BUDGET becomes its own step.
- 3. Line-budget grouping:
- Commits accumulate into a step until source_changed_lines > LINE_BUDGET
- (no commit-count limit — as many commits as fit within the line budget).
-
-The LINE_BUDGET can be controlled via TA_LINE_BUDGET env var (default: 1000).
-
-Output:
- - /steps.json — machine-readable plan
- - /steps//upstream.patch — per-step upstream diff
- - /steps//changed_files.txt — per-step changed files
-"""
-
-from __future__ import annotations
-
-import json
-import os
-from pathlib import Path
-from typing import Any
-
-from TA_main2main_workflow.utils import (
- WORKSPACE_DIR, STEPS_FILE, STEPS_DIR, LINE_BUDGET, SOURCE_DIRS,
- LLVM_HASH_FILE, run_git,
-)
-
-
-def _list_commits(repo: Path, base: str, target: str) -> list[dict[str, str]]:
- """List all commits between base and target, ordered chronologically."""
- log_output = run_git(
- repo, "log", "--reverse", "--format=%H%x1f%s", f"{base}..{target}"
- )
- commits: list[dict[str, str]] = []
- for line in log_output.strip().splitlines():
- if not line.strip():
- continue
- parts = line.split("\x1f", 1)
- commits.append({
- "sha": parts[0].strip(),
- "subject": parts[1].strip() if len(parts) > 1 else "",
- })
- return commits
-
-
-def _source_lines_for_commit(repo: Path, sha: str) -> int:
- """Count changed lines in SOURCE_DIRS for a single commit using diff-tree."""
- total = 0
- for src_dir in SOURCE_DIRS:
- try:
- output = run_git(
- repo, "diff-tree", "--no-commit-id", "-r", "--numstat",
- sha, "--", f":(top){src_dir}",
- )
- except Exception:
- continue
- for line in output.strip().splitlines():
- if not line.strip():
- continue
- parts = line.split("\t")
- if len(parts) >= 3:
- added = int(parts[0]) if parts[0] != "-" else 0
- deleted = int(parts[1]) if parts[1] != "-" else 0
- total += added + deleted
- return total
-
-
-def _commit_changed_llvm_hash(repo: Path, sha: str) -> bool:
- """Check if a single commit modified cmake/llvm-hash.txt.
-
- Uses git diff-tree to list files changed by *sha*, then checks whether
- LLVM_HASH_FILE appears in the output. A commit that touches this file
- must become a solo step regardless of its source-line count.
- """
- try:
- output = run_git(repo, "diff-tree", "--no-commit-id", "--name-only", "-r", sha)
- return LLVM_HASH_FILE in output
- except Exception:
- return False
-
-
-def _make_step(
- index: int,
- commits: list[dict[str, str]],
- start_commit: str,
- total_lines: int,
- line_budget: int,
- reason: str = "line_budget",
-) -> dict[str, Any]:
- """Build a step dict from accumulated commits.
-
- The 'commits' field stores objects with 'sha' and 'subject' keys,
- matching the vllm-ascend main2main_flow format.
-
- *reason* explains why this step was formed:
- - ``"line_budget"`` — normal grouping by line budget
- - ``"llvm_version"`` — solo step because commit changed llvm-hash.txt
- - ``"oversized"`` — solo step because a single commit exceeds budget
- """
- return {
- "index": index,
- "id": f"step-{index}",
- "commits": commits, # list of {"sha": ..., "subject": ...}
- "commit_count": len(commits),
- "start_commit": start_commit,
- "end_commit": commits[-1]["sha"],
- "source_changed_lines": total_lines,
- "line_budget": line_budget,
- "reason": reason,
- }
-
-
-def _plan_steps(
- commits: list[dict[str, str]],
- lines_per_commit: dict[str, int],
- base_commit: str,
- line_budget: int = LINE_BUDGET,
- llvm_commits: set[str] | None = None,
-) -> list[dict[str, Any]]:
- """Group commits into steps with LLVM-aware planning.
-
- Every commit in the range is included — even those that touch zero
- source lines (they contribute 0 to the line budget and don't cause
- step splits on their own).
-
- Algorithm (in priority order):
- 1. **LLVM version change → solo step**: If a commit modifies
- ``cmake/llvm-hash.txt`` it MUST be merged alone, regardless of
- its source-line count. Pending commits are flushed first.
- 2. **Oversized single commit**: A commit whose source lines exceed
- LINE_BUDGET becomes its own step.
- 3. **Line-budget grouping**: Otherwise accumulate commits until
- ``step_lines + commit_lines > line_budget``, then flush.
- No commit-count cap — as many commits as fit within the budget.
- """
- if llvm_commits is None:
- llvm_commits = set()
-
- steps: list[dict[str, Any]] = []
- step_commits: list[dict[str, str]] = []
- step_lines = 0
- start = base_commit
-
- for commit in commits:
- sha = commit["sha"]
- lines = lines_per_commit.get(sha, 0)
- is_llvm_change = sha in llvm_commits
-
- # ── Rule 1.1: LLVM version change → solo step ──
- if is_llvm_change:
- if step_commits:
- steps.append(_make_step(
- len(steps) + 1, step_commits, start, step_lines,
- line_budget, reason="line_budget",
- ))
- start = steps[-1]["end_commit"]
- step_commits = []
- step_lines = 0
- steps.append(_make_step(
- len(steps) + 1, [commit], start, lines, line_budget,
- reason="llvm_version",
- ))
- start = steps[-1]["end_commit"]
- continue
-
- # ── Rule 2.1: Oversized single commit → solo step ──
- if lines > line_budget:
- if step_commits:
- steps.append(_make_step(
- len(steps) + 1, step_commits, start, step_lines,
- line_budget, reason="line_budget",
- ))
- start = steps[-1]["end_commit"]
- step_commits = []
- step_lines = 0
- steps.append(_make_step(
- len(steps) + 1, [commit], start, lines, line_budget,
- reason="oversized",
- ))
- start = steps[-1]["end_commit"]
- continue
-
- # ── Would exceed line budget → flush current step first ──
- if step_lines + lines > line_budget:
- steps.append(_make_step(
- len(steps) + 1, step_commits, start, step_lines,
- line_budget, reason="line_budget",
- ))
- start = steps[-1]["end_commit"]
- step_commits = []
- step_lines = 0
-
- step_commits.append(commit)
- step_lines += lines
-
- # ── Flush remaining ──
- if step_commits:
- steps.append(_make_step(
- len(steps) + 1, step_commits, start, step_lines,
- line_budget, reason="line_budget",
- ))
-
- return steps
-
-
-def _enrich_steps_with_diff(triton_path: Path, steps: list[dict[str, Any]]) -> None:
- """Add upstream diff and changed file list to each step.
-
- Filters to SOURCE_DIRS only so each step's patch is scoped to the
- code that actually needs adaptation (python/triton/, lib/, include/).
- Matches vllm-ascend's approach of filtering to vllm/.
- """
- # Build pathspec arg for git diff filtering: :(top)python/triton/ :(top)lib/ :(top)include/
- pathspec_args: list[str] = []
- for d in SOURCE_DIRS:
- pathspec_args.extend(["--", f":(top){d}"])
-
- for step in steps:
- step["upstream_patch"] = run_git(
- triton_path, "diff",
- f"{step['start_commit']}..{step['end_commit']}",
- *pathspec_args,
- )
- changed_files = run_git(
- triton_path, "diff", "--name-only",
- f"{step['start_commit']}..{step['end_commit']}",
- *pathspec_args,
- )
- step["changed_files"] = changed_files
- step["files_changed"] = sorted(
- f for f in changed_files.strip().splitlines() if f
- )
-
-
-def run_plan(
- triton_path: Path,
- base_commit: str,
- target_commit: str,
- line_budget: int | None = None,
-) -> dict[str, Any]:
- """Main entry point: plan steps and write steps.json + per-step artifacts.
-
- Args:
- triton_path: Path to the upstream Triton git repository.
- base_commit: Merge-base commit (start of the range).
- target_commit: Target upstream commit (end of the range).
- line_budget: Max source lines per step. Reads TA_LINE_BUDGET env var
- if omitted, falls back to LINE_BUDGET (1000).
-
- Steps are determined solely by the line budget — there is no
- commit-count limit. All commits between base and target are included.
-
- Returns:
- Plan dict with keys: base_commit, target_commit, total_commits, steps.
- """
- if line_budget is None:
- line_budget = int(os.getenv("TA_LINE_BUDGET", str(LINE_BUDGET)))
-
- commits = _list_commits(triton_path, base_commit, target_commit)
-
- print(f"[plan] Scanning {len(commits)} upstream commits "
- f"({base_commit[:8]}..{target_commit[:8]})")
- print(f"[plan] Line budget: {line_budget} (no commit-count limit)")
-
- # Count changed source lines per commit + detect LLVM version changes
- lines_per_commit: dict[str, int] = {}
- llvm_commits: set[str] = set()
- source_touching_count = 0
- for i, c in enumerate(commits):
- lines = _source_lines_for_commit(triton_path, c["sha"])
- lines_per_commit[c["sha"]] = lines
- if lines > 0:
- source_touching_count += 1
- # Rule 1.1: check if this commit changed cmake/llvm-hash.txt
- if _commit_changed_llvm_hash(triton_path, c["sha"]):
- llvm_commits.add(c["sha"])
- print(f"[plan] LLVM version change detected: {c['sha'][:8]} {c['subject'][:80]}")
- if (i + 1) % 50 == 0:
- print(f"[plan] ... scanned {i + 1}/{len(commits)} commits")
-
- if source_touching_count < len(commits):
- print(f"[plan] {len(commits) - source_touching_count} commits touch zero "
- f"source lines — included in steps with 0 line contribution")
-
- if llvm_commits:
- print(f"[plan] {len(llvm_commits)} commit(s) changed LLVM hash "
- f"— each will be a solo merge step")
-
- steps = _plan_steps(commits, lines_per_commit, base_commit, line_budget,
- llvm_commits=llvm_commits)
- _enrich_steps_with_diff(triton_path, steps)
-
- plan = {
- "base_commit": base_commit,
- "target_commit": target_commit,
- "line_budget": line_budget,
- "total_source_commits": source_touching_count,
- "total_commits": sum(s["commit_count"] for s in steps),
- "total_steps": len(steps),
- "steps": steps,
- }
-
- # ── Write steps.json ──
- steps_dir = WORKSPACE_DIR / STEPS_DIR
- steps_dir.mkdir(parents=True, exist_ok=True)
- (WORKSPACE_DIR / STEPS_FILE).write_text(
- json.dumps(plan, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
- )
-
- # ── Write per-step artifacts ──
- for step in steps:
- step_dir = steps_dir / step["id"]
- step_dir.mkdir(parents=True, exist_ok=True)
- (step_dir / "upstream.patch").write_text(
- step["upstream_patch"], encoding="utf-8"
- )
- (step_dir / "changed_files.txt").write_text(
- step["changed_files"], encoding="utf-8"
- )
- # Write a human-readable commit list for this step
- commit_list_lines = []
- for c in step["commits"]:
- commit_list_lines.append(f"{c['sha'][:8]} {c['subject']}")
- (step_dir / "commits.txt").write_text(
- "\n".join(commit_list_lines) + "\n", encoding="utf-8"
- )
-
- print(f"[plan] Generated {len(steps)} step(s) totaling "
- f"{plan['total_commits']} source-touching commits")
- for s in steps:
- reason_tag = ""
- if s.get("reason") == "llvm_version":
- reason_tag = " [LLVM VERSION]"
- elif s.get("reason") == "oversized":
- reason_tag = " [OVERSIZED]"
- print(f" {s['id']}: {s['commit_count']} commits, "
- f"{s['source_changed_lines']} lines "
- f"({'OVERSIZED' if s['source_changed_lines'] > line_budget else 'OK'})"
- f"{reason_tag}")
-
- return plan
-
-
-def plan_steps(
- triton_path: Path,
- base_commit: str,
- target_commit: str,
- line_budget: int | None = None,
-) -> list[dict[str, Any]]:
- """Public wrapper: plan steps and return the step list (for testing).
-
- Same as run_plan() but returns just the steps list instead of the full
- plan dict. Does NOT write files to disk — call run_plan() for that.
- """
- if line_budget is None:
- line_budget = int(os.getenv("TA_LINE_BUDGET", str(LINE_BUDGET)))
-
- commits = _list_commits(triton_path, base_commit, target_commit)
-
- lines_per_commit: dict[str, int] = {}
- llvm_commits: set[str] = set()
- for c in commits:
- lines = _source_lines_for_commit(triton_path, c["sha"])
- lines_per_commit[c["sha"]] = lines
- if _commit_changed_llvm_hash(triton_path, c["sha"]):
- llvm_commits.add(c["sha"])
-
- return _plan_steps(commits, lines_per_commit, base_commit, line_budget,
- llvm_commits=llvm_commits)
diff --git a/src/TA_main2main_workflow/scripts/pre_ci_check.py b/src/TA_main2main_workflow/scripts/pre_ci_check.py
deleted file mode 100644
index 0f67183..0000000
--- a/src/TA_main2main_workflow/scripts/pre_ci_check.py
+++ /dev/null
@@ -1,269 +0,0 @@
-#!/usr/bin/env python3
-"""Pre-CI verification for TA_main2main sync steps.
-
-Runs mechanical checks before build/test to catch common issues early:
- 1. Merge conflict marker check: no remaining <<<<<<< / ======= / >>>>>>> markers
- 2. Python syntax check: quick syntax validation on modified .py files
-
-Also provides cleanup_temp_files() to actively remove test artifacts
-(result_profiling/, *.lock, __pycache__/, *.pyc) before committing.
-
-All results are printed to the local console and written to workspace.
-"""
-
-from __future__ import annotations
-
-import ast
-import json
-import subprocess
-from pathlib import Path
-
-from TA_main2main_workflow.utils import (
- WORKSPACE_DIR, PRE_CI_CHECK_FILE, run_git_no_check,
- print_section, print_status, print_info, print_warn,
-)
-
-# Directories to purge (recursively removed if found under repo root)
-_CLEANUP_DIRS = [
- "result_profiling",
- "__pycache__",
- ".pytest_cache",
- ".mypy_cache",
- "*.egg-info",
-]
-
-# File patterns to purge (matched via glob **/*.suffix and exact basename)
-_CLEANUP_SUFFIXES = [
- ".lock",
- ".pyc",
- ".pyo",
- ".orig", # git merge conflict backups
- ".rej", # patch rejection files
- ".log", # log files that may leak into repo
-]
-
-_CLEANUP_BASENAMES = [
- ".DS_Store",
-]
-
-_CONFLICT_MARKERS = [
- "<<<<<<<",
- "=======",
- ">>>>>>>",
-]
-
-
-def _get_modified_files(repo: Path) -> list[str]:
- """Return list of modified (unstaged + staged) files."""
- modified: set[str] = set()
-
- result = run_git_no_check(repo, "diff", "--name-only", "HEAD")
- if result.stdout.strip():
- modified.update(result.stdout.strip().splitlines())
-
- result = run_git_no_check(repo, "diff", "--name-only", "--cached")
- if result.stdout.strip():
- modified.update(result.stdout.strip().splitlines())
-
- result = run_git_no_check(repo, "ls-files", "--others", "--exclude-standard")
- if result.stdout.strip():
- modified.update(result.stdout.strip().splitlines())
-
- return sorted(modified)
-
-
-def _check_conflict_markers(repo: Path, modified_files: list[str]) -> dict:
- """Scan modified files for remaining merge conflict markers."""
- violations: list[dict] = []
- for filepath in modified_files:
- full_path = repo / filepath
- if not full_path.exists() or not full_path.is_file():
- continue
- try:
- content = full_path.read_text(encoding="utf-8", errors="replace")
- except Exception:
- continue
- for lineno, line in enumerate(content.splitlines(), 1):
- for marker in _CONFLICT_MARKERS:
- if line.strip().startswith(marker):
- violations.append({
- "file": filepath,
- "line": lineno,
- "marker": marker,
- "text": line.strip()[:120],
- })
-
- return {
- "name": "conflict_markers",
- "passed": len(violations) == 0,
- "violations": violations,
- "detail": (
- "no remaining conflict markers"
- if len(violations) == 0
- else f"{len(violations)} conflict marker(s) still present"
- ),
- }
-
-
-def cleanup_temp_files(repo: Path) -> dict:
- """Actively remove test artifacts and temp files from the repository.
-
- Deletes directories like result_profiling/, __pycache__/ and files
- matching *.lock, *.pyc, etc. This prevents them from being accidentally
- committed via git add -u or git add -A.
-
- Only operates inside the repo (not outside it). Uses a whitelist of
- known-temp patterns — it will NOT delete arbitrary files.
-
- Returns a dict with counts of what was cleaned.
- """
- import shutil
-
- removed_dirs: list[str] = []
- removed_files: list[str] = []
-
- # ── Remove matching directories (recursively from repo root) ──
- for dirname in _CLEANUP_DIRS:
- for found in repo.rglob(dirname):
- if found.is_dir() and ".git" not in found.parts:
- try:
- shutil.rmtree(found, ignore_errors=True)
- removed_dirs.append(str(found.relative_to(repo)))
- except Exception:
- pass
-
- # ── Remove files by suffix ──
- for suffix in _CLEANUP_SUFFIXES:
- for found in repo.rglob(f"*{suffix}"):
- if found.is_file() and ".git" not in found.parts:
- try:
- found.unlink()
- removed_files.append(str(found.relative_to(repo)))
- except Exception:
- pass
-
- # ── Remove files by exact basename ──
- for name in _CLEANUP_BASENAMES:
- for found in repo.rglob(name):
- if found.is_file() and ".git" not in found.parts:
- try:
- found.unlink()
- removed_files.append(str(found.relative_to(repo)))
- except Exception:
- pass
-
- total = len(removed_dirs) + len(removed_files)
- if total > 0:
- print_info(f"Cleaned up {total} temp artifact(s):")
- for d in removed_dirs:
- print_info(f" rmdir: {d}")
- for f in removed_files:
- print_info(f" rm: {f}")
- else:
- print_info("No temp artifacts to clean up")
-
- return {
- "name": "cleanup_temp_files",
- "passed": True,
- "removed_dirs": removed_dirs,
- "removed_files": removed_files,
- "total_removed": total,
- }
-
-
-def _check_python_syntax(repo: Path, modified_files: list[str]) -> dict:
- """Quick Python syntax check on modified .py files."""
- violations: list[dict] = []
- py_files = [f for f in modified_files if f.endswith(".py")]
-
- for filepath in py_files:
- full_path = repo / filepath
- if not full_path.exists():
- continue
- try:
- source = full_path.read_text(encoding="utf-8")
- ast.parse(source, filename=filepath)
- except SyntaxError as e:
- violations.append({
- "file": filepath,
- "line": e.lineno or 0,
- "msg": str(e.msg),
- })
- except Exception:
- pass
-
- return {
- "name": "python_syntax",
- "passed": len(violations) == 0,
- "violations": violations,
- "detail": (
- f"all {len(py_files)} modified .py files pass syntax check"
- if len(violations) == 0
- else f"{len(violations)} file(s) have syntax errors"
- ),
- }
-
-
-def run_pre_ci_check(repo: Path, step_id: str = "") -> dict:
- """Run all pre-CI checks on the triton-ascend working tree.
-
- Returns a dict with 'all_passed' (bool) and 'checks' (list of check results).
- """
- print_section(f"Pre-CI Check{f' — {step_id}' if step_id else ''}")
-
- try:
- modified_files = _get_modified_files(repo)
- except subprocess.CalledProcessError as exc:
- print_warn(f"Could not list modified files: {exc.stderr}")
- return {"all_passed": True, "checks": [], "error": str(exc.stderr)}
-
- if not modified_files:
- print_info("No modified files — nothing to check")
- return {"all_passed": True, "checks": [], "modified_files_count": 0}
-
- print_info(f"Checking {len(modified_files)} modified file(s)")
-
- # ── Phase 0: active cleanup of known temp artifacts ──
- cleanup_temp_files(repo)
- # Re-scan modified files after cleanup (some may have been removed)
- try:
- modified_files = _get_modified_files(repo)
- except subprocess.CalledProcessError:
- pass
-
- checks: list[dict] = []
- all_passed = True
-
- conflict_check = _check_conflict_markers(repo, modified_files)
- checks.append(conflict_check)
- print_status(conflict_check["passed"], conflict_check["detail"])
- if not conflict_check["passed"]:
- all_passed = False
- for v in conflict_check["violations"]:
- print_warn(f" {v['file']}:{v['line']} — {v['marker']}")
-
- syntax_check = _check_python_syntax(repo, modified_files)
- checks.append(syntax_check)
- print_status(syntax_check["passed"], syntax_check["detail"])
- if not syntax_check["passed"]:
- all_passed = False
- for v in syntax_check["violations"]:
- print_warn(f" {v['file']}:{v['line']} — {v['msg']}")
-
- if all_passed:
- print_status(True, "All pre-CI checks passed")
- else:
- print_status(False, "Pre-CI checks found issues")
-
- result = {
- "all_passed": all_passed,
- "checks": checks,
- "modified_files_count": len(modified_files),
- }
-
- check_path = WORKSPACE_DIR / PRE_CI_CHECK_FILE
- check_path.write_text(
- json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
- )
-
- return result
diff --git a/src/TA_main2main_workflow/scripts/push_to_github.py b/src/TA_main2main_workflow/scripts/push_to_github.py
deleted file mode 100644
index b64038d..0000000
--- a/src/TA_main2main_workflow/scripts/push_to_github.py
+++ /dev/null
@@ -1,612 +0,0 @@
-#!/usr/bin/env python3
-"""Push the sync branch and create a GitHub Pull Request for triton-ascend.
-
-Steps:
- 1. Ensure gh CLI is authenticated.
- 2. Clean up temp files (result_profiling/, __pycache__/, *.lock, etc.).
- 3. Run pre-commit run --from-ref origin/main --to-ref HEAD.
- 4. If pre-commit auto-fixes files, amend the latest commit.
- 5. Push the work branch to origin.
- 6. Open a PR via gh pr create with [user](type) title format.
-
-Environment variables:
- PUSH_TO_GITHUB — must be "true" to proceed
- GITHUB_REPO — target repo "owner/name" (default: TecJesh/triton-ascend)
- GH_TOKEN — GitHub Personal Access Token (CI fallback)
- PR_AUTHOR — user tag in PR title, e.g. "TA" → [TA](sync) ... (default: git user)
- PR_TYPE — conventional commit type in PR title (default: "sync")
-"""
-
-from __future__ import annotations
-
-import json
-import os
-import subprocess
-import sys
-import time
-import urllib.request
-from datetime import datetime
-from pathlib import Path
-
-from TA_main2main_workflow.utils import (
- WORKSPACE_DIR, FINAL_TARGET_PATCH_FILE, FINAL_SUMMARY_FILE,
- run_git, run_git_no_check, print_error,
- ENV_BASE_BRANCH, get_base_branch_ref,
-)
-
-
-def _detect_origin_owner(repo: Path, remote: str = "origin") -> str:
- """Extract the GitHub owner from the origin remote URL.
-
- Handles direct GitHub URLs, SSH URLs, and proxy URLs
- (e.g. gh-proxy.test.osinfra.cn/https://github.com/owner/repo.git).
- """
- try:
- url = run_git(repo, "remote", "get-url", remote).strip()
- # Strip credentials
- if "@" in url:
- url = url.split("@", 1)[-1]
- # If URL is behind a proxy, extract the real GitHub path
- if "github.com/" in url:
- # e.g. gh-proxy.test.osinfra.cn/https://github.com/owner/repo.git
- url = url.split("github.com/", 1)[-1]
- elif "github.com:" in url:
- # e.g. git@github.com:owner/repo.git
- url = url.split("github.com:", 1)[-1]
- # Now url should be owner/repo or owner/repo.git
- url = url.replace("https://", "").replace("git@", "")
- if url.endswith(".git"):
- url = url[:-4]
- parts = url.split("/")
- if parts and parts[0]:
- return parts[0]
- except Exception:
- pass
- return ""
-
-
-def _detect_default_branch(repo: Path, remote: str = "origin") -> str:
- """Detect the default branch of the remote."""
- try:
- ref = run_git(repo, "symbolic-ref", f"refs/remotes/{remote}/HEAD").strip()
- return ref.rsplit("/", 1)[-1]
- except subprocess.CalledProcessError:
- return os.getenv("TA_PR_BASE_BRANCH", "upstream-sync")
-
-
-def _ensure_gh_auth(repo: Path) -> None:
- """Ensure GitHub CLI is ready for authenticated git push.
-
- When GH_TOKEN is set (PAT in CI), gh and git use it directly —
- no explicit login needed. Otherwise fall back to interactive auth.
-
- IMPORTANT: when the runner uses a git proxy (url.insteadOf), the origin
- URL may point to a non-GitHub host (e.g. gh-proxy.test.osinfra.cn).
- We use TWO separate auth paths to handle this:
- 1. gh auth login --with-token → tells gh CLI about github.com directly
- (does NOT look at git remotes — essential for gh pr create)
- 2. Embed token in origin URL → ensures git push works through the proxy
- (gh auth setup-git alone may not work with url.insteadOf rewriting)
- """
- gh_token = os.getenv("GH_TOKEN", "")
- if not gh_token:
- try:
- subprocess.run(
- ["gh", "auth", "status"],
- check=True, capture_output=True, text=True,
- )
- print("[push] gh CLI already authenticated.")
- except subprocess.CalledProcessError:
- print(
- "[push] gh not authenticated and GH_TOKEN not set. "
- "Run 'gh auth login' locally or set GH_TOKEN in CI.",
- file=sys.stderr,
- )
- sys.exit(1)
- subprocess.run(
- ["gh", "auth", "setup-git"],
- check=True, capture_output=True, text=True,
- )
- print("[push] Git credential helper configured (via gh auth setup-git).")
- return
-
- # ── GH_TOKEN is set ──
- print("[push] Using GH_TOKEN from environment")
-
- # Step 1: Explicitly login gh CLI against github.com.
- # This is essential when the git remote points to a proxy host —
- # gh needs to know about github.com independently of git remotes.
- result = subprocess.run(
- ["gh", "auth", "login", "--with-token", "--hostname", "github.com"],
- input=gh_token + "\n", text=True, capture_output=True,
- )
- if result.returncode == 0:
- print("[push] gh auth login --with-token: success")
- else:
- print(f"[push] gh auth login stderr: {result.stderr.strip()}")
-
- # Step 2: Verify the token works
- result = subprocess.run(
- ["gh", "auth", "status", "--hostname", "github.com"],
- capture_output=True, text=True,
- )
- print(f"[push] gh auth status: {result.stdout.strip()}")
- if result.returncode != 0:
- print(f"[push] gh auth status stderr: {result.stderr.strip()}")
-
- # Step 3: Configure git credential helper (best-effort).
- # This may fail when the git remote points to a proxy host that gh
- # doesn't recognize — but it's non-essential because Step 4 embeds
- # the token directly in the origin URL.
- result = subprocess.run(
- ["gh", "auth", "setup-git", "--hostname", "github.com"],
- capture_output=True, text=True,
- )
- if result.returncode == 0:
- print("[push] Git credential helper configured (via gh auth setup-git).")
- else:
- print(f"[push] gh auth setup-git skipped "
- f"(exit {result.returncode}): {result.stderr.strip()}")
-
- # Step 4: Embed token in origin URL so git push works through the proxy.
- # (gh auth setup-git may not help when url.insteadOf rewrites the host.)
- try:
- origin_url = run_git(repo, "remote", "get-url", "origin").strip()
- if origin_url.startswith("https://"):
- clean_url = origin_url.replace("https://", "", 1)
- if "@" in clean_url:
- clean_url = clean_url.split("@", 1)[1]
- new_url = f"https://x-access-token:{gh_token}@{clean_url}"
- run_git(repo, "remote", "set-url", "origin", new_url)
- safe = f"https://x-access-token:***@{clean_url}"
- print(f"[push] origin URL rewritten with token: {safe}")
- except Exception as exc:
- print(f"[push] Note: could not rewrite origin URL: {exc}")
-
-
-def _run_pre_commit_and_amend(repo: Path) -> bool:
- """Run pre-commit and amend the latest commit if auto-fixes were applied.
-
- Steps:
- 1. Clean temp files first (result_profiling/, __pycache__/, *.lock, *.pyc)
- 2. Run: pre-commit run --from-ref --to-ref HEAD
- 3. If pre-commit modified files → git add -u && git commit --amend --no-edit
- 4. Re-clean temp files after amend
-
- Returns True if pre-commit passed (with or without auto-fixes).
- Returns False if pre-commit found unfixable issues.
- """
- from TA_main2main_workflow.scripts.pre_ci_check import cleanup_temp_files
-
- base_ref = get_base_branch_ref()
-
- print("[push] ── Pre-commit check before PR ──")
-
- # ── Step 1: clean temp files ──
- print("[push] Cleaning temp files before pre-commit...")
- cleanup_temp_files(repo)
-
- # ── Step 2: run pre-commit ──
- print(f"[push] Running: pre-commit run --from-ref {base_ref} --to-ref HEAD")
- try:
- pc_proc = subprocess.run(
- ["pre-commit", "run", "--from-ref", base_ref, "--to-ref", "HEAD"],
- cwd=repo,
- capture_output=True,
- text=True,
- timeout=300,
- )
- except subprocess.TimeoutExpired:
- print("[push] ⚠ pre-commit timed out after 300s, continuing anyway")
- return True
- except FileNotFoundError:
- print("[push] ⚠ pre-commit not installed, skipping")
- return True
-
- # Print pre-commit output
- if pc_proc.stdout:
- print(pc_proc.stdout)
- if pc_proc.stderr:
- print(pc_proc.stderr, file=sys.stderr)
-
- precommit_passed = pc_proc.returncode == 0
-
- # ── Step 3: check if pre-commit modified any files ──
- status_proc = run_git_no_check(repo, "status", "--porcelain")
- has_modifications = bool(status_proc.stdout.strip())
-
- if has_modifications:
- print("[push] Pre-commit modified files, amending latest commit...")
- # Stage only tracked files to avoid temp artifacts
- run_git(repo, "add", "-u")
- try:
- run_git(repo, "commit", "--amend", "--no-edit")
- print("[push] Commit amended with pre-commit fixes.")
- except subprocess.CalledProcessError:
- print("[push] Nothing to amend (already clean)")
-
- # ── Step 4: re-clean temp files after amend ──
- cleanup_temp_files(repo)
- else:
- if precommit_passed:
- print("[push] Pre-commit passed, no modifications needed.")
- else:
- print("[push] ⚠ Pre-commit reported issues but no files were modified "
- "(may need manual review).")
-
- return True
-
-
-def _build_pr_title(target_commit: str = "") -> str:
- """Build PR title in conventional commit format.
-
- Example: [Sync](feat) Merge upstream triton commits (abc12345)
-
- Env vars:
- PR_AUTHOR — user tag (default: "Sync")
- PR_TYPE — conventional commit type (default: "feat")
- """
- author = os.getenv("PR_AUTHOR", "Sync").strip()
- pr_type = os.getenv("PR_TYPE", "feat").strip()
- if target_commit:
- return f"[{author}]({pr_type}) Merge upstream triton commits {target_commit[:8]}"
- ts = datetime.now().strftime("%Y%m%d-%H%M%S")
- return f"[{author}]({pr_type}) Merge upstream triton commits {ts}"
-
-
-def _create_pr_via_api(
- github_repo: str,
- title: str,
- body: str,
- head: str,
- base: str,
- token: str,
-) -> str:
- """Create a GitHub PR via the REST API directly.
-
- Uses the GitHub REST API (POST /repos/{owner}/{repo}/pulls) instead of
- gh CLI to avoid host-detection issues when git remotes are rewritten by
- url.insteadOf proxy.
- """
- url = f"https://api.github.com/repos/{github_repo}/pulls"
- payload = json.dumps({
- "title": title,
- "body": body,
- "head": head,
- "base": base,
- }).encode("utf-8")
-
- req = urllib.request.Request(
- url,
- data=payload,
- headers={
- "Accept": "application/vnd.github+json",
- "Authorization": f"Bearer {token}",
- "X-GitHub-Api-Version": "2022-11-28",
- "Content-Type": "application/json",
- },
- method="POST",
- )
-
- try:
- with urllib.request.urlopen(req, timeout=30) as resp:
- result = json.loads(resp.read().decode("utf-8"))
- pr_url = result.get("html_url", "")
- if not pr_url:
- raise RuntimeError(f"API response missing html_url: {result}")
- return pr_url
- except urllib.error.HTTPError as e:
- error_body = e.read().decode("utf-8", errors="replace")
- raise RuntimeError(
- f"GitHub API error {e.code}: {error_body}"
- ) from e
-
-
-def _create_pr_via_gh(
- github_repo: str,
- title: str,
- body: str,
- head_ref: str,
- base_branch: str,
-) -> str:
- """Create a GitHub PR via the gh CLI.
-
- Uses the user's GH_TOKEN (classic PAT with fork write access) to
- authenticate. The auto GITHUB_TOKEN from actions/checkout is scoped
- to the upstream repo only — GH_TOKEN overrides it so the PR can
- reference branches on the user's fork.
- """
- gh_token = os.environ.get("GH_TOKEN") or ""
- gh_cmd = [
- "gh", "pr", "create",
- "--title", title,
- "--body", body,
- "--head", head_ref,
- "--base", base_branch,
- "--repo", github_repo,
- ]
- print(f"[push] Running: GH_HOST=github.com {' '.join(gh_cmd)}")
- result = subprocess.run(
- gh_cmd,
- capture_output=True, text=True, timeout=60,
- env={**os.environ,
- "GITHUB_TOKEN": gh_token,
- "GH_TOKEN": gh_token},
- )
- if result.returncode != 0:
- raise RuntimeError(
- f"gh pr create failed (exit {result.returncode}): "
- f"{result.stderr.strip()}"
- )
- pr_url = result.stdout.strip()
- if not pr_url:
- raise RuntimeError("gh pr create returned empty output")
- return pr_url
-
-
-def push_and_create_pr(
- ascend_path: Path,
- github_repo: str = "triton-lang/triton-ascend",
- work_branch: str = "",
- summary_path: Path | None = None,
- target_commit: str = "",
-) -> str:
- """Push the current work branch and create a GitHub PR.
-
- Flow:
- 1. Authenticate gh CLI
- 2. Run pre-commit --from-ref --to-ref HEAD, amend if needed
- 3. Clean temp files
- 4. Commit any remaining uncommitted changes
- 5. Push work branch
- 6. Create PR with [user](type) title format
-
- Returns the PR URL, or "" on skip/failure.
- """
- repo = Path(ascend_path)
-
- if not work_branch:
- work_branch = run_git(repo, "branch", "--show-current").strip()
-
- # Fork owner for push and PR head. Defaults to TecJesh because
- # in CI origin points to triton-lang/triton-ascend (upstream),
- # so auto-detection would return the wrong owner.
- _fork_owner = os.environ.get("TA_FORK_OWNER") or "TecJesh"
-
- base_ref = get_base_branch_ref()
- try:
- merge_base = run_git(repo, "merge-base", base_ref, "HEAD").strip()
- except subprocess.CalledProcessError:
- merge_base = "HEAD~1"
-
- patch_content = run_git(repo, "diff", merge_base, "HEAD")
- patch_path = WORKSPACE_DIR / FINAL_TARGET_PATCH_FILE
- patch_path.write_text(patch_content, encoding="utf-8")
- print(f"[push] Cumulative patch written to {patch_path}")
-
- summary_file = summary_path or (WORKSPACE_DIR / FINAL_SUMMARY_FILE)
- if not summary_file.exists():
- summary_file.write_text(
- f"# Triton-Ascend Upstream Sync\n\n"
- f"Branch: `{work_branch}`\n"
- f"Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n",
- encoding="utf-8",
- )
-
- _ensure_gh_auth(repo)
-
- # ── Pre-commit check + amend before pushing ──
- _run_pre_commit_and_amend(repo)
-
- # ── Commit any remaining uncommitted changes (after pre-commit amend) ──
- status = run_git(repo, "status", "--porcelain").strip()
- if status:
- print("[push] Staging uncommitted changes...")
- # Use "git add -u" (tracked-only) to avoid staging test artifacts,
- # cache files, or other transient files created during the flow.
- run_git(repo, "add", "-u")
- commit_msg = f"sync: upstream triton merge ({datetime.now().strftime('%Y%m%d-%H%M%S')})"
- try:
- run_git(repo, "commit", "-s", "-m", commit_msg)
- print(f"[push] Committed: {commit_msg}")
- except subprocess.CalledProcessError:
- print("[push] Nothing to commit (already clean)")
-
- # ── Push ──
- print(f"[push] Pushing branch '{work_branch}' to origin...")
-
- # Debug: show what token / URL we're actually using
- print("[push] === DEBUG push environment ===")
- print(f"[push] GH_TOKEN set: {bool(os.getenv('GH_TOKEN'))}")
- print(f"[push] GITHUB_TOKEN set: {bool(os.getenv('GITHUB_TOKEN'))}")
- try:
- remote_url = run_git(repo, "remote", "get-url", "origin").strip()
- # Mask any embedded token
- if "@" in remote_url:
- safe_url = remote_url.split("@")[0].split(":")[-1] + "@" + remote_url.split("@")[1]
- else:
- safe_url = remote_url
- print(f"[push] origin URL: {safe_url}")
- print(f"[push] current branch: {run_git(repo, 'branch', '--show-current').strip()}")
- except Exception:
- pass
- print("[push] ==============================")
-
- # Push to the fork (same pattern as AscendNPU-IR submodule push).
- # Token embedded in the URL so the CI proxy can authenticate.
- _token = os.environ.get("GH_TOKEN") or ""
- if _token and _fork_owner:
- _fork_remote = "ta-fork-push"
- _fork_url = (
- f"https://x-access-token:{_token}@"
- f"gh-proxy.test.osinfra.cn/"
- f"https://github.com/{_fork_owner}/triton-ascend.git"
- )
- _last_push_error = ""
- for _attempt in range(1, 6):
- run_git_no_check(repo, "remote", "remove", _fork_remote)
- run_git(repo, "remote", "add", _fork_remote, _fork_url)
- _push_result = subprocess.run(
- ["git",
- "-c", "http.https://github.com/.extraheader=",
- "push", "--force-with-lease", _fork_remote, work_branch],
- cwd=str(repo), capture_output=True, text=True,
- )
- run_git(repo, "remote", "remove", _fork_remote)
- if _push_result.returncode == 0:
- if _push_result.stdout.strip():
- print(f"[push] stdout:\n{_push_result.stdout.strip()}")
- break
- _last_push_error = _push_result.stderr.strip() or "(no stderr)"
- print_error(
- f"[push] git push attempt {_attempt}/5 FAILED "
- f"(exit {_push_result.returncode}):\n{_last_push_error}"
- )
- if _attempt < 5:
- time.sleep(10 * _attempt)
- else:
- raise RuntimeError(
- f"git push failed after 5 attempts: {_last_push_error}")
- else:
- run_git(repo, "push", "-u", "origin", work_branch)
-
- # ── Create PR via gh CLI ──
- # gh infers the GitHub host from git remotes. In CI origin points
- # to the proxy, so we temporarily swap it to the fork URL (with
- # token) — gh recognizes github.com and GH_HOST isn't needed.
- base_branch = os.getenv("TA_PR_BASE_BRANCH", "upstream-sync")
- pr_description = summary_file.read_text(encoding="utf-8") if summary_file.exists() else ""
-
- _head = f"{_fork_owner}:{work_branch}" if _fork_owner else work_branch
- pr_title = _build_pr_title(target_commit)
-
- print(f"[push] Creating PR via gh CLI:")
- print(f" head = {_head}")
- print(f" base = {base_branch}")
- print(f" repo = {github_repo}")
-
- _saved_origin = run_git(repo, "config", "--get", "remote.origin.url").strip()
- _pr_origin = f"https://x-access-token:{_token}@github.com/{_fork_owner}/triton-ascend.git" if _token else f"https://github.com/{_fork_owner}/triton-ascend.git"
- run_git(repo, "remote", "set-url", "origin", _pr_origin)
-
- _last_pr_error = ""
- for _attempt in range(1, 6):
- try:
- pr_url = _create_pr_via_gh(
- github_repo=github_repo,
- title=pr_title,
- body=pr_description,
- head_ref=_head,
- base_branch=base_branch,
- )
- print(f"[push] PR created: {pr_url}")
- return pr_url
- except Exception as _e:
- _last_pr_error = str(_e)
- print_error(f"[push] PR create attempt {_attempt}/5 FAILED: "
- f"{_last_pr_error}")
- if _attempt < 5:
- time.sleep(10 * _attempt)
- finally:
- run_git(repo, "remote", "set-url", "origin", _saved_origin)
- raise RuntimeError(
- f"gh pr create failed after 5 attempts: {_last_pr_error}")
-
-
-def push_step_progress(
- ascend_path: Path,
- github_repo: str = "triton-lang/triton-ascend",
- work_branch: str = "",
- step_id: str = "",
- step_num: int = 1,
- total_steps: int = 1,
- pr_url: str = "",
-) -> str:
- """Push work-branch progress after a single step and create/update a PR.
-
- Called after each progressive step's commit. On the first call (pr_url
- is empty) it creates a new PR; on subsequent calls it just pushes —
- the existing PR picks up the new commits automatically.
-
- Returns the PR URL (new or existing).
- """
- repo = Path(ascend_path)
-
- if not work_branch:
- work_branch = run_git(repo, "branch", "--show-current").strip()
-
- _ensure_gh_auth(repo)
-
- # ── Generate step-aware patch ──
- patch_content = run_git(repo, "diff", get_base_branch_ref(), "HEAD")
- patch_path = WORKSPACE_DIR / FINAL_TARGET_PATCH_FILE
- patch_path.write_text(patch_content, encoding="utf-8")
-
- # ── Push ──
- print(f"[push] [{step_id}] Pushing branch '{work_branch}' to origin...")
- run_git(repo, "push", "-u", "origin", work_branch)
-
- # ── Create PR on first call only ──
- if not pr_url:
- base_branch = _detect_default_branch(repo)
- ts = datetime.now().strftime("%Y%m%d-%H%M%S")
- pr_title = (
- f"[Step {step_num}/{total_steps}] sync: upstream triton merge ({ts})"
- )
- pr_body = (
- f"## Progressive Sync — Step {step_num}/{total_steps}\n\n"
- f"**Work branch**: `{work_branch}`\n"
- f"**Target repo**: `{github_repo}`\n"
- f"**Generated**: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n\n"
- f"This PR will be updated as subsequent steps complete.\n"
- )
-
- print(f"[push] [{step_id}] Creating PR: {pr_title}")
- gh_cmd = [
- "gh", "pr", "create",
- "--title", pr_title,
- "--body", pr_body,
- "--head", work_branch,
- "--base", base_branch,
- "--repo", github_repo,
- ]
- result = subprocess.run(
- gh_cmd, check=True, capture_output=True, text=True, cwd=str(repo)
- )
- pr_url = result.stdout.strip()
- print(f"[push] [{step_id}] PR created: {pr_url}")
- else:
- print(f"[push] [{step_id}] Pushed to existing PR: {pr_url}")
-
- return pr_url
-
-
-def update_pr_description(
- ascend_path: Path,
- github_repo: str,
- pr_url: str,
- step_descriptions: list[str],
-) -> None:
- """Update the PR body with a summary of all completed steps."""
- if not pr_url:
- return
-
- body = (
- "# Triton-Ascend Progressive Upstream Sync\n\n"
- "## Completed Steps\n\n"
- )
- for desc in step_descriptions:
- body += f"- {desc}\n"
- body += (
- f"\n---\n"
- f"Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n"
- )
-
- try:
- subprocess.run(
- ["gh", "pr", "edit", pr_url, "--body", body, "--repo", github_repo],
- check=True, capture_output=True, text=True, cwd=str(ascend_path),
- )
- print(f"[push] Updated PR description: {pr_url}")
- except subprocess.CalledProcessError as e:
- print(f"[push] Warning: could not update PR description: {e}")
diff --git a/src/TA_main2main_workflow/scripts/update_commit_reference.py b/src/TA_main2main_workflow/scripts/update_commit_reference.py
deleted file mode 100644
index 5da0b8a..0000000
--- a/src/TA_main2main_workflow/scripts/update_commit_reference.py
+++ /dev/null
@@ -1,117 +0,0 @@
-#!/usr/bin/env python3
-"""Update version tracking references after a successful upstream sync.
-
-For Triton-Ascend (a fork of Triton), update the version tracking file
-(version.txt) to record the new upstream commit that was synced. Also
-creates a sync metadata file in the workspace for audit trail.
-
-Output:
- - Updated version.txt in triton-ascend repo
- - workspace/sync_meta.json with sync details
-"""
-
-from __future__ import annotations
-
-import json
-from datetime import datetime
-from pathlib import Path
-
-from TA_main2main_workflow.utils import (
- WORKSPACE_DIR, run_git, run_git_no_check,
- print_section, print_status, print_info, print_key_value,
-)
-
-
-def _read_version_file(repo: Path) -> str | None:
- """Read the current version.txt if it exists."""
- version_path = repo / "version.txt"
- if version_path.exists():
- return version_path.read_text(encoding="utf-8").strip()
- return None
-
-
-def _write_version_file(repo: Path, version: str) -> None:
- """Write the version.txt file."""
- version_path = repo / "version.txt"
- version_path.write_text(version + "\n", encoding="utf-8")
-
-
-def _get_commit_date(repo: Path, commit: str) -> str:
- """Get ISO date of a commit."""
- try:
- return run_git(repo, "log", "-1", "--format=%cI", commit).strip()
- except Exception:
- return ""
-
-
-def run_update(
- ascend_path: Path,
- old_commit: str,
- new_commit: str,
- work_branch: str = "",
-) -> dict:
- """Update version tracking after successful upstream sync.
-
- Args:
- ascend_path: Path to the triton-ascend repository
- old_commit: Previous upstream commit (merge-base before sync)
- new_commit: New upstream commit that was synced
- work_branch: Name of the work branch used for the sync
-
- Returns:
- dict with 'files_updated' list and 'sync_meta'
- """
- print_section("Update Commit Reference")
-
- files_updated: list[str] = []
-
- old_version = _read_version_file(ascend_path)
- short_sha = new_commit[:12]
- sync_date = datetime.now().strftime("%Y-%m-%d")
-
- if old_version:
- print_info(f"Current version.txt: {old_version}")
- else:
- print_info("No version.txt found — creating one")
-
- new_version = f"upstream-triton-{short_sha}-synced-{sync_date}"
- _write_version_file(ascend_path, new_version)
- files_updated.append("version.txt")
- print_status(True, f"version.txt updated: {new_version}")
-
- try:
- run_git(ascend_path, "add", "version.txt")
- except Exception:
- pass
-
- ascend_head = run_git(ascend_path, "rev-parse", "HEAD").strip()
- old_commit_date = _get_commit_date(ascend_path, old_commit)
- new_commit_date = _get_commit_date(ascend_path, new_commit)
-
- sync_meta = {
- "sync_date": sync_date,
- "old_upstream_commit": old_commit,
- "new_upstream_commit": new_commit,
- "old_commit_date": old_commit_date,
- "new_commit_date": new_commit_date,
- "triton_ascend_head": ascend_head,
- "work_branch": work_branch,
- "version_txt": new_version,
- }
-
- meta_path = WORKSPACE_DIR / "sync_meta.json"
- meta_path.write_text(
- json.dumps(sync_meta, indent=2, ensure_ascii=False) + "\n",
- encoding="utf-8",
- )
-
- print_key_value("Old upstream", f"{old_commit[:12]} ({old_commit_date[:10]})")
- print_key_value("New upstream", f"{new_commit[:12]} ({new_commit_date[:10]})")
- print_key_value("Ascend HEAD", ascend_head[:12])
- print_key_value("Sync metadata", str(meta_path))
- print_status(True, f"Updated {len(files_updated)} file(s)")
-
- return {
- "files_updated": files_updated,
- "sync_meta": sync_meta,
- }
diff --git a/src/TA_main2main_workflow/utils.py b/src/TA_main2main_workflow/utils.py
deleted file mode 100644
index bc9912f..0000000
--- a/src/TA_main2main_workflow/utils.py
+++ /dev/null
@@ -1,432 +0,0 @@
-"""Shared constants, git helpers, and console output formatting for TA_main2main_workflow."""
-
-import os
-import shutil
-import subprocess
-import time
-from datetime import datetime
-from pathlib import Path
-from typing import Any
-
-# ── Flow routing signals ─────────────────────────────────────────────────────
-UpgradeCompleted = "UpgradeCompleted"
-UpgradeFailed = "UpgradeFailed"
-HasNewCommits = "HasNewCommits"
-HasNoNewCommits = "HasNoNewCommits"
-MergeSuccess = "MergeSuccess"
-MergeConflict = "MergeConflict"
-TestsPassed = "TestsPassed"
-TestsFailed = "TestsFailed"
-
-# ── Workspace paths ──────────────────────────────────────────────────────────
-_PACKAGE_DIR = Path(__file__).resolve().parent # TA_main2main_workflow package dir
-_WORKSPACE_DEFAULT = _PACKAGE_DIR / "workspace"
-WORKSPACE_DIR = Path(os.getenv("TA_MAIN2MAIN_WORKSPACE", str(_WORKSPACE_DEFAULT)))
-REPOS_DIR_NAME = "repos"
-TRITON_REPO_NAME = "triton"
-TRITON_ASCEND_REPO_NAME = "triton-ascend"
-
-# ── Step-planning constants ──────────────────────────────────────────────────
-LINE_BUDGET = 1000
-BASE_LINE_BUDGET = 1000
-BASE_COMMIT_COUNT_BUDGET = 5 # deprecated — no longer used as a step limit
-# Directories in upstream triton whose changed lines count toward the budget
-SOURCE_DIRS = ["python/triton/", "lib/", "include/"]
-# File that tracks the LLVM version — commits that modify it get solo steps
-LLVM_HASH_FILE = "cmake/llvm-hash.txt"
-# Env var to control the line budget at runtime
-ENV_LINE_BUDGET = "TA_LINE_BUDGET"
-# Env var to control the commit-count budget at runtime
-ENV_COMMIT_BUDGET = "TA_COMMIT_BUDGET"
-# Env var to enable single-step mode (per-step merge → IR → build → test → fix)
-ENV_SINGLE_STEP_MODE = "TA_SINGLE_STEP_MODE"
-# Env var to control the base branch for work branches, diffs, and pre-commit.
-# Defaults to "main". Set to "master", "develop", etc. to use a different base.
-ENV_BASE_BRANCH = "TA_BASE_BRANCH"
-
-
-def get_base_branch_ref(remote: str = "origin") -> str:
- """Return the full base-branch ref (e.g. origin/main).
-
- Controlled by the TA_BASE_BRANCH env var; defaults to 'main'.
- """
- branch = os.getenv(ENV_BASE_BRANCH, "main")
- return f"{remote}/{branch}"
-
-# ── Output file names ────────────────────────────────────────────────────────
-DETECT_FILE = "detect.json"
-STEPS_FILE = "steps.json"
-MERGE_LOG_FILE = "merge.log"
-MERGE_RESULT_FILE = "merge_result.json"
-BUILD_LOG_FILE = "build.log"
-BUILD_RESULT_FILE = "build_result.json"
-TEST_RESULT_FILE = "test_result.json"
-CONFLICT_LOG_DIR = "conflicts"
-FIX_LOG_DIR = "fixes"
-STEPS_DIR = "steps"
-FINAL_SUMMARY_FILE = "final_summary.md"
-FINAL_TARGET_PATCH_FILE = "final_target.patch"
-EACH_STEP_SUMMARY_FILE = "step_summary.md"
-EACH_STEP_TARGET_PATCH_FILE = "step_target.patch"
-PRE_CI_CHECK_FILE = "pre_ci_check.json"
-CODE_STRUCTURE_GUIDE_FILE = "code-structure-guide.md"
-
-# ── IR Analysis & Patch paths ────────────────────────────────────────────────
-IR_ANALYSIS_DIR = "ir-analysis"
-IR_PATCHES_DIR = "ir-patches"
-IR_OPS_REPORT_FILE = "ops_report.json"
-IR_CHANGES_REPORT_FILE = "changes_report.json"
-IR_DIAGNOSIS_FILE = "ir_diagnosis.json"
-
-# ── Per-step LLVM change analysis (single-step mode) ─────────────────────────
-LLVM_CHANGE_ANALYSIS_DIR = "llvm_change_analysis"
-
-# ── IR Patch loop constants ──────────────────────────────────────────────────
-IR_MAX_ITERATIONS = 3
-
-# ── Timing tracker ───────────────────────────────────────────────────────────
-_phase_timers: dict[str, float] = {}
-_flow_start_time: float = 0.0
-
-
-def commit_count_budget(line_budget: int = LINE_BUDGET) -> int:
- """DEPRECATED: Steps are now determined solely by line budget.
-
- Kept for backward compatibility with existing step plan files.
- Returns a derived value from line_budget, but no longer used as a
- hard limit during step planning.
- """
- import math
- import os
- base = int(os.getenv(ENV_COMMIT_BUDGET, str(BASE_COMMIT_COUNT_BUDGET)))
- return max(1, round(base * math.sqrt(line_budget / BASE_LINE_BUDGET)))
-
-
-def _ts() -> str:
- return datetime.now().strftime("%H:%M:%S")
-
-
-# ═══════════════════════════════════════════════════════════════════════════════
-# Console Output Helpers — all progress printed locally, no CrewAI web UI needed
-# ═══════════════════════════════════════════════════════════════════════════════
-
-def print_header(title: str) -> None:
- width = 72
- print(f"\n╔{'═' * width}╗", flush=True)
- print(f"║ {title:^{width}} ║", flush=True)
- print(f"╚{'═' * width}╝", flush=True)
-
-
-def print_section(title: str) -> None:
- print(f"\n{'─' * 60}", flush=True)
- print(f" [{_ts()}] {title}", flush=True)
- print(f"{'─' * 60}", flush=True)
-
-
-def print_step(step_num: int, total: int, name: str) -> None:
- print(f"\n ▸ [{step_num}/{total}] {name} @ {_ts()}", flush=True)
-
-
-def print_status(ok: bool, msg: str) -> None:
- icon = "✔" if ok else "✘"
- print(f" {icon} {msg}", flush=True)
-
-
-def print_info(msg: str) -> None:
- print(f" ℹ {msg}", flush=True)
-
-
-def print_warn(msg: str) -> None:
- print(f" ⚠ {msg}", flush=True)
-
-
-def print_error(msg: str) -> None:
- print(f" ✘ {msg}", flush=True)
-
-
-def print_key_value(key: str, value: Any) -> None:
- print(f" {key}: {value}", flush=True)
-
-
-def print_separator() -> None:
- print(f" {'─' * 56}", flush=True)
-
-
-def print_flow_progress(phase: str, detail: str = "") -> None:
- msg = f"[{_ts()}] [{phase}] {detail}" if detail else f"[{_ts()}] [{phase}]"
- print(msg, flush=True)
-
-
-def start_timer(name: str) -> None:
- global _flow_start_time
- _phase_timers[name] = time.monotonic()
- if not _flow_start_time:
- _flow_start_time = time.monotonic()
-
-
-def stop_timer(name: str) -> float:
- start = _phase_timers.pop(name, None)
- if start is None:
- return 0.0
- elapsed = time.monotonic() - start
- print(f" ⏱ {name} took {elapsed:.1f}s", flush=True)
- return elapsed
-
-
-def print_elapsed_total() -> None:
- if _flow_start_time:
- total = time.monotonic() - _flow_start_time
- print(f"\n ⏱ Total elapsed: {total:.1f}s ({total/60:.1f}m)", flush=True)
-
-
-def print_summary_table(rows: list[tuple[str, str, str]]) -> None:
- status_icons = {"PASS": "✔", "FAIL": "✘", "SKIP": "○", "WARN": "⚠"}
- print(f"\n{'═' * 72}", flush=True)
- print(f" SYNC SUMMARY @ {_ts()}", flush=True)
- print(f"{'═' * 72}", flush=True)
- print(f" {'Phase':<30} {'Status':<8} {'Details'}", flush=True)
- print(f" {'─' * 30} {'─' * 8} {'─' * 32}", flush=True)
- for step, status, detail in rows:
- icon = status_icons.get(status, "?")
- print(f" {step:<30} {icon} {status:<5} {detail}", flush=True)
- print(f"{'═' * 72}", flush=True)
-
-
-def print_conflict_list(files: list[str]) -> None:
- if not files:
- print_info("No conflicts")
- return
- print(f" Conflicted files ({len(files)}):")
- for i, f in enumerate(files, 1):
- print(f" {i}. {f}")
-
-
-def print_ai_call_info(backend: str, mode: str, attempt: int, max_attempts: int) -> None:
- print(f"\n ╭─ AI Call ─────────────────────────────────────────────", flush=True)
- print(f" │ Backend: {backend}", flush=True)
- print(f" │ Mode: {mode}", flush=True)
- print(f" │ Attempt: {attempt}/{max_attempts}", flush=True)
- print(f" │ Time: {_ts()}", flush=True)
- print(f" ╰──────────────────────────────────────────────────────", flush=True)
-
-
-def print_ai_result(ok: bool, modified_files: list[str] = (), summary: str = "") -> None:
- icon = "✔" if ok else "✘"
- print(f"\n ╭─ AI Result ───────────────────────────────────────────", flush=True)
- print(f" │ Status: {icon} {'Success' if ok else 'Failed'}", flush=True)
- if modified_files:
- print(f" │ Modified files ({len(modified_files)}):", flush=True)
- for f in modified_files:
- print(f" │ • {f}", flush=True)
- if summary:
- preview = summary[:500] + "..." if len(summary) > 500 else summary
- print(f" │ Summary: {preview}", flush=True)
- print(f" ╰──────────────────────────────────────────────────────", flush=True)
-
-
-# ═══════════════════════════════════════════════════════════════════════════════
-# Git helpers
-# ═══════════════════════════════════════════════════════════════════════════════
-
-def run_git(repo: Path | str, *args: str) -> str:
- result = subprocess.run(
- ["git", *args],
- cwd=str(repo),
- check=True,
- capture_output=True,
- text=True,
- encoding="utf-8",
- errors="replace",
- )
- return result.stdout
-
-
-def run_git_no_check(repo: Path | str, *args: str) -> subprocess.CompletedProcess:
- return subprocess.run(
- ["git", *args],
- cwd=str(repo),
- capture_output=True,
- text=True,
- encoding="utf-8",
- errors="replace",
- )
-
-
-def is_git_url(path: str) -> bool:
- return path.startswith(("https://", "http://", "git@"))
-
-
-def clone_repo(url: str, target: str) -> None:
- print(f"[init] Cloning {url} → {target}")
- subprocess.run(["git", "clone", url, target], check=True)
-
-
-def resolve_path(raw: str, name: str) -> str:
- if is_git_url(raw):
- target = WORKSPACE_DIR / REPOS_DIR_NAME / name
- if target.exists():
- shutil.rmtree(target)
- target.mkdir(parents=True, exist_ok=True)
- clone_repo(raw, str(target))
- return str(target)
- return raw
-
-
-def get_repo_head(repo: Path) -> str:
- if not repo.exists():
- raise FileNotFoundError(f"Repository path does not exist: {repo}")
- return run_git(repo, "rev-parse", "HEAD").strip()
-
-
-def get_merge_base(repo: Path, commit_a: str, commit_b: str) -> str:
- return run_git(repo, "merge-base", commit_a, commit_b).strip()
-
-
-def has_merge_conflicts(repo: Path) -> bool:
- result = run_git_no_check(repo, "diff", "--name-only", "--diff-filter=U")
- return bool(result.stdout.strip())
-
-
-def get_conflict_files(repo: Path) -> list[str]:
- result = run_git(repo, "diff", "--name-only", "--diff-filter=U")
- return [f for f in result.strip().splitlines() if f]
-
-
-def get_modified_files(repo: Path, base_ref: str = "HEAD") -> list[str]:
- result = run_git(repo, "diff", "--name-only", base_ref)
- return [f for f in result.strip().splitlines() if f]
-
-
-def get_unstaged_diff(repo: Path) -> str:
- return run_git(repo, "diff")
-
-
-def get_staged_diff(repo: Path) -> str:
- return run_git(repo, "diff", "--cached")
-
-
-# ═══════════════════════════════════════════════════════════════════════════════
-# Submodule helpers — AscendNPU-IR
-# ═══════════════════════════════════════════════════════════════════════════════
-
-_ASCENDNPU_IR_SUBMODULE = "third_party/ascend/AscendNPU-IR"
-_ASCENDNPU_IR_REMOTE = "https://github.com/TecJesh/AscendNPU-IR.git"
-_ASCENDNPU_IR_REMOTE_NAME = "npuir-push" # dedicated remote, never overwrite origin
-
-
-def _submodule_path(repo: Path) -> Path:
- """Resolve the AscendNPU-IR submodule path."""
- return repo / _ASCENDNPU_IR_SUBMODULE
-
-
-def submodule_has_changes(repo: Path) -> bool:
- """Check if the AscendNPU-IR submodule has uncommitted changes."""
- sm = _submodule_path(repo)
- if not sm.exists():
- return False
- proc = run_git_no_check(sm, "status", "--porcelain")
- return bool(proc.stdout.strip())
-
-
-def commit_submodule(repo: Path, commit_msg: str) -> bool:
- """Commit uncommitted changes inside the AscendNPU-IR submodule.
-
- Stages ALL changes (including new files) with 'git add -A' and commits
- them. Uses -A (not -u) so AI-created files are not silently dropped.
- Returns True if a new commit was created.
- """
- sm = _submodule_path(repo)
- if not sm.exists():
- print_info("[submodule] AscendNPU-IR submodule not found — skipping")
- return False
-
- if not submodule_has_changes(repo):
- print_info("[submodule] No uncommitted changes in AscendNPU-IR")
- return False
-
- print_section("Commit AscendNPU-IR Submodule")
- try:
- run_git(sm, "add", "-A")
- # Show what will be committed
- staged = run_git(sm, "diff", "--cached", "--name-only").strip()
- if staged:
- print_info(f"[submodule] Files staged ({len(staged.splitlines())}):")
- for f in staged.splitlines()[:10]:
- print_info(f" - {f}")
- run_git(sm, "commit", "-s", "-m", commit_msg)
- new_head = run_git(sm, "rev-parse", "HEAD").strip()
- print_status(True, f"Committed AscendNPU-IR: {new_head[:12]}")
- return True
- except subprocess.CalledProcessError as e:
- stderr = (e.stderr or "").strip() if hasattr(e, 'stderr') else str(e)
- if "nothing to commit" in stderr.lower():
- print_info("[submodule] Nothing to commit")
- return False
- print_warn(f"Could not commit AscendNPU-IR submodule: {stderr[-200:]}")
- return False
-
-
-def push_submodule(
- repo: Path,
- branch: str,
- remote: str = _ASCENDNPU_IR_REMOTE,
- remote_name: str = _ASCENDNPU_IR_REMOTE_NAME,
- force: bool = False,
-) -> bool:
- """Push the AscendNPU-IR submodule branch to its remote.
-
- Sets up the remote if it doesn't exist, then pushes the given branch.
- By default uses force-with-lease for safety; pass force=True for --force.
- Returns True on success.
- """
- sm = _submodule_path(repo)
- if not sm.exists():
- print_warn("[submodule] AscendNPU-IR submodule not found — cannot push")
- return False
-
- print_section("Push AscendNPU-IR Submodule")
-
- # ── Set up dedicated push remote (never touch origin) ──
- # Remove stale npuir-push if it exists, then add fresh
- run_git_no_check(sm, "remote", "remove", remote_name)
- print_info(f"[submodule] Adding push remote '{remote_name}' → {remote}")
- run_git(sm, "remote", "add", remote_name, remote)
-
- # ── Create branch at current HEAD (don't switch, stay on detached HEAD) ──
- current_head = run_git(sm, "rev-parse", "HEAD").strip()
- proc = run_git_no_check(sm, "branch", "-f", branch, current_head)
- if proc.returncode != 0:
- print_error(f"[submodule] Failed to create branch '{branch}': "
- f"{proc.stderr.strip()}")
- return False
-
- # ── Configure auth ──
- gh_token = os.getenv("GH_TOKEN", "")
- if gh_token:
- try:
- current_url = run_git(sm, "remote", "get-url", remote_name).strip()
- if current_url.startswith("https://") and "x-access-token" not in current_url:
- clean_url = current_url.replace("https://", "", 1)
- if "@" in clean_url:
- clean_url = clean_url.split("@", 1)[1]
- new_url = f"https://x-access-token:{gh_token}@{clean_url}"
- run_git(sm, "remote", "set-url", remote_name, new_url)
- safe = f"https://x-access-token:***@{clean_url}"
- print_info(f"[submodule] Remote URL rewritten: {safe}")
- except Exception as exc:
- print_warn(f"[submodule] Could not configure remote auth: {exc}")
-
- # ── Push ──
- try:
- push_args = ["push"]
- if force:
- push_args.append("--force")
- else:
- push_args.append("--force-with-lease")
- push_args.extend([remote_name, branch])
- run_git(sm, *push_args)
- print_status(True, f"Pushed AscendNPU-IR branch '{branch}' to {remote_name}")
- return True
- except Exception as e:
- print_error(f"[submodule] Failed to push AscendNPU-IR: {e}")
- return False
diff --git a/src/TA_main2main_workflow/utils/__init__.py b/src/TA_main2main_workflow/utils/__init__.py
new file mode 100644
index 0000000..6b15903
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/__init__.py
@@ -0,0 +1,82 @@
+"""Utility package for TA_main2main_workflow."""
+
+from __future__ import annotations
+
+import os
+from pathlib import Path
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Workspace paths
+# ═══════════════════════════════════════════════════════════════════════════
+
+WORKSPACE_DIR = Path(os.getenv("TA_MAIN2MAIN_WORKSPACE", str(Path.cwd() / "workspace")))
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Flow routing signals
+# ═══════════════════════════════════════════════════════════════════════════
+
+UpgradeCompleted = "UpgradeCompleted"
+UpgradeFailed = "UpgradeFailed"
+HasNewCommits = "HasNewCommits"
+HasNoNewCommits = "HasNoNewCommits"
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Step-planning constants
+# ═══════════════════════════════════════════════════════════════════════════
+
+LLVM_HASH_FILE = "cmake/llvm-hash.txt"
+LINE_BUDGET = 1000
+SOURCE_DIRS = ["python/triton/", "lib/", "include/"]
+ENV_SINGLE_STEP_MODE = "TA_SINGLE_STEP_MODE"
+ENV_BASE_BRANCH = "TA_BASE_BRANCH"
+
+# Baseline LLVM version that Ascend backend OP usage is built against
+_ASCEND_BASELINE_LLVM_HASH = "b5cc222d7429fe6f18c787f633d5262fac2e676f"
+
+
+def get_base_branch_ref(remote: str = "origin") -> str:
+ branch = os.getenv(ENV_BASE_BRANCH, "upstream_sync")
+ return f"{remote}/{branch}"
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Output file names
+# ═══════════════════════════════════════════════════════════════════════════
+
+DETECT_FILE = "detect.json"
+STEPS_FILE = "steps.json"
+BUILD_LOG_FILE = "build.log"
+BUILD_RESULT_FILE = "build_result.json"
+TEST_RESULT_FILE = "test_result.json"
+MERGE_LOG_FILE = "merge.log"
+MERGE_RESULT_FILE = "merge_result.json"
+FIX_LOG_DIR = "fixes"
+STEPS_DIR = "steps"
+CONFLICT_LOG_DIR = "conflicts"
+FINAL_SUMMARY_FILE = "final_summary.md"
+FINAL_TARGET_PATCH_FILE = "final_target.patch"
+EACH_STEP_SUMMARY_FILE = "step_summary.md"
+EACH_STEP_TARGET_PATCH_FILE = "step_target.patch"
+PRE_CI_CHECK_FILE = "pre_ci_check.json"
+CODE_STRUCTURE_GUIDE_FILE = "code_structure.md"
+
+# ═══════════════════════════════════════════════════════════════════════════
+# IR analysis / patch file names
+# ═══════════════════════════════════════════════════════════════════════════
+
+IR_ANALYSIS_DIR = "ir_analysis"
+IR_PATCHES_DIR = "ir_patches"
+IR_OPS_REPORT_FILE = "ops_report.json"
+IR_CHANGES_REPORT_FILE = "changes_report.json"
+IR_DIAGNOSIS_FILE = "diagnosis.json"
+IR_MAX_ITERATIONS = 3
+LLVM_CHANGE_ANALYSIS_DIR = "llvm_change_analysis"
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Re-exports
+# ═══════════════════════════════════════════════════════════════════════════
+
+from TA_main2main_workflow.utils.config import TAConfig # noqa: F401, E402
+from TA_main2main_workflow.utils.context import WorkflowContext # noqa: F401, E402
+from TA_main2main_workflow.utils.git import run_git, run_git_no_check # noqa: F401, E402
+from TA_main2main_workflow.utils.logging import get_logger, TALogger # noqa: F401, E402
diff --git a/src/TA_main2main_workflow/utils/config.py b/src/TA_main2main_workflow/utils/config.py
new file mode 100644
index 0000000..695d71b
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/config.py
@@ -0,0 +1,204 @@
+"""Configuration for TA_main2main_workflow.
+
+Only user-configurable parameters. Fixed paths inside triton-ascend
+repo are defined where they're used, not here.
+
+Priority: CLI args > env vars > defaults
+"""
+
+from __future__ import annotations
+
+import os
+from dataclasses import dataclass, field
+from pathlib import Path
+from typing import Literal
+
+
+AIBackendChoice = Literal["opencode", "claude", "auto"]
+
+
+@dataclass
+class TAConfig:
+ """User-configurable parameters for a workflow run."""
+
+ # ── Repository ────────────────────────────────────────────────────────
+ triton_ascend_path: str = "" # local path (skip clone if set)
+ triton_ascend_url: str = "https://github.com/triton-lang/triton-ascend.git"
+ triton_path: str = "" # local triton checkout (for offline/separate-history)
+ triton_upstream_url: str = "https://github.com/triton-lang/triton.git"
+ target_commit: str = ""
+
+ # ── AI Backend ────────────────────────────────────────────────────────
+ ai_backend: AIBackendChoice = "auto"
+ ai_timeout_minutes: int = 30
+ ai_stale_seconds: int = 1200
+ ai_max_stale_retries: int = 3
+
+ # ── Retry / Budget ────────────────────────────────────────────────────
+ max_retries: int = 10
+ line_budget: int = 1000
+
+ # ── Build / Test parallelism ──────────────────────────────────────────
+ llvm_install_prefix: str = ""
+ llvm_repo_url: str = "https://github.com/llvm/llvm-project.git"
+ build_procs: int = 32
+ test_procs: int = 16
+
+ # ── Skip flags ────────────────────────────────────────────────────────
+ resume: bool = False # skip steps whose output already exists
+ skip_ai_analysis: bool = False
+ skip_build: bool = False
+ skip_e2e_test: bool = False
+ skip_llvm_rebuild: bool = False # skip LLVM rebuild when version changes (IR patch)
+ skip_baseline_llvm: bool = False # skip initial baseline LLVM build at workflow start
+ skip_ir_patch: bool = False # skip entire IR patch phase (SKIP_IR_PATCH)
+
+ # ── Git / Branch ──────────────────────────────────────────────────────
+ base_branch: str = "upstream_sync"
+ work_branch_base: str = "upstream-ascend"
+ progressive_merge: bool = True
+
+ # ── PR / Push ─────────────────────────────────────────────────────────
+ push_to_github: bool = False
+ github_repo: str = "triton-lang/triton-ascend"
+
+ # ── LLVM workspace ────────────────────────────────────────────────────
+ llvm_project_path: str = "~/llvm-project"
+ llvm_install_prefix_sync: str = "~/llvm-install-sync"
+
+ # ── Conda / Python ────────────────────────────────────────────────────
+ conda_env: str = "ta-upgrade"
+ test_dir: str = "third_party/ascend/unittest/pytest_ut"
+ test_dirs: list[str] = field(default_factory=list) # resolved from env
+ test_command: str = "" # full shell command override (TA_TEST_COMMAND)
+ python_exe: str = ""
+
+ # ── Single-step mode (always enabled) ─────────────────────────────────
+ single_step_mode: bool = True
+
+ # ── IR patch ──────────────────────────────────────────────────────────
+ ir_max_iterations: int = 3
+
+ # ═══════════════════════════════════════════════════════════════════════
+ @classmethod
+ def from_env(cls) -> TAConfig:
+ return cls(
+ triton_ascend_path=os.getenv("TRITON_ASCEND_PATH", ""),
+ triton_ascend_url=os.getenv(
+ "TRITON_ASCEND_URL", "https://github.com/triton-lang/triton-ascend.git"
+ ),
+ triton_path=os.getenv("TRITON_PATH", ""),
+ triton_upstream_url=os.getenv(
+ "TRITON_UPSTREAM_URL", "https://github.com/triton-lang/triton.git"
+ ),
+ target_commit=os.getenv("TRITON_TARGET_COMMIT", ""),
+ ai_backend=_env_choice(
+ "AI_BACKEND", ["opencode", "claude", "auto"], "auto"
+ ),
+ ai_timeout_minutes=_env_int("TA_AI_TIMEOUT_MINUTES", 30),
+ ai_stale_seconds=_env_int("TA_AI_STALE_SECONDS", 1200),
+ ai_max_stale_retries=_env_int("TA_AI_MAX_STALE_RETRIES", 3),
+ max_retries=_env_int("TA_MAX_RETRIES", 10),
+ line_budget=_env_int("TA_LINE_BUDGET", 1000),
+ llvm_install_prefix=os.getenv("LLVM_INSTALL_PREFIX", ""),
+ llvm_repo_url=os.getenv(
+ "LLVM_REPO_URL", "https://github.com/llvm/llvm-project.git"
+ ),
+ build_procs=_env_int_fallback("BUILD_PROCS", "MAX_JOBS", 32),
+ test_procs=_env_int_fallback("TEST_PROCS", "NUM_PROCS", 16),
+ resume=_env_bool("TA_RESUME", False),
+ skip_ai_analysis=_env_bool("SKIP_AI_ANALYSIS", False),
+ skip_build=_env_bool("SKIP_BUILD", False),
+ skip_e2e_test=_env_bool("SKIP_E2E_TEST", False),
+ skip_llvm_rebuild=_env_bool("SKIP_LLVM_REBUILD", False),
+ skip_baseline_llvm=_env_bool("SKIP_BASELINE_LLVM", False),
+ skip_ir_patch=_env_bool("SKIP_IR_PATCH", False),
+ base_branch=os.getenv("TA_BASE_BRANCH", "upstream_sync"),
+ work_branch_base=os.getenv("TA_WORK_BRANCH_BASE", "upstream-ascend"),
+ progressive_merge=_env_bool("TA_PROGRESSIVE_MERGE", True),
+ push_to_github=_env_bool("PUSH_TO_GITHUB", False),
+ github_repo=os.getenv("GITHUB_REPO", "triton-lang/triton-ascend"),
+ llvm_project_path=os.getenv("LLVM_PROJECT_PATH", "~/llvm-project"),
+ llvm_install_prefix_sync=os.getenv("LLVM_INSTALL_PREFIX_SYNC", "~/llvm-install-sync"),
+ conda_env=os.getenv("CONDA_ENV", "ta-upgrade"),
+ test_dir=os.getenv("TA_TEST_DIR", "third_party/ascend/unittest/pytest_ut"),
+ test_dirs=_resolve_test_dirs(),
+ test_command=os.getenv("TA_TEST_COMMAND", ""),
+ python_exe=os.getenv("PYTHON", ""),
+ single_step_mode=_env_bool("TA_SINGLE_STEP_MODE", True),
+ ir_max_iterations=_env_int("TA_IR_MAX_ITERATIONS", 3),
+ )
+
+ @property
+ def llvm_project(self) -> Path:
+ if self.llvm_project_path:
+ return Path(os.path.expanduser(self.llvm_project_path))
+ return Path(os.path.expanduser(os.getenv("LLVM_PROJECT_PATH", "~/llvm-project")))
+
+ @property
+ def llvm_install(self) -> Path:
+ if self.llvm_install_prefix_sync:
+ return Path(os.path.expanduser(self.llvm_install_prefix_sync))
+ return Path(os.path.expanduser(
+ os.getenv("LLVM_INSTALL_PREFIX_SYNC", "~/llvm-install-sync")))
+
+
+def _env_bool(name: str, default: bool) -> bool:
+ val = os.getenv(name, "").lower()
+ if val in ("true", "1", "yes"):
+ return True
+ if val in ("false", "0", "no"):
+ return False
+ return default
+
+
+def _env_int(name: str, default: int) -> int:
+ try:
+ return int(os.getenv(name, str(default)))
+ except (TypeError, ValueError):
+ return default
+
+
+def _env_int_fallback(name: str, legacy_name: str, default: int) -> int:
+ """Read *name* first; if unset, fall back to *legacy_name*."""
+ val = os.getenv(name, "")
+ if val:
+ try:
+ return int(val)
+ except (TypeError, ValueError):
+ pass
+ val = os.getenv(legacy_name, "")
+ if val:
+ try:
+ return int(val)
+ except (TypeError, ValueError):
+ pass
+ return default
+
+
+def _resolve_test_dirs(primary: str = "", extra: str = "") -> list[str]:
+ """Build the ordered list of test directories to run.
+
+ ``TA_TEST_DIR`` provides the primary directory (default
+ ``third_party/ascend/unittest/pytest_ut``). ``TA_EXTRA_TEST_DIRS``
+ adds additional directories, comma- or space-separated.
+
+ When called from main.py with CLI args, *primary* and *extra* override
+ the env vars.
+ """
+ _primary = primary or os.getenv("TA_TEST_DIR", "third_party/ascend/unittest/pytest_ut")
+ dirs = [_primary] if _primary else []
+
+ extra_raw = extra or os.getenv("TA_EXTRA_TEST_DIRS", "")
+ if extra_raw:
+ for part in extra_raw.replace(",", " ").split():
+ part = part.strip()
+ if part and part not in dirs:
+ dirs.append(part)
+
+ return dirs
+
+
+def _env_choice(name: str, choices: list[str], default: str) -> str:
+ val = os.getenv(name, default).lower()
+ return val if val in choices else default
diff --git a/src/TA_main2main_workflow/utils/context.py b/src/TA_main2main_workflow/utils/context.py
new file mode 100644
index 0000000..41e1d06
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/context.py
@@ -0,0 +1,106 @@
+"""WorkflowContext — shared state carrier between pipeline steps.
+
+A flat dataclass that each pipeline step reads from and returns an updated
+copy of. Steps never mutate the context in place.
+"""
+
+from __future__ import annotations
+
+from dataclasses import dataclass, field, replace
+from pathlib import Path
+from typing import Any
+
+
+@dataclass
+class WorkflowContext:
+ """All mutable state that flows through the sync pipeline.
+
+ Each step function takes a ``WorkflowContext``, reads what it needs,
+ and returns a **new** instance with updated fields (via
+ :meth:`copy_with`). This makes data flow explicit and testable.
+ """
+
+ # ── Input configuration (set once at start) ────────────────────────────
+ triton_ascend_path: str = ""
+ triton_path: str = ""
+
+ # ── Remote names (set by prepare step) ─────────────────────────────────
+ origin_remote: str = "origin"
+ upstream_remote: str = "triton-upstream"
+
+ # ── Git state ──────────────────────────────────────────────────────────
+ merge_base: str = ""
+ ascend_head: str = ""
+ target_commit: str = ""
+ work_branch: str = ""
+ original_branch: str = ""
+
+ # ── Detection results (produced by detect step) ───────────────────────
+ upstream_commits: list[dict] = field(default_factory=list)
+ upstream_commits_count: int = 0
+ changed_files_count: int = 0
+ changed_lines_total: int = 0
+ has_new_commits: bool = False
+
+ # ── Step plan (produced by plan step) ──────────────────────────────────
+ steps: list[dict] = field(default_factory=list)
+ total_steps: int = 0
+ current_step: int = 0
+ step_start_ascend_head: str = ""
+
+ # ── Merge results (produced by merge step) ────────────────────────────
+ merge_has_conflicts: bool = False
+ conflict_files: list[str] = field(default_factory=list)
+
+ # ── Build / test results ──────────────────────────────────────────────
+ build_passed: bool = False
+ test_passed: bool = False
+ pytest_passed: bool = False
+ fix_errors: list[str] = field(default_factory=list)
+ test_log_dir: str = ""
+ test_failures_by_python: dict = field(default_factory=dict)
+
+ # ── Fix tracking ──────────────────────────────────────────────────────
+ build_fix_count: int = 0
+ test_fix_count: int = 0
+ conflict_files_resolved: int = 0
+ retry_count: int = 0
+ fix_attempts: list[dict] = field(default_factory=list)
+
+ # ── IR patch state ────────────────────────────────────────────────────
+ ir_analysis_done: bool = False
+ ir_ops_report: dict = field(default_factory=dict)
+ ir_changes_report: dict = field(default_factory=dict)
+ ir_patches: list = field(default_factory=list)
+ ir_patch_iteration: int = 0
+ ir_max_iterations: int = 3
+ ir_issues_found: int = 0
+ ir_fix_count: int = 0
+ llvm_hash_changed: bool = False
+ ir_loop_details: list[dict] = field(default_factory=list)
+
+ # ── Step tracking / reporting ─────────────────────────────────────────
+ step_details: list[dict] = field(default_factory=list)
+ step_pr_descriptions: list[str] = field(default_factory=list)
+ summary_rows: list[tuple] = field(default_factory=list)
+
+ # ── Final state ───────────────────────────────────────────────────────
+ final_status: str = ""
+ pr_url: str = ""
+
+ # ═══════════════════════════════════════════════════════════════════════
+ # Helpers
+ # ═══════════════════════════════════════════════════════════════════════
+
+ def copy_with(self, **kwargs: Any) -> WorkflowContext:
+ """Return a new WorkflowContext with the given fields updated.
+
+ Usage::
+
+ ctx = ctx.copy_with(build_passed=True, retry_count=1)
+ """
+ return replace(self, **kwargs)
+
+ @property
+ def ascend_path(self) -> Path:
+ return Path(self.triton_ascend_path)
diff --git a/src/TA_main2main_workflow/utils/git.py b/src/TA_main2main_workflow/utils/git.py
new file mode 100644
index 0000000..a3ee663
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/git.py
@@ -0,0 +1,86 @@
+"""Git helpers with automatic retry for transient failures.
+
+``run_git`` raises on non-zero exit (after retries for
+fetch/clone/push/pull). ``run_git_no_check`` returns the
+``CompletedProcess`` and never raises.
+"""
+
+from __future__ import annotations
+
+import subprocess
+import time
+from pathlib import Path
+
+RETRYABLE_OPERATIONS = ("fetch", "clone", "push", "pull", "ls-remote")
+MAX_RETRIES = 5
+RETRY_DELAY = 3 # seconds
+
+
+def run_git(repo: Path | str, *args: str) -> str:
+ """Run a git command in *repo*, return stdout, raise on failure.
+
+ Commands that start with fetch/clone/push/pull/ls-remote are
+ automatically retried up to *MAX_RETRIES* times on failure.
+ """
+ repo = Path(repo)
+ cmd = ["git", "-C", str(repo), *args]
+
+ is_retryable = args and args[0] in RETRYABLE_OPERATIONS
+
+ for attempt in range(1, MAX_RETRIES + 1 if is_retryable else 2):
+ result = subprocess.run(
+ cmd, capture_output=True, text=True, timeout=600,
+ )
+ if result.returncode == 0:
+ return result.stdout
+ if not is_retryable or attempt == MAX_RETRIES:
+ raise RuntimeError(
+ f"git {args[0]} failed (exit {result.returncode}):\n"
+ f"{result.stderr.strip()}"
+ )
+ time.sleep(RETRY_DELAY)
+ return "" # unreachable
+
+
+def run_git_no_check(repo: Path | str, *args: str) -> subprocess.CompletedProcess:
+ """Run a git command in *repo*, return ``CompletedProcess``, never raise."""
+ repo = Path(repo)
+ cmd = ["git", "-C", str(repo), *args]
+ return subprocess.run(cmd, capture_output=True, text=True, timeout=600)
+
+
+def stream_cmd(cmd: list[str], cwd: Path, log_fh, timeout: int,
+ label: str = "", env: dict | None = None) -> int:
+ """Stream subprocess output line-by-line to console and log file.
+
+ Each output line is:
+ - written in full to *log_fh* (for post-mortem debugging)
+ - printed to the terminal as a single self-updating ``\\r`` line showing
+ the last non-empty line (real-time progress)
+
+ If *env* is given, it is merged on top of the parent environment
+ (os.environ) before the subprocess is launched.
+
+ Returns the process exit code. Does NOT raise on non-zero.
+ """
+ import os as _os
+ import sys
+ proc_env = _os.environ.copy()
+ if env:
+ proc_env.update(env)
+ proc = subprocess.Popen(
+ cmd, cwd=str(cwd), env=proc_env,
+ stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True,
+ )
+ assert proc.stdout is not None
+ last_line = ""
+ for line in proc.stdout:
+ log_fh.write(line)
+ stripped = line.rstrip()
+ if stripped:
+ last_line = stripped
+ print(f"\r {stripped[:140]}\033[K", end="", file=sys.stderr, flush=True)
+ proc.wait(timeout=timeout)
+ if last_line:
+ print(file=sys.stderr) # final newline after \r lines
+ return proc.returncode
diff --git a/src/TA_main2main_workflow/utils/logging.py b/src/TA_main2main_workflow/utils/logging.py
new file mode 100644
index 0000000..8c9a904
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/logging.py
@@ -0,0 +1,140 @@
+"""Logging setup for TA_main2main_workflow.
+
+Uses Python's standard ``logging`` module with a custom formatter that
+preserves the visual style of the old ``console.py`` (headers, sections,
+status icons) while routing everything through the logging framework.
+
+Usage::
+
+ from TA_main2main_workflow.utils.logging import get_logger
+ log = get_logger(__name__)
+ log.info("Starting sync...")
+ log.header("Phase 1: Detect") # boxed header
+ log.section("Build Triton-Ascend") # section divider
+ log.step(1, 3, "AI fix") # step indicator
+ log.status(True, "Build passed") # ✔ / ✘
+ log.key_value("target", "abc123") # key: value
+ log.table(rows) # summary table
+"""
+
+from __future__ import annotations
+
+import logging
+import sys
+from datetime import datetime
+from typing import Any
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Custom logger class
+# ═══════════════════════════════════════════════════════════════════════════
+
+
+class TALogger(logging.getLoggerClass()):
+ """Logger with extra formatting methods for workflow output."""
+
+ def header(self, title: str) -> None:
+ width = 72
+ self.info(f"\n╔{'═' * width}╗")
+ self.info(f"║ {title:^{width}} ║")
+ self.info(f"╚{'═' * width}╝")
+
+ def section(self, title: str) -> None:
+ ts = datetime.now().strftime("%H:%M:%S")
+ self.info(f"\n{'─' * 60}")
+ self.info(f" [{ts}] {title}")
+ self.info(f"{'─' * 60}")
+
+ def step(self, num: int, total: int, name: str) -> None:
+ ts = datetime.now().strftime("%H:%M:%S")
+ self.info(f"\n ▸ [{num}/{total}] {name} @ {ts}")
+
+ def status(self, ok: bool, msg: str) -> None:
+ icon = "✔" if ok else "✘"
+ self.info(f" {icon} {msg}")
+
+ def warn(self, msg: str, *args, **kwargs) -> None:
+ # Override to use consistent prefix
+ super().warning(f" ⚠ {msg}", *args, **kwargs)
+
+ def error(self, msg: str, *args, **kwargs) -> None:
+ super().error(f" ✘ {msg}", *args, **kwargs)
+
+ def key_value(self, key: str, value: Any) -> None:
+ self.info(f" {key}: {value}")
+
+ def flow_progress(self, phase: str, detail: str = "") -> None:
+ ts = datetime.now().strftime("%H:%M:%S")
+ msg = f"[{ts}] [{phase}] {detail}" if detail else f"[{ts}] [{phase}]"
+ self.info(msg)
+
+ def conflict_list(self, files: list[str]) -> None:
+ if not files:
+ self.info(" ℹ No conflicts")
+ return
+ self.info(f" Conflicted files ({len(files)}):")
+ for i, f in enumerate(files, 1):
+ self.info(f" {i}. {f}")
+
+ def ai_call(self, backend: str, mode: str, attempt: int, max_attempts: int) -> None:
+ ts = datetime.now().strftime("%H:%M:%S")
+ self.info(f"\n ╭─ AI Call ─────────────────────────────────────────────")
+ self.info(f" │ Backend: {backend}")
+ self.info(f" │ Mode: {mode}")
+ self.info(f" │ Attempt: {attempt}/{max_attempts}")
+ self.info(f" │ Time: {ts}")
+ self.info(f" ╰──────────────────────────────────────────────────────")
+
+ def ai_result(
+ self, ok: bool, modified_files: list[str] = (), summary: str = ""
+ ) -> None:
+ icon = "✔" if ok else "✘"
+ self.info(f"\n ╭─ AI Result ───────────────────────────────────────────")
+ self.info(f" │ Status: {icon} {'Success' if ok else 'Failed'}")
+ if modified_files:
+ self.info(f" │ Modified files ({len(modified_files)}):")
+ for f in modified_files:
+ self.info(f" │ • {f}")
+ if summary:
+ preview = summary[:500] + "..." if len(summary) > 500 else summary
+ self.info(f" │ Summary: {preview}")
+ self.info(f" ╰──────────────────────────────────────────────────────")
+
+ def table(self, rows: list[tuple[str, str, str]]) -> None:
+ ts = datetime.now().strftime("%H:%M:%S")
+ status_icons = {"PASS": "✔", "FAIL": "✘", "SKIP": "○", "WARN": "⚠"}
+ self.info(f"\n{'═' * 72}")
+ self.info(f" SYNC SUMMARY @ {ts}")
+ self.info(f"{'═' * 72}")
+ self.info(f" {'Phase':<30} {'Status':<8} {'Details'}")
+ self.info(f" {'─' * 30} {'─' * 8} {'─' * 32}")
+ for step, status, detail in rows:
+ icon = status_icons.get(status, "?")
+ self.info(f" {step:<30} {icon} {status:<5} {detail}")
+ self.info(f"{'═' * 72}")
+
+ def elapsed(self, seconds: float) -> None:
+ self.info(f"\n ⏱ Total elapsed: {seconds:.1f}s ({seconds / 60:.1f}m)")
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Setup
+# ═══════════════════════════════════════════════════════════════════════════
+
+logging.setLoggerClass(TALogger)
+
+
+def get_logger(name: str) -> TALogger:
+ """Return a configured TALogger for *name*."""
+ log = logging.getLogger(name)
+ if not log.handlers:
+ handler = logging.StreamHandler(sys.stdout)
+ handler.setFormatter(logging.Formatter("%(message)s"))
+ log.addHandler(handler)
+ log.setLevel(logging.INFO)
+ log.propagate = False
+ return log # type: ignore[return-value]
+
+
+# Default logger for simple imports
+default_logger = get_logger("ta-workflow")
diff --git a/src/TA_main2main_workflow/utils/submodule.py b/src/TA_main2main_workflow/utils/submodule.py
new file mode 100644
index 0000000..9aa0f83
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/submodule.py
@@ -0,0 +1,105 @@
+"""AscendNPU-IR submodule helpers.
+
+Detects changes in ``third_party/ascend/AscendNPU-IR`` and handles
+commit + push for the submodule before the parent repo is committed.
+"""
+
+from __future__ import annotations
+
+import os
+import subprocess
+from pathlib import Path
+
+from TA_main2main_workflow.utils.git import run_git, run_git_no_check
+from TA_main2main_workflow.utils.logging import get_logger
+
+log = get_logger(__name__)
+
+_SUBMODULE_DIR = "third_party/ascend/AscendNPU-IR"
+_NPUIR_REMOTE = "npuir-push"
+
+
+def _submodule_path(repo: Path) -> Path:
+ """Return the absolute path to the AscendNPU-IR submodule."""
+ return repo / _SUBMODULE_DIR
+
+
+def submodule_has_changes(repo: Path) -> bool:
+ """Return True if the AscendNPU-IR submodule has uncommitted changes."""
+ sp = _submodule_path(repo)
+ if not sp.exists():
+ return False
+ result = run_git_no_check(sp, "status", "--porcelain")
+ return bool(result.stdout.strip())
+
+
+def commit_submodule(repo: Path, commit_msg: str) -> bool:
+ """Stage all changes and commit in the AscendNPU-IR submodule.
+
+ Returns True if a commit was created, False if there was nothing to commit.
+ """
+ sp = _submodule_path(repo)
+ if not sp.exists():
+ log.info("AscendNPU-IR submodule not found — skipping")
+ return False
+
+ if not submodule_has_changes(repo):
+ return False
+
+ log.info(f"Committing AscendNPU-IR submodule changes...")
+ try:
+ run_git(sp, "add", "-A")
+ run_git(sp, "commit", "-s", "-m", commit_msg)
+ log.status(True, "AscendNPU-IR submodule committed")
+ return True
+ except Exception as e:
+ if "nothing to commit" not in str(getattr(e, "stderr", "")):
+ log.warning(f"Submodule commit failed: {e}")
+ return False
+
+
+def push_submodule(repo: Path, branch: str | None = None,
+ force: bool = True) -> bool:
+ """Push the AscendNPU-IR submodule to its dedicated remote.
+
+ Creates/updates a branch at current HEAD and pushes using
+ GH_TOKEN for authentication. By default uses ``--force-with-lease``.
+
+ Returns True on success.
+ """
+ sp = _submodule_path(repo)
+ if not sp.exists():
+ log.info("AscendNPU-IR submodule not found — skipping push")
+ return False
+
+ # Ensure the npuir-push remote exists
+ result = run_git_no_check(sp, "remote")
+ if _NPUIR_REMOTE not in result.stdout:
+ npuir_url = os.getenv("ASCENDNPU_IR_PUSH_URL", "")
+ if not npuir_url:
+ log.warning("ASCENDNPU_IR_PUSH_URL not set — cannot push submodule")
+ return False
+ # Embed GH_TOKEN in URL if available
+ token = os.getenv("GH_TOKEN", "")
+ if token and "@" not in npuir_url and npuir_url.startswith("https://"):
+ npuir_url = npuir_url.replace("https://", f"https://{token}@")
+ run_git(sp, "remote", "add", _NPUIR_REMOTE, npuir_url)
+
+ if branch is None:
+ branch = f"sync-{run_git(sp, 'rev-parse', '--short', 'HEAD').strip()}"
+
+ # Create/update branch at current HEAD
+ run_git(sp, "checkout", "-B", branch)
+
+ log.info(f"Pushing AscendNPU-IR branch '{branch}'...")
+ try:
+ push_args = ["push"]
+ if force:
+ push_args.append("--force-with-lease")
+ push_args.extend([_NPUIR_REMOTE, branch])
+ run_git(sp, *push_args)
+ log.status(True, f"AscendNPU-IR pushed to {branch}")
+ return True
+ except Exception as e:
+ log.error(f"Submodule push failed: {e}")
+ return False
diff --git a/src/TA_main2main_workflow/utils/tracker.py b/src/TA_main2main_workflow/utils/tracker.py
new file mode 100644
index 0000000..42a27a6
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/tracker.py
@@ -0,0 +1,40 @@
+"""Simple phase timer for the workflow pipeline.
+
+Usage::
+
+ from TA_main2main_workflow.utils.tracker import timed, total_elapsed
+
+ with timed("build"):
+ ... # build work
+
+ print(f"Total: {total_elapsed():.1f}s")
+"""
+
+from __future__ import annotations
+
+import time
+from contextlib import contextmanager
+
+_flow_start_time: float | None = None
+_phase_times: dict[str, float] = {}
+
+
+@contextmanager
+def timed(name: str):
+ """Context manager that records elapsed wall-clock time for *name*."""
+ global _flow_start_time
+ if _flow_start_time is None:
+ _flow_start_time = time.time()
+ start = time.time()
+ try:
+ yield
+ finally:
+ elapsed = time.time() - start
+ _phase_times[name] = elapsed
+
+
+def total_elapsed() -> float:
+ """Return total seconds since the first ``timed()`` call."""
+ if _flow_start_time is None:
+ return 0.0
+ return time.time() - _flow_start_time