diff --git a/docs/workflow.md b/docs/workflow.md
index 44f0316..3d211b9 100644
--- a/docs/workflow.md
+++ b/docs/workflow.md
@@ -1,23 +1,85 @@
+# 单步模式工作流
+
```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"]
+ B --> C["Phase 1: Detect"]
+ C -->|无新 commit| D["Done"]
+ C -->|有新 commit| E["Phase 2: Plan"]
+ E --> F["Phase 3: Per-Step Loop"]
+
+ subgraph STEP["每个步骤"]
+ F1["Merge"] --> F2{冲突?}
+ F2 -->|是| F3["AI Resolve"]
+ F2 -->|否| F4{LLVM 变更?}
+ F3 --> F4
+ F4 -->|是| F5["IR Patch Pipeline"]
+ F4 -->|否| F6["Build + Fix"]
+ F5 --> F7{通过?}
+ F6 --> F7
+ F7 -->|否| FAIL["UpgradeFailed"]
+ F7 -->|是| F8["Test + Fix"]
+ F8 -->|失败| FAIL
+ F8 -->|通过| F9["Commit"]
+ F9 --> F10["current_step += 1"]
+ end
+
+ STEP --> G["Phase 4: Finalize"]
+ G -->|push| PR["Create PR"]
+
+ style FAIL fill:#d73
+ style D fill:#4a9
+```
+
+## IR Patch Pipeline(LLVM hash 变更时)
+
+```mermaid
+flowchart TD
+ MERGE["Merge 后发现 LLVM hash 变更"] --> P1["1. 切换到目标 LLVM commit
确保工作区干净"]
+ P1 --> P2["2. 直接应用现有补丁
llvm_patch_f6ded0b.patch"]
+ P2 --> P3{应用成功?}
+ P3 -->|否| P4["AI 分析失败原因
根据当前 LLVM commit 调整补丁"]
+ P4 --> P2
+ P3 -->|是| P5["3. 编译 LLVM
(编译失败则 AI 修复补丁)"]
+ P5 --> P6{编译成功?}
+ P6 -->|否| P7["AI 根据编译错误修复补丁"]
+ P7 --> P5
+ P6 -->|是| P8["4. 编译 TA + AI 修复编译错误"]
+ P8 --> P9["5. Run pytest"]
+ P9 --> P10{测试通过?}
+ P10 -->|是| DONE["Done"]
+ P10 -->|否| P11["AI 分类: IR 问题 vs 代码问题"]
+ P11 -->|IR 问题| P12["AI 在现有补丁基础上
补充缺失 OP 的适配
(不重新生成)"]
+ P12 --> P13["重新编译 LLVM"]
+ P13 --> P8
+ P11 -->|代码问题| P14["AI 修复代码 -> Rebuild"]
+ P14 --> P9
+
+ style DONE fill:#4a9
+```
+
+## 关键变化(vs 旧流程)
+
+| 旧流程 | 新流程 |
+|--------|--------|
+| 先做 OP 分析 + 变更分析 | 直接应用现有补丁 |
+| AI 从零生成补丁 | AI 在现有补丁上调整/补充 |
+| 分析阶段在补丁之前 | 分析阶段在测试发现问题后 |
+| 每次重新生成完整补丁 | 保留已有内容,只补充缺失 OP |
+
+## 修复代码评价机制
+
+```
+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)
diff --git a/src/TA_main2main_workflow/agent/prompt.md b/src/TA_main2main_workflow/agent/prompt.md
index a143add..a118539 100644
--- a/src/TA_main2main_workflow/agent/prompt.md
+++ b/src/TA_main2main_workflow/agent/prompt.md
@@ -745,6 +745,46 @@ The active mode is: {mode}
Output ONLY by modifying `{ascend_patch_file}` in-place — the existing
patch file. Do NOT write analysis.md, step_summary.md, or review.md.
+── ir_fix_patch_apply mode ──────────────────────────────────────────────
+
+ Trigger: {mode} is "ir_fix_patch_apply" (fix patch that failed to apply).
+
+ The existing `{ascend_patch_file}` failed to apply to LLVM at
+ `{target_llvm_hash}`. Analyze the error and adjust the patch:
+
+ 1. Read the current patch file: `{ascend_patch_file}`
+ 2. Read the apply error: {patch_error_msg}
+ 3. For each failing hunk, use git show to see the target file:
+ git -C {llvm_project_path} show {target_llvm_hash}:
+ 4. Adjust the patch to match the target LLVM's actual code:
+ - Fix line numbers and context lines
+ - If the OP/func being patched was removed, remove that hunk
+ - If the API changed, adapt the patch code
+ 5. Write the fixed patch to `{ascend_patch_file}`
+
+ Reference: {reference_dir}/05-ir-patch-generation-guide.md
+
+── ir_supplement_patch mode ────────────────────────────────────────────
+
+ Trigger: {mode} is "ir_supplement_patch" (add missing OPs to existing patch).
+
+ Tests revealed IR compatibility issues that the current patch does NOT
+ cover. The current patch at `{ascend_patch_file}` already handles some
+ OPs. You must SUPPLEMENT it — keep all existing content and ADD
+ adaptations for the missing OPs identified in the diagnosis.
+
+ 1. Read the current patch to understand what's already covered
+ 2. Read `{step_dir}/ir_diagnosis.json` for the IR issues found
+ 3. For each missing OP, check its definition at the target LLVM:
+ git -C {llvm_project_path} show {target_llvm_hash}:
+ 4. ADD the new OP adaptations to the patch file (do NOT remove
+ existing content)
+ 5. Apply the same patch patterns (per the guide) to the new OPs
+ 6. Write the supplemented patch to `{ascend_patch_file}`
+
+ Reference: {reference_dir}/05-ir-patch-generation-guide.md
+ {reference_dir}/04-ir-compatibility-and-backend-adaptation.md
+
── For ir_diagnose mode ──
Output ONLY `{step_dir}/ir_diagnosis.json` — the structured JSON specified
in the ir_diagnose section above. Do NOT write analysis.md, step_summary.md,
diff --git a/src/TA_main2main_workflow/flow.py b/src/TA_main2main_workflow/flow.py
index 03af27d..b03dbf0 100644
--- a/src/TA_main2main_workflow/flow.py
+++ b/src/TA_main2main_workflow/flow.py
@@ -1,4446 +1,143 @@
-"""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::
-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 -> for each step:
+ merge -> [resolve] -> build (with fix loop) ->
+ [ir_patch (if LLVM changed)] -> test (with 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 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 __future__ import annotations
+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 (
- 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,
)
-
-_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
- 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()
- return UpgradeFailed
-
- # ── 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)
-
- # 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
+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
+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
+from TA_main2main_workflow.pipeline.test import test
+from TA_main2main_workflow.pipeline.ir_patch import per_step_ir_patch
+from TA_main2main_workflow.pipeline.commit import commit_step
+from TA_main2main_workflow.pipeline.finalize import finalize
+
+log = get_logger(__name__)
+
+
+class TA_Main2MainFlow:
+ """Orchestrator -- builds context, runs pipeline steps, handles PR."""
+
+ 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")
+ log.key_value("AI Backend", self.config.ai_backend)
+ log.key_value("Max Retries", str(self.config.max_retries))
+
+ # 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")
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]
- step_id = step["id"]
- self.state.retry_count = 0
-
- print_header(
- f"Single-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])
+ 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 3: Per-step loop
+ while ctx.current_step < ctx.total_steps:
+ step = ctx.steps[ctx.current_step]
+ sid = step["id"]
reason = step.get("reason", "line_budget")
- print_key_value("step reason", reason)
-
- self._print_workspace_info(f"Single-Step Mode — {step_id}")
-
- # 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
- return UpgradeFailed
-
- # ── Step C: IR patch if LLVM hash changed in this step ──
- # Covers LLVM rebuild + TA build + test+fix (with IR retry embedded)
- 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
- return UpgradeFailed
- else:
- # ── 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
+ log.header(f"Step {ctx.current_step + 1}/{ctx.total_steps}: {sid}")
+ log.key_value("commits", str(step["commit_count"]))
+ log.key_value("end commit", step["end_commit"][:12])
+ log.key_value("reason", reason)
+ ctx = ctx.copy_with(retry_count=0)
+
+ # 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 {sid}")
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
-
- # ── Step F: Commit step progress ──
- self._do_commit_step(step)
-
- # Record step description for 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', '?')}, "
- f"reason={reason}"
- )
- self.state.step_pr_descriptions.append(desc)
-
- # Record per-step detail for sync report
- 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": len(self.state.conflict_files),
- "build_fixes": self.state.build_fix_count,
- "test_fixes": self.state.test_fix_count,
- "retries": self.state.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})")
-
- # ── Phase 3: Finalize ──
- print_header("Finalize — Generate Summary & Push")
- self._do_finalize()
-
- # ── Phase 4: Push to GitHub + create PR ──
- self.push_to_github()
-
- self.state.summary_rows.append(
- ("Single-Step Sync", "PASS",
- f"{self.state.total_steps} step(s), branch: {self.state.work_branch}")
- )
- print_summary_table(self.state.summary_rows)
- print_elapsed_total()
-
- self.state.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}")
-
- # 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)")
+ log.status(True, "Conflicts resolved")
+
+ # Step C: Build + fix (with optional IR patch for LLVM steps)
+ with timed("build"):
+ if reason == "llvm_version":
+ log.section(f"LLVM Version Change in {sid} -- IR Patch Pipeline")
+ ctx = per_step_ir_patch(ctx, self.config, step)
+ if not ctx.build_passed or not ctx.test_passed:
+ log.error(f"IR patch pipeline failed for {sid}")
+ return UpgradeFailed
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"
- )
- 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)
- 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
+ ctx = build(ctx, self.config)
+ if not ctx.build_passed:
+ log.error(f"Build failed for {sid}")
+ return UpgradeFailed
+
+ # Step D: Test + fix (non-LLVM steps only; LLVM handled in ir_patch)
+ if reason != "llvm_version":
+ with timed("test"):
+ ctx = test(ctx, self.config)
+ if not ctx.test_passed:
+ log.error(f"Tests failed for {sid}")
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 E: Commit
+ ctx = commit_step(ctx, self.config)
- # ── Step D: commit step progress ──
- self._do_commit_step(step)
-
- # ── Record step description for final PR body ──
+ # Record step description
desc = (
- f"✅ **{step_id}**: {step['commit_count']} commits, "
+ f"**{sid}**: {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,
- })
+ f"source lines={step.get('source_changed_lines', '?')}, "
+ f"reason={reason}")
+ ctx = ctx.copy_with(
+ step_pr_descriptions=ctx.step_pr_descriptions + [desc],
+ current_step=ctx.current_step + 1)
+ log.status(True, f"Step {sid} completed ({ctx.current_step}/{ctx.total_steps})")
+
+ # Phase 4: Finalize
+ ctx = finalize(ctx, self.config)
+
+ # Phase Push: Push + PR
+ if self.config.push_to_github:
+ from TA_main2main_workflow.pipeline.push_pr import push_and_create_pr
+ try:
+ pr_url = push_and_create_pr(ctx, self.config)
+ log.status(True, f"PR created: {pr_url}")
+ ctx = ctx.copy_with(pr_url=pr_url)
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
- attempt = 0
-
- while attempt <= self.state.max_retries:
- 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", "")
- # ── Validate fix: only third_party/ascend/ files allowed ──
- fix_valid, fix_reason = self._validate_fix(modified_files, ascend_path)
- if not fix_valid:
- print_error(f"Fix rejected: {fix_reason}")
- print_warn(
- f"Fix modified files outside third_party/ascend/ — "
- f"changes reverted, this attempt will NOT count, "
- f"retrying fix with rejection feedback...")
- # Write rejection feedback so AI sees it next round
- rejection_file = step_dir / "fix_rejection.txt"
- rejection_file.write_text(
- f"PREVIOUS FIX WAS REJECTED: {fix_reason}\n"
- f"Only files under {ascend_path}/third_party/ascend/ "
- f"may be modified for compile-error fixes.\n",
- encoding="utf-8")
- self.state.fix_errors.append(str(rejection_file))
- if hasattr(self, '_last_ai_result'):
- self._last_ai_result["modified_files"] = []
- self._last_ai_result["step_summary"] = (
- f"REJECTED: {fix_reason}")
- continue # don't count this attempt, retry
- # 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
+ log.error(f"Failed to create PR: {e}")
- # 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}")
- attempt += 1
- 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 _detect_ascend_npu_ir_errors(self) -> bool:
- """Check whether the build log contains AscendNPU-IR compile errors.
-
- AscendNPU-IR (bishengir) is at third_party/ascend/AscendNPU-IR/.
- LLVM version changes often break its compilation — error patterns
- include bishengir paths, dialect registration failures, and MLIR
- API incompatibilities.
- """
- build_log = WORKSPACE_DIR / BUILD_LOG_FILE
- if not build_log.exists():
- return False
- try:
- content = build_log.read_text(encoding="utf-8", errors="replace")
- except Exception:
- return False
- # Patterns indicating AscendNPU-IR compilation failures
- npu_ir_markers = [
- "AscendNPU-IR",
- "bishengir",
- "bishengir-",
- "NPUIR",
- "HACC/IR",
- "HFusion/IR",
- "HIVM/IR",
- "third_party/ascend/",
- "AscendNPU",
- ]
- for marker in npu_ir_markers:
- if marker in content:
- return True
- return False
-
- def _detect_oom_in_tests(self) -> bool:
- """Check whether test failures include NPU/GPU OOM errors.
-
- OOM errors are transient resource exhaustion — they should trigger
- a full test-suite rerun with reduced concurrency instead of an AI
- code fix.
- """
- test_log_dir = WORKSPACE_DIR / "test-logs"
- oom_markers = [
- "out of memory",
- ]
- # Scan .log and .xml files (pytest JUnit XML captures test failure messages)
- if test_log_dir.exists():
- try:
- for log_file in test_log_dir.rglob("*"):
- if log_file.suffix not in (".log", ".xml"):
- continue
- content = log_file.read_text(encoding="utf-8", errors="replace")
- for marker in oom_markers:
- if marker.lower() in content.lower():
- return True
- except Exception:
- pass
- # Also check test result JSON
- test_result = WORKSPACE_DIR / TEST_RESULT_FILE
- if test_result.exists():
- try:
- data = json.loads(test_result.read_text(encoding="utf-8"))
- error_msg = json.dumps(data) # search the whole JSON
- for marker in oom_markers:
- if marker.lower() in error_msg.lower():
- return True
- except Exception:
- pass
- return False
-
- def _rerun_tests_reduced_concurrency(self, ascend_path: Path, max_reruns: int = 5) -> bool | None:
- """Rerun tests with halved concurrency on OOM, restoring it after.
-
- Returns True if tests pass, None if SKIP_E2E_TEST, False if still failing.
- """
- original_procs = self.state.num_procs
- reduced = max(1, original_procs // 2)
- self.state.num_procs = reduced
- print_warn(
- f"Reducing pytest concurrency: {original_procs} → {reduced} "
- f"(to avoid OOM)")
- try:
- for rerun in range(1, max_reruns + 1):
- print_info(f"OOM rerun {rerun}/{max_reruns} (procs={reduced})")
- result = self._do_test(ascend_path)
- if result is None or result:
- return result
- if not self._detect_oom_in_tests():
- print_info("OOM resolved — remaining failures are not memory-related")
- return False
- return False
- finally:
- self.state.num_procs = original_procs
- print_info(f"Restored pytest concurrency to {original_procs}")
-
- def _validate_fix(self, modified_files: list[str], ascend_path: Path) -> tuple[bool, str]:
- """Validate that an AI fix only touches allowed files.
-
- Checks:
- 1. All modified files are under third_party/ascend/ (hard rule)
-
- When validation fails, the illegal changes are reverted via
- git checkout so the next fix attempt starts from a clean state.
-
- Returns (passed: bool, reason: str).
- """
- if not modified_files:
- return False, "No files were modified"
-
- illegal_files: list[str] = []
- ascend_root = str(ascend_path / "third_party" / "ascend")
- for f in modified_files:
- f_abs = str(Path(f).resolve()) if not Path(f).is_absolute() else f
- if ascend_root not in f_abs:
- illegal_files.append(f)
-
- if illegal_files:
- # ── Revert ALL working-tree changes since the fix is invalid ──
- print_warn(f"Reverting invalid fix changes in {ascend_path}...")
- try:
- subprocess.run(
- ["git", "checkout", "--", "."],
- cwd=str(ascend_path),
- capture_output=True, text=True, timeout=30,
- )
- subprocess.run(
- ["git", "clean", "-fd"],
- cwd=str(ascend_path),
- capture_output=True, text=True, timeout=30,
- )
- print_status(True, "Reverted — working tree is clean")
- except Exception as e:
- print_error(f"Failed to revert changes: {e}")
- return False, (
- f"Fix modified files OUTSIDE third_party/ascend/: "
- + ", ".join(illegal_files)
- + ". Changes have been reverted. "
- + "Next fix MUST only modify files under "
- + f"{ascend_path}/third_party/ascend/")
-
- print_status(True,
- f"Fix validation: {len(modified_files)} file(s) all within "
- f"third_party/ascend/")
- return True, "All modified files are within third_party/ascend/"
-
- def _do_ai_fix(self, ascend_path: Path, step_dir: Path, attempt: int,
- ascend_npu_ir_fix: bool = False) -> 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,
- "ascend_npu_ir_fix": str(ascend_npu_ir_fix).lower(),
- "ascend_npu_ir_compat_ref": str(
- Path(__file__).parent / "reference"
- / "AscendNPU-IR_LLVM_VERSION_COMPAT.md"),
- })
-
- 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():
- print_error(
- "LLVM patch apply/rebuild failed after all retries — "
- "cannot proceed without a working LLVM build. "
- "Terminating IR patch loop.")
- self.state.summary_rows.append(
- ("IR Patch Loop", "FATAL", "LLVM rebuild exhausted"))
- return False
-
- # ── [4.1] Build TA ──
- print_info("Step 4.1: Building Triton-Ascend with patched LLVM...")
- build_ok = self._do_build(ascend_path, clean=True)
- if not build_ok:
- if os.getenv("SKIP_AI_ANALYSIS", "false").lower() == "true":
- return False
- # ── AscendNPU-IR compile-error fix loop ──
- # LLVM version changes often break AscendNPU-IR compilation.
- # Loop: detect errors → AI fix with NPU-IR reference docs →
- # rebuild until the build passes or retries exhausted.
- for fix_attempt in range(1, self.state.max_retries + 1):
- self.state.fix_errors = [str(WORKSPACE_DIR / BUILD_RESULT_FILE)]
- self.state.build_fix_count += 1
- # Check if errors are AscendNPU-IR related
- is_npu_ir = self._detect_ascend_npu_ir_errors()
- if is_npu_ir:
- print_warn(
- f"AscendNPU-IR compile errors detected — "
- f"AI will reference AscendNPU-IR_LLVM_VERSION_COMPAT.md "
- f"(attempt {fix_attempt}/{self.state.max_retries})")
- else:
- print_warn(
- f"Build failed after IR patches — AI fix "
- f"(attempt {fix_attempt}/{self.state.max_retries})")
- self._do_ai_fix(ascend_path, WORKSPACE_DIR, fix_attempt,
- ascend_npu_ir_fix=is_npu_ir)
- if self._do_build(ascend_path, clean=False):
- build_ok = True
- break
- print_warn(f"Build still failing after fix attempt {fix_attempt}")
- if not build_ok:
- print_error(
- f"Build still failing after {self.state.max_retries} "
- f"fix attempts in IR patch iteration {iteration + 1}")
- self.state.ir_loop_details.append({
- "iteration": iteration + 1,
- "result": "BUILD_FIX_EXHAUSTED",
- })
- return False
-
- # ── [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] = []
- ascend_root = ascend_path / "third_party" / "ascend"
- scan_dirs = [
- ascend_root,
- 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", "--exclude-dir=patch",
- "--exclude-dir=cmake", 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...")
-
- # Read target LLVM hash (same as _do_ir_change_analysis uses)
- 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("Baseline LLVM", f"{_ASCEND_BASELINE_LLVM_HASH[:12]}")
- print_key_value("Target LLVM", f"{target_llvm_hash[:12]}")
-
- 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,
- "target_llvm_hash": target_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.
-
- Retry loop (max 10): if patch apply fails or LLVM build fails,
- AI fixes the patch and we retry from scratch (clean → checkout →
- apply → build).
- """
- 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
-
- # 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])
-
- from TA_main2main_workflow.scripts.build_test import (
- apply_llvm_patches, build_llvm)
-
- _MAX_PATCH_RETRIES = 10
-
- for retry in range(_MAX_PATCH_RETRIES + 1):
- is_retry = retry > 0
- if is_retry:
- print_header(
- f"Patch Apply/Rebuild Retry {retry}/{_MAX_PATCH_RETRIES}")
-
- # ── Ensure llvm-project workspace is clean ──
- if not self._ensure_llvm_workspace_clean(reason="ir-apply-patches"):
- print_error("Cannot clean llvm-project workspace")
- self.state.summary_rows.append(
- ("IR Apply+Rebuild", "FAIL", "workspace not clean"))
- return False
-
- # ── [3.4] Apply patch ──
- 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"]
- error_msg = failed[0]['error'][:500] if failed else "unknown"
- print_error(f"LLVM patch apply failed: {error_msg}")
- if retry < _MAX_PATCH_RETRIES:
- print_warn(
- f"Patch apply failed — AI will fix the patch "
- f"(retry {retry + 1}/{_MAX_PATCH_RETRIES})")
- self._do_ir_fix_patch(
- ascend_path, ascend_patch, target_llvm_hash,
- error_type="apply", error_msg=error_msg,
- retry=retry + 1)
- continue
- self.state.summary_rows.append(
- ("IR Apply+Rebuild", "FAIL",
- f"patch apply failed after {_MAX_PATCH_RETRIES} retries"))
- return False
-
- print_status(True, f"{ascend_patch.name} applied to llvm-project")
-
- # ── Show git status after patch ──
- 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 ──
- try:
- print_info("Step 3.5: Rebuilding LLVM (this takes ~15-30 minutes)...")
- llvm_install = Path(os.path.expanduser(
- os.getenv("LLVM_INSTALL_PREFIX_SYNC", "~/llvm-install-sync")))
- llvm_prefix = build_llvm(
- llvm_project, llvm_install,
- 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"
- + (f" (after {retry} retries)" if is_retry else "")))
- return True
- except Exception as e:
- build_error = str(e)[:500]
- # Also capture tail of build log for AI context
- build_log = WORKSPACE_DIR / "llvm_build.log"
- if build_log.exists():
- try:
- log_tail = build_log.read_text(
- encoding="utf-8", errors="replace")[-3000:]
- build_error = (
- f"Build exception: {e}\n\n"
- f"Build log tail:\n{log_tail}")
- except Exception:
- pass
- print_error(f"LLVM rebuild failed: {e}")
- if retry < _MAX_PATCH_RETRIES:
- print_warn(
- f"LLVM build failed — AI will fix the patch "
- f"(retry {retry + 1}/{_MAX_PATCH_RETRIES})")
- self._do_ir_fix_patch(
- ascend_path, ascend_patch, target_llvm_hash,
- error_type="build", error_msg=build_error,
- retry=retry + 1)
- continue
- self.state.summary_rows.append(
- ("IR Apply+Rebuild", "FAIL",
- f"LLVM build failed after {_MAX_PATCH_RETRIES} retries"))
- return False
-
- return False
-
- def _do_ir_fix_patch(self, ascend_path: Path, ascend_patch: Path,
- target_llvm_hash: str, error_type: str,
- error_msg: str, retry: int) -> None:
- """Invoke AI to fix a broken IR compatibility patch.
-
- Called when patch apply or LLVM build fails. AI re-examines the
- target LLVM commit and IR compatibility references, then fixes
- the patch in-place.
- """
- print_info(f"Invoking AI to fix patch ({error_type} failure, retry {retry})...")
- try:
- ai_result = run_opencode_adapter({
- "step_id": f"ir-fix-patch-{retry}",
- "previous_step_id": "ir-generate-patch",
- "previous_step_summary_path": "",
- "is_last_step": "false",
- "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([], 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_llvm_hash,
- "ascend_patch_file": str(ascend_patch),
- "patch_error_type": error_type,
- "patch_error_msg": error_msg,
- })
- _ = ai_result
- except Exception as e:
- print_error(f"AI patch fix failed: {e}")
-
- 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 LLVM update pipeline: compile-error fix → IR patch → test.
-
- Called from _run_single_step_mode() when a step's merge included an
- LLVM hash change.
-
- Pipeline:
- 1. Build new LLVM (clean, no patches) + build TA + fix compile errors
- — resolve all LLVM version-related build issues first.
- 2. IR OP analysis → IR change analysis → generate patches
- 3. Apply patches + rebuild LLVM + build TA
- 4. Test + AI fix loop:
- - Code issues → AI fix → rebuild TA → retest
- - IR issues → regenerate patches → rebuild LLVM → build TA → retest
- """
- step_id = step["id"]
- ascend_path = Path(self.state.triton_ascend_path)
-
- # ── 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
-
- # Read target LLVM hash
- 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()
-
- # ── 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))
-
- from TA_main2main_workflow.scripts.build_test import build_llvm
- llvm_install = Path(os.path.expanduser(
- os.getenv("LLVM_INSTALL_PREFIX_SYNC", "~/llvm-install-sync")))
-
- # ═══════════════════════════════════════════════════════════════
- # Phase 1: Build new LLVM (clean, no patches) + fix TA compile errors
- # ═══════════════════════════════════════════════════════════════
- print_header(f"Phase 1: Build new LLVM + Fix TA Compile Errors — {step_id}")
- print_key_value("Baseline LLVM", _ASCEND_BASELINE_LLVM_HASH[:12])
- print_key_value("Target LLVM", target_llvm_hash[:12])
-
- # 1a. Clean llvm-project and checkout target commit
- if not self._ensure_llvm_workspace_clean(reason="pre-ir-build"):
- print_error("Cannot clean llvm-project workspace")
- return False
- try:
- subprocess.run(
- ["git", "checkout", target_llvm_hash],
- cwd=str(_llvm_project_path()),
- capture_output=True, text=True, timeout=120,
- )
- except Exception as e:
- print_error(f"Failed to checkout target LLVM: {e}")
- return False
-
- # 1b. Build LLVM (no IR patches)
- print_info("Building LLVM at target commit (no IR patches)...")
- try:
- llvm_prefix = build_llvm(
- _llvm_project_path(), llvm_install,
- required_hash=target_llvm_hash,
- )
- if llvm_prefix and not self.state.llvm_prefix:
- self.state.llvm_prefix = llvm_prefix
- print_status(True, "Baseline LLVM build complete (no IR patches)")
- except Exception as e:
- print_error(f"Baseline LLVM build failed: {e}")
- return False
-
- # 1c. Build TA and fix compile errors (no IR patches yet)
- print_info("Building Triton-Ascend with new LLVM (no IR patches)...")
- build_ok = self._do_build(ascend_path, clean=True)
- if not build_ok:
- if os.getenv("SKIP_AI_ANALYSIS", "false").lower() == "true":
- return False
- print_warn("Build failed with new LLVM — entering compile-error fix loop")
- for fix_attempt in range(1, self.state.max_retries + 1):
- self.state.fix_errors = [str(WORKSPACE_DIR / BUILD_RESULT_FILE)]
- self.state.build_fix_count += 1
- is_npu_ir = self._detect_ascend_npu_ir_errors()
- print_warn(
- f"Compile error fix attempt {fix_attempt}/{self.state.max_retries}"
- f"{' (AscendNPU-IR)' if is_npu_ir else ''}")
- self._do_ai_fix(ascend_path, WORKSPACE_DIR, fix_attempt,
- ascend_npu_ir_fix=is_npu_ir)
- if self._do_build(ascend_path, clean=False):
- build_ok = True
- break
- if not build_ok:
- print_error(
- f"TA build still failing after {self.state.max_retries} fixes "
- f"— cannot proceed to IR patch generation")
- self.state.summary_rows.append(
- ("Phase 1", "FATAL", "compile errors not resolved"))
- return False
-
- print_status(True, "TA builds successfully with new LLVM — compile errors resolved")
- self.state.summary_rows.append(
- ("Phase 1", "PASS", "TA builds (no IR patches)"))
-
- # ═══════════════════════════════════════════════════════════════
- # Phase 2: IR patch generation → apply → rebuild → test + fix loop
- # ═══════════════════════════════════════════════════════════════
- print_header(f"Phase 2: IR Patch Generation & Test — {step_id}")
- self._print_workspace_info("Phase 2: 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 — {step_id} "
- f"(iter {iteration + 1}/{self.state.ir_max_iterations})"
- )
-
- # [2.1 + 2.2] OP analysis (first iteration only)
- if iteration == 0:
- print_info("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
-
- # [2.3] Generate patches
- if not self._do_ir_generate_patches():
- return False
-
- # [2.4 + 2.5] Apply patches + rebuild LLVM (retry on patch failure)
- 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
- print_warn(
- f"LLVM rebuild failed (patch attempt {patch_attempt + 1}) — "
- f"retrying patch generation")
- self._stash_and_drop_llvm_patch()
- if not self._do_ir_generate_patches():
- break
-
- if not rebuild_ok:
- print_warn(f"LLVM rebuild failed in iteration {iteration + 1}")
- continue
-
- print_status(True, f"IR patch + LLVM rebuild OK for {step_id}")
-
- # Build TA with patched LLVM
- print_info("Building Triton-Ascend with patched LLVM...")
- build_ok = self._do_build(ascend_path, clean=(iteration == 0))
- if not build_ok:
- for fix_attempt in range(1, self.state.max_retries + 1):
- self.state.fix_errors = [str(WORKSPACE_DIR / BUILD_RESULT_FILE)]
- self.state.build_fix_count += 1
- is_npu_ir = self._detect_ascend_npu_ir_errors()
- print_warn(
- f"Build failed after IR patch — AI fix "
- f"{fix_attempt}/{self.state.max_retries}"
- f"{' (AscendNPU-IR)' if is_npu_ir else ''}")
- self._do_ai_fix(ascend_path, WORKSPACE_DIR, fix_attempt,
- ascend_npu_ir_fix=is_npu_ir)
- if self._do_build(ascend_path, clean=False):
- build_ok = True
- break
- if not build_ok:
- print_error(f"Build still failing after {self.state.max_retries} fixes")
- continue
-
- # Test + fix loop (with IR retry embedded)
- test_ok = self._do_test_and_fix_with_ir_retry(
- step, ascend_path, iteration)
- if test_ok:
- self.state.ir_loop_details.append({
- "step_id": step_id,
- "iteration": iteration + 1,
- "result": "ALL_PASS",
- })
- return True
-
- print_warn(f"IR patch iteration {iteration + 1} — "
- f"IR issues remain, retrying outer loop")
-
- print_error(f"IR patch loop exhausted {self.state.ir_max_iterations} "
- f"iterations for {step_id}")
- return False
-
- def _do_test_and_fix_with_ir_retry(
- self, step: dict, ascend_path: Path, ir_iteration: int) -> bool:
- """Test + AI fix loop with embedded IR patch retry.
-
- Runs tests, classifies failures (IR vs code), and fixes them:
- - OOM errors → automatic full-suite rerun (up to 5), no AI fix
- - Code issues → AI fix → rebuild → retest (up to max_retries)
- - IR issues → regenerate patches → rebuild LLVM → build TA → retest
- (up to MAX_IR_RETRIES within this test loop)
-
- Returns True when all tests pass, False if IR retries exhausted.
- """
- _MAX_IR_RETRIES = 3
- _MAX_OOM_RERUNS = 5
- step_id = step["id"]
- step_dir = WORKSPACE_DIR / STEPS_DIR / step_id
- step_dir.mkdir(parents=True, exist_ok=True)
-
- ir_retries = 0
- code_fix_attempt = 0
-
- while ir_retries <= _MAX_IR_RETRIES and code_fix_attempt <= self.state.max_retries:
- # ── Run tests ──
- test_result = self._do_test(ascend_path)
- if test_result is None:
- # SKIP_E2E_TEST — treat as pass
- return True
- if test_result:
- print_status(True, f"All tests pass for {step_id}")
- return True
-
- # ── OOM detection: rerun with reduced concurrency, skip AI ──
- if self._detect_oom_in_tests():
- print_warn("NPU/CUDA OOM detected — rerunning with reduced concurrency")
- oom_result = self._rerun_tests_reduced_concurrency(
- ascend_path, max_reruns=_MAX_OOM_RERUNS)
- if oom_result is None or oom_result:
- return True if oom_result else None # SKIP or pass
- if not self._detect_oom_in_tests():
- print_info("OOM resolved — classifying remaining failures")
- else:
- print_error(f"OOM persists after {_MAX_OOM_RERUNS} reruns")
- return False
-
- # ── Classify failures: IR vs code ──
- print_warn(f"Tests failed — classifying failures (IR vs code)...")
- has_ir_issues = self._do_ir_diagnose_failures()
- if has_ir_issues:
- ir_retries += 1
- print_warn(
- f"IR compatibility issues detected "
- f"(IR retry {ir_retries}/{_MAX_IR_RETRIES}) — "
- f"regenerating IR patches...")
- # Regenerate patches + apply + rebuild LLVM
- if not self._do_ir_generate_patches():
- print_error("IR patch regeneration failed")
- return False
- # Clean llvm workspace, apply patches, rebuild
- for patch_attempt in range(IR_MAX_ITERATIONS):
- if self._do_ir_apply_patches_and_rebuild():
- break
- self._stash_and_drop_llvm_patch()
- if not self._do_ir_generate_patches():
- break
- else:
- print_error("LLVM rebuild failed after IR retry")
- continue
- # Rebuild TA
- if not self._do_build(ascend_path, clean=False):
- print_warn("TA build failed after IR retry — "
- "will fix in next iteration")
- continue
-
- # ── Code issues → AI fix ──
- code_fix_attempt += 1
- print_warn(
- f"Code issues detected — AI fix attempt "
- f"{code_fix_attempt}/{self.state.max_retries}")
- self.state.fix_errors = self._collect_test_error_logs()
- if self.state.fix_errors:
- self._do_ai_fix(ascend_path, step_dir, code_fix_attempt)
- self.state.test_fix_count += 1
- if not self._do_build(ascend_path, clean=False):
- print_warn("Build failed after code fix")
- else:
- print_warn("No test error logs found — cannot fix")
- break
-
- if ir_retries > _MAX_IR_RETRIES:
- print_error(f"IR retries exhausted ({_MAX_IR_RETRIES}) — IR issues unresolved")
- else:
- print_error(f"Code fix attempts exhausted ({self.state.max_retries})")
- 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.
-
- OOM errors trigger automatic full-suite reruns (up to 5) without
- consuming AI fix attempts. Fix validation rejections also don't
- consume attempts — changes are reverted and AI retries.
-
- 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)
-
- _MAX_OOM_RERUNS = 5
- test_passed = False
- attempt = 0
-
- while attempt <= self.state.max_retries:
- is_fix_attempt = attempt > 0
- self.state.retry_count = attempt
-
- # AI fix test failures (skip on first round)
- if is_fix_attempt:
- # ── OOM detection: rerun with reduced concurrency, skip AI ──
- if self._detect_oom_in_tests():
- print_warn("NPU OOM detected — rerunning with reduced concurrency")
- oom_result = self._rerun_tests_reduced_concurrency(
- ascend_path, max_reruns=_MAX_OOM_RERUNS)
- if oom_result is None:
- test_passed = True
- break
- if oom_result:
- test_passed = True
- break
- if not self._detect_oom_in_tests():
- print_info("OOM resolved — remaining failures need AI fix")
- else:
- print_error(
- f"OOM persists after {_MAX_OOM_RERUNS} reruns — "
- f"resource issue, cannot continue")
- self.state.summary_rows.append(
- ("Tests", "FAIL", f"OOM after {_MAX_OOM_RERUNS} reruns"))
- return False
-
- 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)
- # 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", "")
- # ── Validate fix: only third_party/ascend/ files allowed ──
- fix_valid, fix_reason = self._validate_fix(modified_files, ascend_path)
- if not fix_valid:
- print_error(f"Fix rejected: {fix_reason}")
- print_warn(
- f"Fix modified files outside third_party/ascend/ — "
- f"changes reverted, this attempt will NOT count, "
- f"retrying fix with rejection feedback...")
- # Write rejection feedback so AI sees it next round
- rejection_file = step_dir / "fix_rejection.txt"
- rejection_file.write_text(
- f"PREVIOUS FIX WAS REJECTED: {fix_reason}\n"
- f"For test fixes, prefer files under "
- f"{ascend_path}/third_party/ascend/. "
- f"Upstream files may only be modified when root "
- f"cause analysis confirms no Ascend-side workaround.\n",
- encoding="utf-8")
- self.state.fix_errors.append(str(rejection_file))
- if hasattr(self, '_last_ai_result'):
- self._last_ai_result["modified_files"] = []
- self._last_ai_result["step_summary"] = (
- f"REJECTED: {fix_reason}")
- continue # don't count this attempt, retry
- 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})")
- attempt += 1
- 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})")
- attempt += 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
+ return UpgradeCompleted
diff --git a/src/TA_main2main_workflow/main.py b/src/TA_main2main_workflow/main.py
index d735cf2..5c977cb 100644
--- a/src/TA_main2main_workflow/main.py
+++ b/src/TA_main2main_workflow/main.py
@@ -1,175 +1,100 @@
#!/usr/bin/env python3
-"""CLI entrypoint for TA_main2main_workflow — Triton-Ascend upstream sync.
+"""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)
-
-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_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
- 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)
+ ta-kickoff Run the main2main sync flow
+
+Configuration is via environment variables (see TAConfig.from_env() for
+the full list). CLI args override env vars for common options.
+
+Key environment variables:
+ TRITON_ASCEND_PATH -- path to triton-ascend repo
+ TRITON_PATH -- path to upstream triton repo
+ TRITON_TARGET_COMMIT -- specific upstream commit to sync to
+ TA_MAX_RETRIES -- max AI fix retries (default: 10)
+ SKIP_AI_ANALYSIS -- set to "true" to skip AI
+ SKIP_BUILD -- set to "true" to skip build
+ SKIP_E2E_TEST -- set to "true" to skip tests
+ PUSH_TO_GITHUB -- set to "true" to auto-create PR
+ NUM_PROCS -- pytest parallel workers (default: 16)
+ LLVM_PROJECT_PATH -- path to llvm-project repo
+ LLVM_INSTALL_PREFIX_SYNC -- LLVM install prefix
+ TA_BASE_BRANCH -- base branch for sync (default: upstream_sync)
"""
import argparse
-import os
import sys
-from pathlib import Path
from TA_main2main_workflow.flow import TA_Main2MainFlow
+from TA_main2main_workflow.utils.config import TAConfig
from TA_main2main_workflow.utils import UpgradeFailed
-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")
-
+def _print_banner(config: TAConfig) -> None:
+ ai = "NO (SKIP_AI_ANALYSIS=true)" if config.skip_ai_analysis else "YES"
print(f"╔{'═' * 60}╗")
- print(f"║ TA_main2main_workflow — Triton-Ascend Upstream Sync ║")
+ 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"║ AI Backend: {config.ai_backend:<44}║")
+ print(f"║ AI Enabled: {ai:<44}║")
+ print(f"║ Skip Build: {str(config.skip_build):<44}║")
+ print(f"║ Skip Test: {str(config.skip_e2e_test):<44}║")
+ print(f"║ Max Retries: {config.max_retries:<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.")
+ if config.skip_ai_analysis:
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
+ print(" ⚠ WARNING: SKIP_AI_ANALYSIS=true -- AI will NOT be called!")
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 Flow")
parser.add_argument(
"--triton-ascend-path", default=None,
- help="Local path to the triton-ascend repository (default: current directory)"
- )
- parser.add_argument(
- "--triton-path", default=None,
- help="Local path to the upstream triton repository (default: uses remote)"
- )
+ help="Local path to the triton-ascend repository")
parser.add_argument(
"--target-commit", default=None,
- help="Upstream triton commit SHA to merge (default: upstream HEAD)"
- )
+ help="Upstream triton commit SHA to merge (default: upstream HEAD)")
parser.add_argument(
"--llvm-prefix", default=None,
- help="LLVM install prefix path for building"
- )
+ help="LLVM install prefix path for building (LLVM_INSTALL_PREFIX_SYNC)")
parser.add_argument(
- "--conda-env", default=None,
- help="Conda environment name (default: ta-upgrade)"
- )
+ "--max-retries", type=int, default=None,
+ help="Max AI fix retries (default: 10)")
parser.add_argument(
"--num-procs", type=int, default=None,
- help="Number of parallel pytest workers (default: 16)"
- )
+ help="Number of parallel pytest workers (default: 16)")
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 = {}
+ # Build config from env, then overlay CLI args
+ config = TAConfig.from_env()
if args.triton_ascend_path:
- inputs["triton_ascend_path"] = args.triton_ascend_path
- if args.triton_path:
- inputs["triton_path"] = args.triton_path
+ config.triton_ascend_path = args.triton_ascend_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
- if args.conda_env:
- inputs["conda_env"] = args.conda_env
- if args.num_procs:
- inputs["num_procs"] = args.num_procs
+ config.llvm_install_prefix_sync = args.llvm_prefix
+ if args.max_retries is not None:
+ config.max_retries = args.max_retries
+ if args.num_procs is not None:
+ config.test_procs = args.num_procs
- flow = TA_Main2MainFlow()
+ _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}")
+ import traceback
+ traceback.print_exc()
sys.exit(1)
- if _is_failed(result):
+ if result == UpgradeFailed:
print(f"\n{'=' * 60}")
- print(f" WORKFLOW FAILED — exiting with code 1")
+ print(f" WORKFLOW FAILED -- exiting with code 1")
print(f"{'=' * 60}")
sys.exit(1)
@@ -178,17 +103,5 @@ def kickoff():
print(f"{'=' * 60}")
-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}")
-
-
if __name__ == "__main__":
kickoff()
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..0006568
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/build.py
@@ -0,0 +1,132 @@
+"""Pipeline step: Build LLVM + Triton-Ascend with retry/fix loop."""
+
+from __future__ import annotations
+import os, 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, run_git_no_check
+from TA_main2main_workflow.utils import (
+ BUILD_RESULT_FILE, BUILD_LOG_FILE, WORKSPACE_DIR, STEPS_DIR,
+)
+from TA_main2main_workflow.pipeline.fix import ai_fix, validate_fix, get_last_ai_result
+
+log = get_logger(__name__)
+
+
+def build(ctx: WorkflowContext, config: TAConfig,
+ do_ir_patch: bool = False) -> WorkflowContext:
+ """Build phase: LLVM setup -> TA build -> fix loop.
+
+ When do_ir_patch is True, IR patch generation + LLVM rebuild
+ happens inside the loop (per-step IR patch mode).
+ """
+ if config.skip_build:
+ log.info("SKIP_BUILD=true -- skipping build")
+ return ctx.copy_with(build_passed=True)
+
+ ascend_path = Path(ctx.triton_ascend_path)
+ step = ctx.steps[ctx.current_step] if 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)
+
+ build_passed = False
+ attempt = 0
+
+ while attempt <= config.max_retries:
+ is_fix_attempt = attempt > 0
+ ctx = ctx.copy_with(retry_count=attempt)
+
+ if is_fix_attempt:
+ log.header(f"Build Fix Attempt {attempt}/{config.max_retries}")
+ is_npu_ir = _detect_ascend_npu_ir_errors()
+ ctx = ai_fix(ctx, config, attempt=attempt, ascend_npu_ir_fix=is_npu_ir)
+
+ # Validate fix
+ modified_files = get_last_ai_result().modified_files
+ fix_valid, fix_reason = validate_fix(modified_files, ascend_path)
+ if not fix_valid:
+ log.error(f"Fix rejected: {fix_reason}")
+ rejection_file = step_dir / "fix_rejection.txt"
+ rejection_file.write_text(
+ f"PREVIOUS FIX WAS REJECTED: {fix_reason}\n"
+ f"Only files under {ascend_path}/third_party/ascend/ "
+ f"may be modified for compile-error fixes.\n",
+ encoding="utf-8")
+ ctx = ctx.copy_with(
+ fix_errors=ctx.fix_errors + [str(rejection_file)])
+ continue # don't count this attempt
+
+ # Build triton-ascend
+ with timed("build-triton"):
+ ctx = _build_triton(ctx, config, clean=(attempt == 0))
+ if ctx.build_passed:
+ build_passed = True
+ break
+
+ log.info(f"Build failed (attempt {attempt + 1}) -- retrying")
+ ctx = ctx.copy_with(
+ fix_errors=[str(WORKSPACE_DIR / BUILD_RESULT_FILE)])
+ attempt += 1
+
+ return ctx.copy_with(build_passed=build_passed)
+
+
+def _build_triton(ctx: WorkflowContext, config: TAConfig,
+ clean: bool = False, python_exe: str = "python3") -> WorkflowContext:
+ """Build triton-ascend."""
+ from TA_main2main_workflow.scripts.build_test import build_triton_ascend
+
+ ascend_path = Path(ctx.triton_ascend_path)
+ llvm_prefix = config.llvm_install_prefix_sync or ctx.llvm_prefix
+ if not llvm_prefix:
+ llvm_prefix = os.path.expanduser(
+ os.getenv("LLVM_INSTALL_PREFIX_SYNC", "~/llvm-install-sync"))
+
+ python_exe = python_exe or os.getenv("PYTHON", "python3.10")
+
+ log.section("Build Triton-Ascend")
+ try:
+ build_result = build_triton_ascend(
+ ascend_path,
+ llvm_prefix=str(llvm_prefix),
+ conda_env=config.conda_env,
+ clean_build=clean,
+ python_exe=python_exe,
+ )
+ passed = build_result.get("all_passed", False)
+ if passed:
+ log.status(True, "Build passed")
+ return ctx.copy_with(build_passed=True)
+ log.error("Build FAILED")
+ return ctx.copy_with(
+ build_passed=False,
+ fix_errors=[str(WORKSPACE_DIR / BUILD_RESULT_FILE)])
+ except Exception as e:
+ log.error(f"Build FAILED: {e}")
+ return ctx.copy_with(
+ build_passed=False,
+ fix_errors=[str(WORKSPACE_DIR / BUILD_RESULT_FILE)])
+
+
+def _detect_ascend_npu_ir_errors() -> bool:
+ """Check whether the build log contains AscendNPU-IR compile errors."""
+ build_log = WORKSPACE_DIR / BUILD_LOG_FILE
+ if not build_log.exists():
+ return False
+ try:
+ content = build_log.read_text(encoding="utf-8", errors="replace")
+ except Exception:
+ return False
+ markers = [
+ "AscendNPU-IR", "bishengir", "bishengir-", "NPUIR",
+ "HACC/IR", "HFusion/IR", "HIVM/IR",
+ "third_party/ascend/", "AscendNPU",
+ ]
+ for marker in markers:
+ if marker in content:
+ return True
+ return False
diff --git a/src/TA_main2main_workflow/pipeline/commit.py b/src/TA_main2main_workflow/pipeline/commit.py
new file mode 100644
index 0000000..74db4b9
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/commit.py
@@ -0,0 +1,41 @@
+"""Pipeline step: Commit step progress."""
+
+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, submodule_has_changes, commit_submodule
+
+log = get_logger(__name__)
+
+
+def commit_step(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ ascend_path = Path(ctx.triton_ascend_path)
+ step = ctx.steps[ctx.current_step]
+ step_id = step["id"]
+
+ if submodule_has_changes(ascend_path):
+ commit_submodule(
+ ascend_path,
+ f"[Sync](fix) AI fix for {ctx.target_commit[:12]}\n")
+
+ if not run_git(ascend_path, "status", "--porcelain").strip():
+ log.info(f"[{step_id}] Nothing to commit")
+ return ctx
+
+ end_short = step["end_commit"][:12]
+ msg = (
+ f"sync: merge upstream commits for step {step_id}\n\n"
+ f"Upstream range: {step.get('start_commit', '?')[:12]}..{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..69a7536
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/detect.py
@@ -0,0 +1,120 @@
+"""Pipeline step: Detect upstream commits to merge."""
+
+from __future__ import annotations
+import json
+import os
+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 import (
+ DETECT_FILE, WORKSPACE_DIR, ENV_BASE_BRANCH, get_base_branch_ref,
+)
+
+log = get_logger(__name__)
+
+
+def run_detect(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ 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:
+ ascend_path = Path(ctx.triton_ascend_path)
+ target = ctx.target_commit
+
+ # ── Fetch latest from the configured base branch (private fork) ──
+ base_branch = os.getenv(ENV_BASE_BRANCH, "main")
+ base_ref = get_base_branch_ref()
+ try:
+ run_git(ascend_path, "fetch", "origin", base_branch)
+ log.info(f"Fetched {base_ref} (private fork)")
+ except Exception:
+ log.warning(f"Could not fetch {base_ref}, using local refs")
+
+ # ── Resolve ascend_head from the base ref (may have been updated by fetch) ──
+ try:
+ ascend_head = run_git(ascend_path, "rev-parse", base_ref).strip()
+ except Exception:
+ ascend_head = ctx.ascend_head
+ log.warning(f"{base_ref} not available, using ascend_head from prepare")
+
+ 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]}).")
+
+ commits = _list_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 = _count_changed_lines(ascend_path, merge_base, target)
+
+ result = {
+ "ascend_head": ascend_head,
+ "target_commit": target,
+ "merge_base": merge_base,
+ "upstream_commits_count": len(commits),
+ "upstream_commits": commits,
+ "changed_lines": changed_lines,
+ "changed_files": changed_files,
+ "changed_files_count": len(changed_files),
+ }
+ (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=len(changed_files), changed_lines_total=changed_lines,
+ has_new_commits=has_new,
+ )
+
+
+def _list_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:
+ try:
+ output = run_git(repo, "diff", "--shortstat", merge_base, target)
+ except Exception:
+ return 0
+ 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..f175954
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/finalize.py
@@ -0,0 +1,64 @@
+"""Pipeline step: Finalize — generate cumulative patch and summary."""
+
+from __future__ import annotations
+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
+from TA_main2main_workflow.utils.tracker import total_elapsed
+from TA_main2main_workflow.utils import (
+ FINAL_SUMMARY_FILE, FINAL_TARGET_PATCH_FILE, WORKSPACE_DIR,
+)
+
+log = get_logger(__name__)
+
+
+def finalize(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ log.header("Finalize & Summary")
+ ascend_path = Path(ctx.triton_ascend_path)
+
+ # Summary
+ summary_path = WORKSPACE_DIR / FINAL_SUMMARY_FILE
+ summary_path.write_text(
+ f"# Triton-Ascend Upstream Sync\n\n"
+ f"- **Target**: `{ctx.target_commit[:12]}`\n"
+ f"- **Steps**: {ctx.total_steps}\n"
+ f"- **Upstream commits**: {ctx.upstream_commits_count}\n"
+ f"- **Status**: Success\n"
+ f"- **Date**: {time.strftime('%Y-%m-%d %H:%M:%S')}\n",
+ encoding="utf-8")
+ log.info(f"Final summary: {summary_path}")
+
+ # Cumulative patch
+ try:
+ patch = run_git(ascend_path, "diff", ctx.ascend_head, "HEAD")
+ (WORKSPACE_DIR / FINAL_TARGET_PATCH_FILE).write_text(patch, encoding="utf-8")
+ log.info(f"Cumulative patch: {len(patch)} bytes")
+ except Exception as e:
+ log.warning(f"Could not generate patch: {e}")
+
+ # PR body
+ pr_body_path = WORKSPACE_DIR / "pr_body.md"
+ lines = [
+ "## Triton-Ascend Upstream Sync\n\n",
+ f"**Target commit**: `{ctx.target_commit[:12]}`\n\n",
+ f"**Steps**: {ctx.total_steps}\n\n",
+ f"**Commits merged**: {ctx.upstream_commits_count}\n\n",
+ "### Step Details\n\n",
+ ]
+ for desc in ctx.step_pr_descriptions:
+ lines.append(f"- {desc}\n")
+ lines.append("\n---\nGenerated with [Claude Code](https://claude.com/claude-code)\n")
+ pr_body_path.write_text("".join(lines), encoding="utf-8")
+
+ log.header("Sync Complete!")
+ elapsed = total_elapsed()
+ log.elapsed(elapsed)
+ rows = [
+ ("Finalize", "PASS", f"{ctx.total_steps} step(s)"),
+ ("OVERALL", "PASS", f"{ctx.total_steps} step(s)"),
+ ]
+ log.table(rows)
+ return ctx
diff --git a/src/TA_main2main_workflow/pipeline/fix.py b/src/TA_main2main_workflow/pipeline/fix.py
new file mode 100644
index 0000000..8f15c43
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/fix.py
@@ -0,0 +1,124 @@
+"""Shared AI fix invocation — used by both build and test loops."""
+
+from __future__ import annotations
+import json, os, subprocess
+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 import (
+ FIX_LOG_DIR, STEPS_DIR, WORKSPACE_DIR,
+)
+
+log = get_logger(__name__)
+_REF = str(Path(__file__).parent.parent / "reference")
+
+
+class AIFixResult:
+ """Transient per-call result, stored on the flow orchestrator."""
+ def __init__(self):
+ self.modified_files: list[str] = []
+ self.step_summary: str = ""
+ self.is_noop: bool = False
+ self.elapsed_seconds: float = 0.0
+
+
+_ai_fix_result = AIFixResult()
+
+
+def get_last_ai_result() -> AIFixResult:
+ return _ai_fix_result
+
+
+def ai_fix(ctx: WorkflowContext, config: TAConfig, attempt: int = 1,
+ ascend_npu_ir_fix: bool = False) -> WorkflowContext:
+ """Invoke AI to fix build or test failures."""
+ global _ai_fix_result
+ _ai_fix_result = AIFixResult()
+
+ 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)
+
+ log.step(attempt, config.max_retries, "AI fix")
+ try:
+ result = run_opencode_adapter({
+ "step_id": f"{step_id}-fix-{attempt}",
+ "previous_step_id": step_id,
+ "previous_step_summary_path": str(step_dir / "step_summary.md"),
+ "is_last_step": "false",
+ "step_index": f"{ctx.current_step + 1}/{ctx.total_steps}",
+ "step_dir": str(step_dir),
+ "fix_dir": str(fix_dir),
+ "conflict_dir": "",
+ "ascend_path": str(ascend_path),
+ "triton_path": ctx.triton_path,
+ "reference_dir": _REF,
+ "mode": "fix",
+ "error_logs": json.dumps(ctx.fix_errors, ensure_ascii=False),
+ "target_commit": ctx.target_commit,
+ "ascend_npu_ir_fix": str(ascend_npu_ir_fix).lower(),
+ "ascend_npu_ir_compat_ref": str(Path(_REF) / "AscendNPU-IR_LLVM_VERSION_COMPAT.md"),
+ })
+ _ai_fix_result.modified_files = list(result.modified_files)
+ _ai_fix_result.step_summary = result.step_summary or ""
+ _ai_fix_result.is_noop = result.is_noop
+ _ai_fix_result.elapsed_seconds = result.elapsed_seconds
+ log.ai_result(
+ bool(result.modified_files),
+ result.modified_files,
+ (result.step_summary or "")[:500],
+ )
+ except Exception as e:
+ log.error(f"AI fix failed: {e}")
+
+ return ctx
+
+
+def validate_fix(modified_files: list[str], ascend_path: Path) -> tuple:
+ """Validate that an AI fix only touches allowed files.
+
+ When validation fails, illegal changes are reverted via git checkout.
+ Returns (passed: bool, reason: str).
+ """
+ if not modified_files:
+ return False, "No files were modified"
+
+ illegal_files: list[str] = []
+ ascend_root = str(ascend_path / "third_party" / "ascend")
+ for f in modified_files:
+ f_abs = str(Path(f).resolve()) if not Path(f).is_absolute() else f
+ if ascend_root not in f_abs:
+ illegal_files.append(f)
+
+ if illegal_files:
+ log.warning(f"Reverting invalid fix changes in {ascend_path}...")
+ try:
+ subprocess.run(
+ ["git", "checkout", "--", "."],
+ cwd=str(ascend_path), capture_output=True, text=True, timeout=30)
+ subprocess.run(
+ ["git", "clean", "-fd"],
+ cwd=str(ascend_path), capture_output=True, text=True, timeout=30)
+ log.status(True, "Reverted -- working tree is clean")
+ except Exception as e:
+ log.error(f"Failed to revert changes: {e}")
+ return False, (
+ f"Fix modified files OUTSIDE third_party/ascend/: "
+ + ", ".join(illegal_files)
+ + ". Changes have been reverted. "
+ + f"Next fix MUST only modify files under "
+ + f"{ascend_path}/third_party/ascend/")
+
+ log.status(True,
+ f"Fix validation: {len(modified_files)} file(s) all within third_party/ascend/")
+ return True, "All modified files are within third_party/ascend/"
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..f6306e3
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/ir_patch.py
@@ -0,0 +1,490 @@
+"""Pipeline step: IR compatibility patch — apply existing, supplement on test failure.
+
+New flow:
+ 1. Switch to target LLVM commit, clean workspace
+ 2. Directly apply existing llvm_patch_f6ded0b.patch (AI fix if needed)
+ 3. Build LLVM with patch
+ 4. Build TA + fix compile errors
+ 5. Run tests:
+ - Pass → done
+ - IR issues → AI supplements the existing patch → rebuild LLVM → retest
+ - Code issues → AI fixes → retest
+"""
+
+from __future__ import annotations
+import json, os, subprocess
+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.tracker import timed
+from TA_main2main_workflow.utils.git import run_git, run_git_no_check
+from TA_main2main_workflow.utils import (
+ WORKSPACE_DIR, BUILD_RESULT_FILE, BUILD_LOG_FILE,
+ IR_ANALYSIS_DIR, IR_OPS_REPORT_FILE, IR_CHANGES_REPORT_FILE,
+ IR_DIAGNOSIS_FILE, IR_MAX_ITERATIONS, STEPS_DIR,
+)
+from TA_main2main_workflow.pipeline.fix import ai_fix, validate_fix, get_last_ai_result
+
+log = get_logger(__name__)
+_REF = str(Path(__file__).parent.parent / "reference")
+_ASCEND_BASELINE_LLVM_HASH = "b5cc222d7429fe6f18c787f633d5262fac2e676f"
+
+
+def _llvm_project_path() -> Path:
+ return Path(os.path.expanduser(
+ os.getenv("LLVM_PROJECT_PATH", "~/llvm-project")))
+
+
+def per_step_ir_patch(ctx: WorkflowContext, config: TAConfig,
+ step: dict) -> WorkflowContext:
+ """IR patch pipeline when LLVM hash changed.
+
+ 1. Switch to target LLVM, clean workspace
+ 2. Apply existing llvm_patch_f6ded0b.patch (AI fix if failed)
+ 3. Build LLVM → Build TA + fix compile errors
+ 4. Test + supplement loop (IR issues → add to patch → rebuild → retest)
+ """
+ step_id = step["id"]
+ ascend_path = Path(ctx.triton_ascend_path)
+
+ if not _llvm_hash_did_change(ascend_path):
+ log.info(f"[{step_id}] LLVM hash unchanged -- skipping IR patch")
+ return ctx
+
+ llvm_hash_file = ascend_path / "cmake" / "llvm-hash.txt"
+ target_llvm_hash = llvm_hash_file.read_text(encoding="utf-8").strip()
+ llvm_project = _llvm_project_path()
+ llvm_install = Path(os.path.expanduser(
+ config.llvm_install_prefix_sync or "~/llvm-install-sync"))
+ ascend_patch = ascend_path / "third_party" / "ascend" / "patch" / "llvm_patch_f6ded0b.patch"
+
+ from TA_main2main_workflow.scripts.build_test import build_llvm, apply_llvm_patches
+ from TA_main2main_workflow.pipeline.build import _build_triton, _detect_ascend_npu_ir_errors
+
+ log.header(f"IR Patch Pipeline -- {step_id}")
+ log.key_value("Target LLVM", target_llvm_hash[:12])
+ log.key_value("Patch file", str(ascend_patch))
+
+ # ═══════════════════════════════════════════════════════════════════
+ # 1. Switch to target LLVM, clean workspace
+ # ═══════════════════════════════════════════════════════════════════
+ if not _ensure_llvm_workspace_clean():
+ log.error("Cannot clean llvm-project workspace")
+ return ctx.copy_with(build_passed=False)
+ try:
+ subprocess.run(
+ ["git", "checkout", target_llvm_hash],
+ cwd=str(llvm_project), capture_output=True, text=True, timeout=120)
+ log.info(f"Checked out target LLVM: {target_llvm_hash[:12]}")
+ except Exception as e:
+ log.error(f"Failed to checkout target LLVM: {e}")
+ return ctx.copy_with(build_passed=False)
+
+ # ═══════════════════════════════════════════════════════════════════
+ # 2. Apply existing patch directly (AI fix if needed)
+ # ═══════════════════════════════════════════════════════════════════
+ log.info(f"Applying existing patch: {ascend_patch.name}")
+ patch_ok = _apply_patch_with_retry(
+ ctx, config, ascend_path, ascend_patch, llvm_project,
+ target_llvm_hash, step_id, patch_error_type="apply")
+ if not patch_ok:
+ log.error("Patch apply failed after all retries")
+ return ctx.copy_with(build_passed=False)
+
+ # ═══════════════════════════════════════════════════════════════════
+ # 3. Build LLVM with patch
+ # ═══════════════════════════════════════════════════════════════════
+ log.info("Building LLVM with IR patch...")
+ try:
+ llvm_prefix = build_llvm(llvm_project, llvm_install,
+ required_hash=target_llvm_hash)
+ ctx = ctx.copy_with(llvm_prefix=str(llvm_prefix))
+ log.status(True, "LLVM build with patch complete")
+ except Exception as e:
+ log.error(f"LLVM build failed: {e}")
+ # AI fixes the patch based on build error, then retry
+ build_log = WORKSPACE_DIR / "llvm_build.log"
+ build_error = str(e)[:500]
+ if build_log.exists():
+ try:
+ log_tail = build_log.read_text(
+ encoding="utf-8", errors="replace")[-3000:]
+ build_error = f"Build exception: {e}\n\nBuild log tail:\n{log_tail}"
+ except Exception:
+ pass
+ log.warning("LLVM build failed — AI will fix the patch")
+ _fix_patch_and_retry(ctx, config, ascend_path, ascend_patch,
+ target_llvm_hash, "build", build_error, 1)
+ # Retry once after fix
+ if not _ensure_llvm_workspace_clean():
+ return ctx.copy_with(build_passed=False)
+ patch_ok = _apply_patch_with_retry(
+ ctx, config, ascend_path, ascend_patch, llvm_project,
+ target_llvm_hash, step_id, patch_error_type="build")
+ if not patch_ok:
+ return ctx.copy_with(build_passed=False)
+ try:
+ llvm_prefix = build_llvm(llvm_project, llvm_install,
+ required_hash=target_llvm_hash)
+ ctx = ctx.copy_with(llvm_prefix=str(llvm_prefix))
+ log.status(True, "LLVM rebuild after patch fix succeeded")
+ except Exception as e2:
+ log.error(f"LLVM build still failing after patch fix: {e2}")
+ return ctx.copy_with(build_passed=False)
+
+ # ═══════════════════════════════════════════════════════════════════
+ # 4. Build TA + fix compile errors
+ # ═══════════════════════════════════════════════════════════════════
+ log.info("Building TA with patched LLVM...")
+ ctx = _build_triton(ctx, config, clean=True)
+ if not ctx.build_passed:
+ if config.skip_ai_analysis:
+ return ctx.copy_with(build_passed=False)
+ log.warning("Build failed — entering compile-error fix loop")
+ for fix_attempt in range(1, config.max_retries + 1):
+ is_npu_ir = _detect_ascend_npu_ir_errors()
+ ctx = ctx.copy_with(fix_errors=[str(WORKSPACE_DIR / BUILD_RESULT_FILE)])
+ ctx = ai_fix(ctx, config, attempt=fix_attempt, ascend_npu_ir_fix=is_npu_ir)
+ modified_files = get_last_ai_result().modified_files
+ fix_valid, fix_reason = validate_fix(modified_files, ascend_path)
+ if not fix_valid:
+ log.error(f"Fix rejected: {fix_reason}")
+ continue
+ ctx = _build_triton(ctx, config, clean=False)
+ if ctx.build_passed:
+ break
+ if not ctx.build_passed:
+ log.error(f"TA build still failing after {config.max_retries} fixes")
+ return ctx.copy_with(build_passed=False, test_passed=False)
+ log.status(True, "TA builds successfully")
+
+ # ═══════════════════════════════════════════════════════════════════
+ # 5. Test + supplement patch loop
+ # ═══════════════════════════════════════════════════════════════════
+ return _test_and_supplement_loop(ctx, config, ascend_path, step, step_id,
+ target_llvm_hash, ascend_patch,
+ llvm_project, llvm_install)
+
+
+def _apply_patch_with_retry(ctx: WorkflowContext, config: TAConfig,
+ ascend_path: Path, ascend_patch: Path,
+ llvm_project: Path, target_llvm_hash: str,
+ step_id: str,
+ patch_error_type: str = "apply") -> bool:
+ """Try to apply the patch. If it fails, AI fixes it and retry.
+
+ Returns True if patch applied successfully.
+ """
+ from TA_main2main_workflow.scripts.build_test import apply_llvm_patches
+
+ for attempt in range(IR_MAX_ITERATIONS + 1):
+ if attempt > 0:
+ log.header(f"Patch Apply Retry {attempt}/{IR_MAX_ITERATIONS}")
+
+ # Clean + checkout before each attempt
+ if not _ensure_llvm_workspace_clean():
+ return False
+ subprocess.run(
+ ["git", "checkout", target_llvm_hash],
+ cwd=str(llvm_project), capture_output=True, text=True, timeout=120)
+
+ patch_result = apply_llvm_patches(
+ ascend_patch.parent, llvm_project,
+ target_hash=target_llvm_hash, patch_file=ascend_patch)
+
+ if patch_result["all_ok"]:
+ log.status(True, f"Patch applied: {ascend_patch.name}")
+ return True
+
+ failed = patch_result.get("failed", [])
+ error_msg = failed[0]['error'][:500] if failed else "unknown"
+ log.error(f"Patch apply failed: {error_msg}")
+
+ if attempt < IR_MAX_ITERATIONS:
+ log.warning("AI will analyze and fix the patch")
+ _fix_patch_for_apply(ctx, config, ascend_path, ascend_patch,
+ target_llvm_hash, error_msg, attempt + 1,
+ ascend_path / "third_party" / "ascend" / "patch")
+
+ log.error(f"Patch apply failed after {IR_MAX_ITERATIONS} retries")
+ return False
+
+
+def _fix_patch_for_apply(ctx: WorkflowContext, config: TAConfig,
+ ascend_path: Path, ascend_patch: Path,
+ target_llvm_hash: str, error_msg: str,
+ retry: int, step_dir: Path) -> None:
+ """AI analyzes why patch failed and adjusts it for the current LLVM commit."""
+ log.info(f"AI analyzing patch apply failure (retry {retry})...")
+ try:
+ run_opencode_adapter({
+ "step_id": f"ir-fix-patch-apply-{retry}",
+ "previous_step_id": "",
+ "previous_step_summary_path": "",
+ "is_last_step": "false",
+ "step_index": "ir",
+ "step_dir": str(step_dir),
+ "fix_dir": str(step_dir),
+ "conflict_dir": "",
+ "ascend_path": str(ascend_path),
+ "triton_path": ctx.triton_path,
+ "reference_dir": _REF,
+ "mode": "ir_fix_patch_apply",
+ "error_logs": json.dumps([], ensure_ascii=False),
+ "target_commit": ctx.target_commit,
+ "llvm_project_path": str(_llvm_project_path()),
+ "baseline_llvm_hash": _ASCEND_BASELINE_LLVM_HASH,
+ "target_llvm_hash": target_llvm_hash,
+ "ascend_patch_file": str(ascend_patch),
+ "patch_error_msg": error_msg,
+ })
+ except Exception as e:
+ log.error(f"AI patch fix failed: {e}")
+
+
+def _test_and_supplement_loop(ctx: WorkflowContext, config: TAConfig,
+ ascend_path: Path, step: dict, step_id: str,
+ target_llvm_hash: str, ascend_patch: Path,
+ llvm_project: Path,
+ llvm_install: Path) -> WorkflowContext:
+ """Run tests. On IR issues, supplement the existing patch and retry.
+
+ Does NOT regenerate from scratch — AI adds missing OP adaptations
+ to the existing patch file.
+ """
+ from TA_main2main_workflow.pipeline.build import _build_triton
+ from TA_main2main_workflow.pipeline.test import (
+ _run_pytest, _detect_oom_in_tests, _rerun_tests_reduced_concurrency,
+ _collect_test_error_logs)
+ from TA_main2main_workflow.scripts.build_test import build_llvm
+
+ for iteration in range(config.ir_max_iterations):
+ ctx = ctx.copy_with(ir_patch_iteration=iteration)
+ log.header(f"IR Test Loop — Iteration {iteration + 1}/{config.ir_max_iterations}")
+
+ # Run tests
+ ctx = _run_pytest(ctx, config)
+ if ctx.test_passed:
+ log.status(True, f"All tests pass for {step_id}")
+ return ctx.copy_with(test_passed=True)
+
+ # OOM handling
+ if _detect_oom_in_tests():
+ log.warning("NPU OOM detected — rerunning with reduced concurrency")
+ oom_result = _rerun_tests_reduced_concurrency(ascend_path, config)
+ if oom_result is None or oom_result:
+ return ctx.copy_with(test_passed=True)
+ if not _detect_oom_in_tests():
+ log.info("OOM resolved — classifying remaining failures")
+ else:
+ return ctx.copy_with(test_passed=False)
+
+ # Classify failures: IR vs code
+ log.warning("Tests failed — classifying failures (IR vs code)...")
+ has_ir_issues = _do_ir_diagnose_failures(ctx, ascend_path)
+
+ if has_ir_issues:
+ log.warning(
+ f"IR compatibility issues detected — "
+ f"AI will supplement the existing patch with missing OP adaptations "
+ f"(iteration {iteration + 1}/{config.ir_max_iterations})")
+
+ # AI supplements the existing patch (NOT from scratch)
+ ir_dir = WORKSPACE_DIR / IR_ANALYSIS_DIR
+ ir_dir.mkdir(parents=True, exist_ok=True)
+ _supplement_patch(ctx, config, ascend_path, ascend_patch,
+ target_llvm_hash, ir_dir, iteration + 1)
+
+ # Rebuild LLVM with supplemented patch
+ if not _ensure_llvm_workspace_clean():
+ return ctx.copy_with(test_passed=False)
+ subprocess.run(
+ ["git", "checkout", target_llvm_hash],
+ cwd=str(llvm_project), capture_output=True, text=True, timeout=120)
+ if not _apply_patch_with_retry(
+ ctx, config, ascend_path, ascend_patch, llvm_project,
+ target_llvm_hash, step_id, patch_error_type="apply"):
+ continue
+ try:
+ llvm_prefix = build_llvm(llvm_project, llvm_install,
+ required_hash=target_llvm_hash)
+ ctx = ctx.copy_with(llvm_prefix=str(llvm_prefix))
+ log.status(True, "LLVM rebuild with supplemented patch complete")
+ except Exception as e:
+ log.error(f"LLVM rebuild failed after patch supplement: {e}")
+ continue
+
+ # Rebuild TA
+ ctx = _build_triton(ctx, config, clean=False)
+ if not ctx.build_passed:
+ log.warning("TA build failed after patch supplement — will retry")
+ continue
+ continue # loop back to test
+
+ # Code issues → AI fix
+ log.warning("Code issues detected — AI fix")
+ ctx = ctx.copy_with(fix_errors=_collect_test_error_logs())
+ ctx = ai_fix(ctx, config, attempt=1)
+ ctx = _build_triton(ctx, config, clean=False)
+
+ log.error(f"IR test loop exhausted {config.ir_max_iterations} iterations")
+ return ctx.copy_with(test_passed=False)
+
+
+def _supplement_patch(ctx: WorkflowContext, config: TAConfig,
+ ascend_path: Path, ascend_patch: Path,
+ target_llvm_hash: str, ir_dir: Path,
+ iteration: int) -> None:
+ """AI supplements the existing patch with missing OP adaptations.
+
+ Key: the AI adds to the EXISTING patch, not generating from scratch.
+ It should:
+ 1. Read the current patch to understand what's already covered
+ 2. Diagnose which OPs still have IR compatibility issues
+ 3. Add missing OP adaptations to the patch
+ """
+ log.info(f"AI supplementing existing patch (iteration {iteration})...")
+ try:
+ run_opencode_adapter({
+ "step_id": f"ir-supplement-patch-{iteration}",
+ "previous_step_id": "ir-diagnose",
+ "previous_step_summary_path": str(ir_dir / IR_DIAGNOSIS_FILE),
+ "is_last_step": "false",
+ "step_index": "ir",
+ "step_dir": str(ascend_patch.parent),
+ "fix_dir": str(ascend_patch.parent),
+ "conflict_dir": "",
+ "ascend_path": str(ascend_path),
+ "triton_path": ctx.triton_path,
+ "reference_dir": _REF,
+ "mode": "ir_supplement_patch",
+ "error_logs": json.dumps(
+ [str(ir_dir / IR_DIAGNOSIS_FILE)], ensure_ascii=False),
+ "target_commit": ctx.target_commit,
+ "llvm_project_path": str(_llvm_project_path()),
+ "baseline_llvm_hash": _ASCEND_BASELINE_LLVM_HASH,
+ "target_llvm_hash": target_llvm_hash,
+ "ascend_patch_file": str(ascend_patch),
+ })
+ except Exception as e:
+ log.error(f"AI patch supplement failed: {e}")
+
+
+def _fix_patch_and_retry(ctx: WorkflowContext, config: TAConfig,
+ ascend_path: Path, ascend_patch: Path,
+ target_llvm_hash: str, error_type: str,
+ error_msg: str, retry: int) -> None:
+ """AI fixes a broken patch based on build error (kept for backward compat)."""
+ log.info(f"AI fixing patch ({error_type} failure, retry {retry})...")
+ try:
+ run_opencode_adapter({
+ "step_id": f"ir-fix-patch-{retry}",
+ "previous_step_id": "",
+ "previous_step_summary_path": "",
+ "is_last_step": "false",
+ "step_index": "ir",
+ "step_dir": str(ascend_patch.parent),
+ "fix_dir": str(ascend_patch.parent),
+ "conflict_dir": "",
+ "ascend_path": str(ascend_path),
+ "triton_path": ctx.triton_path,
+ "reference_dir": _REF,
+ "mode": "ir_fix_patch_apply",
+ "error_logs": json.dumps([], ensure_ascii=False),
+ "target_commit": ctx.target_commit,
+ "llvm_project_path": str(_llvm_project_path()),
+ "baseline_llvm_hash": _ASCEND_BASELINE_LLVM_HASH,
+ "target_llvm_hash": target_llvm_hash,
+ "ascend_patch_file": str(ascend_patch),
+ "patch_error_type": error_type,
+ "patch_error_msg": error_msg,
+ })
+ except Exception as e:
+ log.error(f"AI patch fix failed: {e}")
+
+
+def _do_ir_diagnose_failures(ctx: WorkflowContext, ascend_path: Path) -> bool:
+ """AI classifies test failures as IR vs code."""
+ ir_dir = WORKSPACE_DIR / IR_ANALYSIS_DIR
+ ir_dir.mkdir(parents=True, exist_ok=True)
+ from TA_main2main_workflow.pipeline.test import _collect_test_error_logs
+ error_log_paths = _collect_test_error_logs()
+
+ if not error_log_paths:
+ return False
+
+ try:
+ run_opencode_adapter({
+ "step_id": "ir-diagnose",
+ "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": ctx.triton_path,
+ "reference_dir": _REF,
+ "mode": "ir_diagnose",
+ "error_logs": json.dumps(error_log_paths, ensure_ascii=False),
+ })
+ except Exception:
+ return False
+
+ diagnosis_file = ir_dir / IR_DIAGNOSIS_FILE
+ if diagnosis_file.exists():
+ try:
+ data = json.loads(diagnosis_file.read_text(encoding="utf-8"))
+ return data.get("summary", {}).get("has_ir_issues", False)
+ except Exception:
+ pass
+ return False
+
+
+def _llvm_hash_did_change(ascend_path: Path) -> bool:
+ llvm_hash_file = ascend_path / "cmake" / "llvm-hash.txt"
+ if not llvm_hash_file.exists():
+ return False
+ current_hash = llvm_hash_file.read_text(encoding="utf-8").strip()
+ return current_hash != _ASCEND_BASELINE_LLVM_HASH
+
+
+def _ensure_llvm_workspace_clean() -> bool:
+ """Ensure the llvm-project working tree is clean."""
+ llvm_project = _llvm_project_path()
+ if not llvm_project.exists():
+ return True
+ try:
+ status = subprocess.run(
+ ["git", "status", "--porcelain"],
+ cwd=str(llvm_project), capture_output=True, text=True, timeout=15,
+ ).stdout.strip()
+ except Exception:
+ return True
+ if not status:
+ return True
+
+ log.warning("llvm-project has uncommitted changes -- cleaning...")
+ try:
+ subprocess.run(
+ ["git", "stash", "push", "-u", "-m", "ta-auto-clean"],
+ 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)
+ log.status(True, "Workspace cleaned")
+ return True
+ except Exception:
+ try:
+ subprocess.run(
+ ["git", "checkout", "--", "."],
+ cwd=str(llvm_project), capture_output=True, text=True, timeout=30)
+ subprocess.run(
+ ["git", "clean", "-fd"],
+ cwd=str(llvm_project), capture_output=True, text=True, timeout=30)
+ return True
+ 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..06fdc20
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/merge.py
@@ -0,0 +1,67 @@
+"""Pipeline step: Execute git merge of upstream commits into triton-ascend."""
+
+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.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
+
+log = get_logger(__name__)
+
+_MERGE_RESULT = "merge_result.json"
+
+
+def merge_upstream_commit(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ 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
+
+ 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", []),
+ )
+
+ if (ascend_path / ".git" / "MERGE_HEAD").exists():
+ try:
+ run_git(ascend_path, "merge", "--abort")
+ except Exception:
+ run_git(ascend_path, "reset", "--hard", "HEAD")
+
+ 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")
+
+ return ctx.copy_with(
+ merge_has_conflicts=has_conflicts, conflict_files=conflict_files
+ )
diff --git a/src/TA_main2main_workflow/pipeline/plan.py b/src/TA_main2main_workflow/pipeline/plan.py
new file mode 100644
index 0000000..e9467bf
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/plan.py
@@ -0,0 +1,167 @@
+"""Pipeline step: Plan merge steps — split upstream commits by line budget."""
+
+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.git import run_git
+from TA_main2main_workflow.utils.logging import get_logger
+from TA_main2main_workflow.utils import (
+ WORKSPACE_DIR, STEPS_FILE, STEPS_DIR, LLVM_HASH_FILE,
+)
+
+log = get_logger(__name__)
+
+
+def run_plan(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ 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:
+ ascend_path = Path(ctx.triton_ascend_path)
+ base = ctx.merge_base
+ target = ctx.target_commit
+ commits = ctx.upstream_commits
+
+ log.info(f"Scanning {len(commits)} upstream commits ({base[:8]}..{target[:8]})")
+
+ if config.progressive_merge and len(commits) > 1:
+ lines_per_commit, llvm_commits = _scan_commits(ascend_path, commits)
+ steps = _build_steps(commits, lines_per_commit, base,
+ config.line_budget, llvm_commits)
+ _enrich_steps(ascend_path, steps)
+ plan = {
+ "base_commit": base, "target_commit": target,
+ "line_budget": config.line_budget,
+ "total_steps": len(steps), "steps": steps,
+ }
+ _write_plan(plan)
+ log.info(f"Generated {len(steps)} step(s)")
+ return ctx.copy_with(steps=steps, total_steps=len(steps))
+
+ return ctx.copy_with(
+ total_steps=1,
+ steps=[{
+ "index": 1, "id": "step-1",
+ "commit_count": len(commits),
+ "start_commit": base, "end_commit": target,
+ "source_changed_lines": ctx.changed_lines_total,
+ }],
+ )
+
+
+# ── Internal helpers ────────────────────────────────────────────────────────
+
+def _source_lines_for_commit(repo: Path, sha: str) -> int:
+ try:
+ output = run_git(repo, "diff-tree", "--no-commit-id", "--shortstat", sha, quiet=True)
+ except Exception:
+ return 0
+ 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 _commit_changed_llvm_hash(repo: Path, sha: str) -> bool:
+ try:
+ output = run_git(repo, "diff-tree", "--no-commit-id", "--name-only", "-r", sha, quiet=True)
+ return LLVM_HASH_FILE in output
+ except Exception:
+ return False
+
+
+def _scan_commits(repo: Path, commits: list[dict]) -> tuple[dict, set]:
+ lines_per_commit: dict[str, int] = {}
+ llvm_commits: set[str] = set()
+ for i, c in enumerate(commits):
+ lines = _source_lines_for_commit(repo, c["sha"])
+ lines_per_commit[c["sha"]] = lines
+ if _commit_changed_llvm_hash(repo, c["sha"]):
+ llvm_commits.add(c["sha"])
+ if (i + 1) % 50 == 0:
+ log.info(f" ... scanned {i + 1}/{len(commits)} commits")
+ return lines_per_commit, llvm_commits
+
+
+def _build_steps(commits: list[dict], lines_per_commit: dict, base: str,
+ budget: int, llvm_commits: set) -> list[dict]:
+ steps: list[dict] = []
+ step_commits: list[dict] = []
+ step_lines = 0
+ start = base
+
+ for commit in commits:
+ sha = commit["sha"]
+ lines = lines_per_commit.get(sha, 0)
+
+ 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
+ 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
+ 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 _make_step(index: int, commits: list[dict], start: str, lines: int,
+ budget: int, reason: str = "line_budget") -> dict:
+ 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 _enrich_steps(repo: Path, steps: list[dict]) -> None:
+ for step in steps:
+ step["upstream_patch"] = run_git(
+ repo, "diff", f"{step['start_commit']}..{step['end_commit']}", quiet=True)
+ step["changed_files"] = run_git(
+ repo, "diff", "--name-only", f"{step['start_commit']}..{step['end_commit']}", quiet=True)
+
+
+def _write_plan(plan: dict) -> 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/prepare.py b/src/TA_main2main_workflow/pipeline/prepare.py
new file mode 100644
index 0000000..22a7e22
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/prepare.py
@@ -0,0 +1,121 @@
+"""Pipeline step: Prepare workspace — configure remotes, fetch, set up refs."""
+
+from __future__ import annotations
+import os
+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, ENV_BASE_BRANCH, get_base_branch_ref
+
+log = get_logger(__name__)
+
+
+def prepare(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Set up workspace: remotes, fetch, resolve ascend_head.
+
+ Does NOT force origin to point to any specific URL — the user may have
+ origin configured as their private fork (e.g. TecJesh/triton-ascend).
+ ascend_head is resolved from origin/{TA_BASE_BRANCH} so that previously
+ merged commits on the fork are reflected in the merge-base calculation.
+ """
+ WORKSPACE_DIR.mkdir(parents=True, exist_ok=True)
+
+ ascend_path = _ensure_repo(config)
+ _ensure_remote(ascend_path, "triton-upstream", config.triton_upstream_url)
+
+ # ── Record original branch before any checkout ──
+ try:
+ original_branch = run_git(ascend_path, "branch", "--show-current").strip()
+ if not original_branch:
+ original_branch = run_git(ascend_path, "rev-parse", "HEAD").strip()
+ except Exception:
+ original_branch = ""
+
+ # ── Abort any stale merge from a previous crashed run ──
+ merge_head = ascend_path / ".git" / "MERGE_HEAD"
+ if merge_head.exists():
+ log.warning("Found stale MERGE_HEAD from previous run, aborting it")
+ try:
+ run_git(ascend_path, "merge", "--abort")
+ log.info("Stale merge aborted successfully")
+ except Exception:
+ log.warning("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()
+
+ # ── Fetch origin (user's fork or upstream) ──
+ log.info("Fetching origin ...")
+ run_git(ascend_path, "fetch", "origin")
+
+ # ── Fetch triton-upstream ──
+ log.info("Fetching triton-upstream ...")
+ run_git(ascend_path, "fetch", "triton-upstream")
+
+ # ── Resolve ascend_head from the configured base branch ──
+ # Uses origin/{TA_BASE_BRANCH} so that when origin is a private fork
+ # (e.g. TecJesh/triton-ascend), previously merged commits are reflected
+ # in the merge-base calculation. Falls back to checkout HEAD gracefully.
+ 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:
+ log.warning(f"Could not fetch {base_ref}, using checkout HEAD as base")
+ try:
+ ascend_head = run_git(ascend_path, "rev-parse", base_ref).strip()
+ except Exception:
+ ascend_head = run_git(ascend_path, "rev-parse", "HEAD").strip()
+ log.warning(f"{base_ref} not available, using checkout HEAD")
+
+ # ── Resolve target commit ──
+ target_commit = config.target_commit
+ if not target_commit:
+ upstream_ref = "triton-upstream/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}'.")
+
+ log.section("Workspace ready")
+ log.key_value("triton-ascend", str(ascend_path))
+ log.key_value("base ref", base_ref)
+ log.key_value("ascend HEAD", ascend_head[:12])
+ log.key_value("target commit", target_commit[:12])
+ log.key_value("original branch", original_branch)
+
+ return ctx.copy_with(
+ triton_ascend_path=str(ascend_path),
+ target_commit=target_commit,
+ ascend_head=ascend_head,
+ original_branch=original_branch,
+ )
+
+
+def _ensure_repo(config: TAConfig) -> Path:
+ 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_DIR / "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_DIR, "clone", config.triton_ascend_url, str(target))
+ return target
+
+
+def _ensure_remote(repo: Path, name: str, url: str) -> None:
+ result = run_git_no_check(repo, "remote")
+ if name not in result.stdout:
+ run_git(repo, "remote", "add", name, url)
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..edd7f3a
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/push_pr.py
@@ -0,0 +1,51 @@
+"""Pipeline step: Push to GitHub and create PR."""
+
+from __future__ import annotations
+import os, 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.git import run_git
+from TA_main2main_workflow.utils import WORKSPACE_DIR
+
+log = get_logger(__name__)
+
+
+def push_and_create_pr(ctx: WorkflowContext, config: TAConfig) -> str:
+ """Push work branch and create GitHub PR. Returns PR URL."""
+ ascend_path = Path(ctx.triton_ascend_path)
+ branch = ctx.work_branch
+ repo = config.github_repo
+ base = config.pr_base_branch
+ target_short = ctx.target_commit[:12]
+
+ if config.gh_token:
+ os.environ["GH_TOKEN"] = config.gh_token
+ if "GH_HOST" not in os.environ:
+ os.environ["GH_HOST"] = "github.com"
+
+ log.info(f"Pushing {branch} ...")
+ run_git(ascend_path, "push", "-u", "origin", branch)
+
+ title = f"[Sync] Merge upstream {target_short}"
+ pr_body_path = WORKSPACE_DIR / "pr_body.md"
+ body = pr_body_path.read_text(encoding="utf-8") if pr_body_path.exists() else ""
+
+ log.info(f"Creating PR: {title}")
+ for attempt in range(1, 4):
+ try:
+ result = subprocess.run(
+ ["gh", "pr", "create",
+ "--base", base, "--head", branch,
+ "--title", title, "--body", body,
+ "--repo", repo],
+ capture_output=True, text=True, timeout=60)
+ if result.returncode == 0:
+ pr_url = result.stdout.strip()
+ log.status(True, f"PR created: {pr_url}")
+ return pr_url
+ log.warning(f"gh pr create failed (attempt {attempt}/3): {result.stderr.strip()[-200:]}")
+ except Exception as e:
+ log.warning(f"gh pr create failed (attempt {attempt}/3): {e}")
+ raise RuntimeError("Failed to create PR after 3 attempts")
diff --git a/src/TA_main2main_workflow/pipeline/resolve.py b/src/TA_main2main_workflow/pipeline/resolve.py
new file mode 100644
index 0000000..062a4a9
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/resolve.py
@@ -0,0 +1,79 @@
+"""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.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:
+ 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)
+
+ 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}",
+ "step_dir": str(step_dir),
+ "conflict_dir": str(step_dir),
+ "ascend_path": str(ascend_path),
+ "triton_path": ctx.triton_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
+
+ try:
+ from TA_main2main_workflow.scripts.pre_ci_check import run_pre_ci_check
+ run_pre_ci_check(ascend_path, step_id="conflict-resolution")
+ except Exception:
+ pass
+
+ try:
+ run_git(ascend_path, "add", "-A")
+ run_git(ascend_path, "commit", "--no-edit", "-s")
+ log.status(True, "Committed resolution")
+ except Exception:
+ pass
+ return ctx.copy_with(merge_has_conflicts=False)
diff --git a/src/TA_main2main_workflow/pipeline/test.py b/src/TA_main2main_workflow/pipeline/test.py
new file mode 100644
index 0000000..f6969ab
--- /dev/null
+++ b/src/TA_main2main_workflow/pipeline/test.py
@@ -0,0 +1,228 @@
+"""Pipeline step: Run pytest + test fix loop with OOM handling."""
+
+from __future__ import annotations
+import json, os, shutil, subprocess, time, 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.fix import ai_fix, validate_fix, get_last_ai_result
+
+log = get_logger(__name__)
+_MAX_OOM_RERUNS = 5
+
+
+def test(ctx: WorkflowContext, config: TAConfig) -> WorkflowContext:
+ """Test phase: run pytest with retry+fix loop."""
+ if config.skip_e2e_test:
+ log.info("SKIP_E2E_TEST=true -- treating tests as passed")
+ return ctx.copy_with(test_passed=True)
+
+ ascend_path = Path(ctx.triton_ascend_path)
+ step = ctx.steps[ctx.current_step] if 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)
+
+ test_passed = False
+ attempt = 0
+
+ while attempt <= config.max_retries:
+ is_fix_attempt = attempt > 0
+ ctx = ctx.copy_with(retry_count=attempt)
+
+ if is_fix_attempt:
+ # OOM detection: rerun with reduced concurrency
+ if _detect_oom_in_tests():
+ log.warning("NPU OOM detected -- rerunning with reduced concurrency")
+ oom_result = _rerun_tests_reduced_concurrency(
+ ascend_path, config, max_reruns=_MAX_OOM_RERUNS)
+ if oom_result is None or oom_result:
+ return ctx.copy_with(test_passed=True)
+ if not _detect_oom_in_tests():
+ log.info("OOM resolved -- remaining failures need AI fix")
+ else:
+ log.error(f"OOM persists after {_MAX_OOM_RERUNS} reruns")
+ return ctx.copy_with(test_passed=False)
+
+ log.header(f"Test Fix Attempt {attempt}/{config.max_retries}")
+ ctx = ctx.copy_with(fix_errors=_collect_test_error_logs())
+ ctx = ai_fix(ctx, config, attempt=attempt)
+
+ # Validate fix
+ modified_files = get_last_ai_result().modified_files
+ fix_valid, fix_reason = validate_fix(modified_files, ascend_path)
+ if not fix_valid:
+ log.error(f"Fix rejected: {fix_reason}")
+ rejection_file = step_dir / "fix_rejection.txt"
+ rejection_file.write_text(
+ f"PREVIOUS FIX WAS REJECTED: {fix_reason}\n"
+ f"For test fixes, prefer files under "
+ f"{ascend_path}/third_party/ascend/.\n",
+ encoding="utf-8")
+ ctx = ctx.copy_with(
+ fix_errors=ctx.fix_errors + [str(rejection_file)])
+ continue
+
+ # Rebuild after fix
+ from TA_main2main_workflow.pipeline.build import _build_triton
+ ctx = _build_triton(ctx, config, clean=False)
+ if not ctx.build_passed:
+ attempt += 1
+ continue
+
+ # Run tests
+ with timed("test"):
+ ctx = _run_pytest(ctx, config)
+ if ctx.test_passed:
+ test_passed = True
+ break
+
+ log.info(f"Tests failed (attempt {attempt + 1}) -- retrying")
+ attempt += 1
+
+ if config.skip_ai_analysis:
+ log.warning("SKIP_AI_ANALYSIS=true -- stopping test fix loop")
+ break
+
+ return ctx.copy_with(test_passed=test_passed)
+
+
+def _run_pytest(ctx: WorkflowContext, config: TAConfig,
+ python_exe: str = "") -> WorkflowContext:
+ """Execute pytest and return updated ctx."""
+ ascend_path = Path(ctx.triton_ascend_path)
+ test_log_dir = WORKSPACE_DIR / "test-logs"
+ test_log_dir.mkdir(parents=True, exist_ok=True)
+
+ test_dir_path = (ascend_path / config.test_dir).resolve()
+ python_exe = python_exe or os.getenv("PYTHON", "python3.10")
+
+ if not test_dir_path.exists():
+ log.warning(f"Test directory not found: {test_dir_path}")
+ return ctx.copy_with(test_passed=True)
+
+ junit_xml = test_log_dir / "pytest-junit.xml"
+ pytest_bin = shutil.which("pytest")
+ cmd = (
+ [pytest_bin, str(test_dir_path)]
+ if pytest_bin
+ else [python_exe, "-m", "pytest", str(test_dir_path)]
+ )
+ cmd += ["-n", str(config.test_procs), f"--junitxml={junit_xml}"]
+
+ log.section("Run Tests")
+ log.info(f"cmd: {' '.join(cmd)}")
+ _start = time.time()
+ try:
+ result = subprocess.run(cmd, cwd=ascend_path, timeout=1000)
+ rc = result.returncode
+ except subprocess.TimeoutExpired:
+ rc = -1
+ log.warning("pytest timed out after 1000s")
+
+ elapsed = time.time() - _start
+ log.info(f"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
+ 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,
+ }
+ (WORKSPACE_DIR / TEST_RESULT_FILE).write_text(
+ json.dumps(summary, indent=2, ensure_ascii=False) + "\n", encoding="utf-8")
+
+ if not passed:
+ log.error(f"Tests FAILED ({pf} failed, {pe} errors)")
+ return ctx.copy_with(
+ test_passed=False,
+ fix_errors=[str(WORKSPACE_DIR / TEST_RESULT_FILE)])
+ log.status(True, f"All tests passed ({tp} passed)")
+ return ctx.copy_with(test_passed=True)
+
+
+def _detect_oom_in_tests() -> bool:
+ """Check test logs for OOM errors."""
+ test_log_dir = WORKSPACE_DIR / "test-logs"
+ oom_markers = ["out of memory"]
+ if test_log_dir.exists():
+ try:
+ for log_file in test_log_dir.rglob("*"):
+ if log_file.suffix not in (".log", ".xml"):
+ continue
+ content = log_file.read_text(encoding="utf-8", errors="replace")
+ for marker in oom_markers:
+ if marker.lower() in content.lower():
+ return True
+ except Exception:
+ pass
+ test_result = WORKSPACE_DIR / TEST_RESULT_FILE
+ if test_result.exists():
+ try:
+ data = json.loads(test_result.read_text(encoding="utf-8"))
+ for marker in oom_markers:
+ if marker.lower() in json.dumps(data).lower():
+ return True
+ except Exception:
+ pass
+ return False
+
+
+def _rerun_tests_reduced_concurrency(ascend_path: Path, config: TAConfig,
+ max_reruns: int = 5) -> bool | None:
+ """Rerun tests with halved concurrency. Restores original after."""
+ original_procs = config.test_procs
+ reduced = max(1, original_procs // 2)
+ log.warning(f"Reducing pytest concurrency: {original_procs} -> {reduced}")
+ # Temporarily override config
+ object.__setattr__(config, 'test_procs', reduced)
+ try:
+ for rerun in range(1, max_reruns + 1):
+ log.info(f"OOM rerun {rerun}/{max_reruns} (procs={reduced})")
+ # Create a temp ctx just for the test run
+ temp_ctx = WorkflowContext(triton_ascend_path=str(ascend_path))
+ temp_ctx = _run_pytest(temp_ctx, config)
+ if temp_ctx.test_passed:
+ return True
+ if not _detect_oom_in_tests():
+ log.info("OOM resolved -- remaining failures are not memory-related")
+ return False
+ return False
+ finally:
+ object.__setattr__(config, 'test_procs', original_procs)
+ log.info(f"Restored pytest concurrency to {original_procs}")
+
+
+def _collect_test_error_logs() -> list[str]:
+ """Collect test failure log paths."""
+ error_logs: 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_logs.append(str(log_file))
+ for xml_file in sorted(test_log_dir.rglob("*.xml")):
+ error_logs.append(str(xml_file))
+ test_result = WORKSPACE_DIR / TEST_RESULT_FILE
+ if test_result.exists():
+ error_logs.append(str(test_result))
+ return error_logs
diff --git a/src/TA_main2main_workflow/utils/__init__.py b/src/TA_main2main_workflow/utils/__init__.py
new file mode 100644
index 0000000..3cace9d
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/__init__.py
@@ -0,0 +1,55 @@
+"""TA_main2main_workflow utilities — config, context, git, logging, tracking."""
+
+from __future__ import annotations
+
+import os
+from pathlib import Path
+
+# ── Workspace path ─────────────────────────────────────────────────────────
+WORKSPACE_DIR = Path(os.getenv("TA_MAIN2MAIN_WORKSPACE", str(Path.cwd() / "workspace")))
+
+# ── Flow routing signals ───────────────────────────────────────────────────
+UpgradeCompleted = "UpgradeCompleted"
+UpgradeFailed = "UpgradeFailed"
+
+# ── Constants ──────────────────────────────────────────────────────────────
+LLVM_HASH_FILE = "cmake/llvm-hash.txt"
+ENV_BASE_BRANCH = "TA_BASE_BRANCH"
+ENV_SINGLE_STEP_MODE = "ENV_SINGLE_STEP_MODE"
+IR_MAX_ITERATIONS = 10 # max retries for patch apply/rebuild
+LLVM_CHANGE_ANALYSIS_DIR = "llvm_change_analysis"
+
+# ── Output file names ──────────────────────────────────────────────────────
+DETECT_FILE = "detect.json"
+STEPS_FILE = "steps.json"
+BUILD_RESULT_FILE = "build_result.json"
+BUILD_LOG_FILE = "build.log"
+TEST_RESULT_FILE = "test_result.json"
+FIX_LOG_DIR = "fixes"
+STEPS_DIR = "steps"
+FINAL_SUMMARY_FILE = "final_summary.md"
+FINAL_TARGET_PATCH_FILE = "final_target.patch"
+PRE_CI_CHECK_FILE = "pre_ci_check.json"
+EACH_STEP_SUMMARY_FILE = "step_summary.md"
+
+# IR analysis constants
+IR_ANALYSIS_DIR = "ir-analysis"
+IR_OPS_REPORT_FILE = "ops_report.json"
+IR_CHANGES_REPORT_FILE = "changes_report.json"
+IR_DIAGNOSIS_FILE = "ir_diagnosis.json"
+
+# ── Helpers ────────────────────────────────────────────────────────────────
+
+
+def get_base_branch_ref(remote: str = "origin") -> str:
+ branch = os.getenv(ENV_BASE_BRANCH, "main")
+ return f"{remote}/{branch}"
+
+
+# ── Re-exports ─────────────────────────────────────────────────────────────
+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, run_git_no_check, submodule_has_changes, commit_submodule)
+from TA_main2main_workflow.utils.logging import get_logger, TALogger
+from TA_main2main_workflow.utils.submodule import push_submodule
diff --git a/src/TA_main2main_workflow/utils/config.py b/src/TA_main2main_workflow/utils/config.py
new file mode 100644
index 0000000..cc777df
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/config.py
@@ -0,0 +1,146 @@
+"""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 typing import Literal
+
+
+AIBackendChoice = Literal["opencode", "claude", "auto"]
+
+
+@dataclass
+class TAConfig:
+ """User-configurable parameters for a workflow run."""
+
+ # ── Repository ────────────────────────────────────────────────────────
+ triton_ascend_path: str = ""
+ triton_ascend_url: str = "https://github.com/triton-lang/triton-ascend.git"
+ triton_path: str = ""
+ 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
+
+ # ── IR Patch ──────────────────────────────────────────────────────────
+ ir_max_iterations: int = 3
+ ascend_baseline_llvm_hash: str = (
+ "b5cc222d7429fe6f18c787f633d5262fac2e676f")
+
+ # ── Build / Test parallelism ──────────────────────────────────────────
+ llvm_install_prefix_sync: str = "" # LLVM_INSTALL_PREFIX_SYNC
+ llvm_project_path: str = "" # LLVM_PROJECT_PATH
+ llvm_repo_url: str = "https://github.com/llvm/llvm-project.git"
+ build_procs: int = 32
+ test_procs: int = 16 # pytest -n
+ test_dir: str = "third_party/ascend/unittest/pytest_ut"
+ conda_env: str = ""
+
+ # ── Skip flags ────────────────────────────────────────────────────────
+ resume: bool = False
+ skip_ai_analysis: bool = False
+ skip_build: bool = False
+ skip_e2e_test: bool = False
+ skip_llvm_rebuild: bool = False
+ skip_ir_patch: bool = False
+
+ # ── Git / Branch ──────────────────────────────────────────────────────
+ base_branch: str = "main"
+ progressive_merge: bool = True
+ single_step_mode: bool = True
+
+ # ── PR / Push ─────────────────────────────────────────────────────────
+ push_to_github: bool = False
+ github_repo: str = "triton-lang/triton-ascend"
+ gh_token: str = ""
+ pr_base_branch: str = "upstream-sync"
+
+ # ── Workspace ─────────────────────────────────────────────────────────
+ workspace_dir: str = ""
+
+ # ═══════════════════════════════════════════════════════════════════════
+ @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),
+ ir_max_iterations=_env_int("IR_MAX_ITERATIONS", 3),
+ ascend_baseline_llvm_hash=os.getenv(
+ "ASCEND_BASELINE_LLVM_HASH",
+ "b5cc222d7429fe6f18c787f633d5262fac2e676f"),
+ llvm_install_prefix_sync=os.getenv(
+ "LLVM_INSTALL_PREFIX_SYNC", ""),
+ llvm_project_path=os.getenv("LLVM_PROJECT_PATH", ""),
+ llvm_repo_url=os.getenv(
+ "LLVM_REPO_URL",
+ "https://github.com/llvm/llvm-project.git"),
+ build_procs=_env_int("BUILD_PROCS", 32),
+ test_procs=_env_int("NUM_PROCS", 16),
+ test_dir=os.getenv(
+ "TEST_DIR", "third_party/ascend/unittest/pytest_ut"),
+ conda_env=os.getenv("CONDA_DEFAULT_ENV", ""),
+ 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_ir_patch=_env_bool("SKIP_IR_PATCH", False),
+ base_branch=os.getenv("TA_BASE_BRANCH", "main"),
+ progressive_merge=_env_bool("TA_PROGRESSIVE_MERGE", True),
+ single_step_mode=_env_bool("ENV_SINGLE_STEP_MODE", True),
+ push_to_github=_env_bool("PUSH_TO_GITHUB", False),
+ github_repo=os.getenv("GITHUB_REPO", "triton-lang/triton-ascend"),
+ gh_token=os.getenv("GH_TOKEN", ""),
+ pr_base_branch=os.getenv("TA_PR_BASE_BRANCH", "upstream-sync"),
+ workspace_dir=os.getenv("TA_MAIN2MAIN_WORKSPACE", ""),
+ )
+
+
+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_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..cc5c6b9
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/context.py
@@ -0,0 +1,99 @@
+"""WorkflowContext — shared state carrier between pipeline steps."""
+
+from __future__ import annotations
+
+from dataclasses import dataclass, field, replace
+from pathlib import Path
+
+
+@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 paths (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 ──────────────────────────────────────────────
+ 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 ──────────────────────────────────────────────────────
+ steps: list[dict] = field(default_factory=list)
+ total_steps: int = 0
+ current_step: int = 0
+ step_start_ascend_head: str = ""
+ step_pr_descriptions: list = field(default_factory=list)
+
+ # ── Merge results ──────────────────────────────────────────────────
+ merge_has_conflicts: bool = False
+ conflict_files: list[str] = field(default_factory=list)
+ conflict_files_resolved: int = 0
+
+ # ── Build / test results ───────────────────────────────────────────
+ build_passed: bool = False
+ test_passed: bool = False
+ fix_errors: list[str] = field(default_factory=list)
+
+ # ── Retry tracking ─────────────────────────────────────────────────
+ retry_count: int = 0
+ build_fix_count: int = 0
+ test_fix_count: int = 0
+ fix_attempts: list = field(default_factory=list)
+ step_details: list = field(default_factory=list)
+
+ # ── LLVM / IR patch state ──────────────────────────────────────────
+ llvm_prefix: str = ""
+ llvm_hash_changed: bool = False
+ 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_issues_found: int = 0
+ ir_fix_count: int = 0
+ ir_loop_details: list = field(default_factory=list)
+
+ # ── Pytest state ───────────────────────────────────────────────────
+ pytest_passed: bool = False
+ test_failures_by_python: dict = field(default_factory=dict)
+ test_log_dir: str = ""
+ test_dir: str = "third_party/ascend/unittest/pytest_ut"
+ num_procs: int = 16
+ conda_env: str = ""
+
+ # ── Status / output ────────────────────────────────────────────────
+ final_status: str = ""
+ pr_url: str = ""
+ summary_rows: list = field(default_factory=list)
+
+ # ═══════════════════════════════════════════════════════════════════
+ # Helpers
+ # ═══════════════════════════════════════════════════════════════════
+
+ def copy_with(self, **kwargs) -> WorkflowContext:
+ """Return a new WorkflowContext with the given fields updated."""
+ 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..2539ce8
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/git.py
@@ -0,0 +1,93 @@
+"""Git utilities with auto-retry for network operations."""
+
+from __future__ import annotations
+
+import subprocess
+import time
+from pathlib import Path
+
+from TA_main2main_workflow.utils.logging import get_logger
+
+log = get_logger("git")
+
+_RETRIES = 3
+_RETRY_DELAY = 5
+_RETRY_OPS = {"fetch", "clone", "push", "pull", "remote", "ls-remote"}
+
+
+def run_git(repo: Path | str, *args: str, quiet: bool = False) -> str:
+ """Run a git command in *repo*, return stdout.
+
+ Auto-retries for network operations (fetch, clone, push, etc.).
+ Raises subprocess.CalledProcessError on failure.
+ """
+ repo_path = Path(repo)
+ repo_name = repo_path.name if repo_path.is_dir() else str(repo)
+ cmd = ["git", "-C", str(repo_path), *args]
+ op = args[0] if args else ""
+
+ if not quiet:
+ log.info(f"[{repo_name}] $ git {' '.join(args)}")
+
+ last_exc = None
+ for attempt in range(1, (_RETRIES if op in _RETRY_OPS else 1) + 1):
+ try:
+ result = subprocess.run(
+ cmd, capture_output=True, text=True, check=True, timeout=300)
+ # Log preview on success
+ if not quiet and result.stdout.strip():
+ preview = result.stdout.strip()[:200]
+ if len(result.stdout.strip()) > 200:
+ preview += "..."
+ log.info(f" → {preview}")
+ return result.stdout
+ except subprocess.CalledProcessError as e:
+ last_exc = e
+ if op in _RETRY_OPS and attempt < _RETRIES:
+ delay = _RETRY_DELAY * attempt
+ log.warning(
+ f"git {op} failed (attempt {attempt}/{_RETRIES}) — "
+ f"retrying in {delay}s: {e.stderr.strip()[-200:]}")
+ time.sleep(delay)
+ else:
+ log.error(f"git {op} failed: {e.stderr.strip()[-500:]}")
+ raise
+
+ raise last_exc # type: ignore[misc]
+
+
+def run_git_no_check(repo: Path | str, *args: str) -> subprocess.CompletedProcess:
+ """Run a git command, never raise on non-zero exit. Returns CompletedProcess."""
+ repo_path = Path(repo)
+ cmd = ["git", "-C", str(repo_path), *args]
+ return subprocess.run(cmd, capture_output=True, text=True, timeout=300)
+
+
+def submodule_has_changes(repo: Path) -> bool:
+ """Check whether the AscendNPU-IR submodule has uncommitted changes."""
+ ascend_npu_ir = repo / "third_party" / "ascend" / "AscendNPU-IR"
+ if not ascend_npu_ir.exists():
+ return False
+ result = subprocess.run(
+ ["git", "status", "--porcelain"],
+ cwd=str(ascend_npu_ir),
+ capture_output=True, text=True, timeout=30,
+ )
+ return bool(result.stdout.strip())
+
+
+def commit_submodule(repo: Path, message: str) -> None:
+ """Commit changes inside the AscendNPU-IR submodule."""
+ ascend_npu_ir = repo / "third_party" / "ascend" / "AscendNPU-IR"
+ if not submodule_has_changes(repo):
+ return
+ subprocess.run(
+ ["git", "add", "-A"], cwd=str(ascend_npu_ir),
+ capture_output=True, text=True, timeout=30,
+ )
+ subprocess.run(
+ ["git", "commit", "-s", "-m", message],
+ cwd=str(ascend_npu_ir),
+ capture_output=True, text=True, timeout=30,
+ )
+ log.info("Committed AscendNPU-IR submodule changes")
diff --git a/src/TA_main2main_workflow/utils/logging.py b/src/TA_main2main_workflow/utils/logging.py
new file mode 100644
index 0000000..782a0b2
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/logging.py
@@ -0,0 +1,109 @@
+"""Rich logging for pipeline steps — TALogger with formatted output."""
+
+from __future__ import annotations
+
+import logging
+import sys
+import time
+from datetime import datetime
+from typing import Any
+
+
+class TALogger(logging.getLoggerClass()): # type: ignore
+ """Logger with pipeline-specific formatting methods."""
+
+ def header(self, title: str) -> None:
+ width = max(68, len(title) + 6)
+ self.info("")
+ self.info("╔" + "═" * (width - 2) + "╗")
+ self.info(f"║ {title}".ljust(width - 1) + "║")
+ self.info("╚" + "═" * (width - 2) + "╝")
+
+ def section(self, title: str) -> None:
+ ts = datetime.now().strftime("%H:%M:%S")
+ self.info(f"\n── {title} [{ts}] ──")
+
+ 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: Any, **kwargs: Any) -> None:
+ self.warning(f" ⚠ {msg}", *args, **kwargs)
+
+ def error(self, msg: str, *args: Any, **kwargs: Any) -> None:
+ self.error_raw(f" ✘ {msg}", *args, **kwargs)
+
+ def error_raw(self, msg: str, *args: Any, **kwargs: Any) -> None:
+ super().error(msg, *args, **kwargs)
+
+ def key_value(self, key: str, value: Any) -> None:
+ self.info(f" {key}: {value}")
+
+ def info(self, msg: str = "", *args: Any, **kwargs: Any) -> None:
+ super().info(msg, *args, **kwargs)
+
+ def ai_call(self, backend: str, mode: str, attempt: int, max_attempts: int) -> None:
+ self.info("╭─ AI Call ──────────────────────────────────────────╮")
+ self.info(f"│ Backend: {backend} Mode: {mode} Attempt: {attempt}/{max_attempts}")
+ self.info("╰────────────────────────────────────────────────────╯")
+
+ def ai_result(self, ok: bool, modified_files: list = (), summary: str = "") -> None:
+ self.info("╭─ AI Result ────────────────────────────────────────╮")
+ self.info(f"│ {'✔' if ok else '✘'} modified: {len(modified_files) if modified_files else 0} file(s)")
+ if summary:
+ for line in summary[:500].splitlines()[:5]:
+ self.info(f"│ {line[:72]}")
+ self.info("╰────────────────────────────────────────────────────╯")
+
+ def table(self, rows: list[tuple[str, str, str]]) -> None:
+ self.info("")
+ self.info(" SYNC SUMMARY")
+ self.info(" " + "─" * 50)
+ for phase, status, detail in rows:
+ icon = "✔" if status == "PASS" else ("✘" if status == "FAIL" else "⚠")
+ self.info(f" {icon} {phase:<25} {detail}")
+ self.info(" " + "─" * 50)
+
+ def elapsed(self, seconds: float) -> None:
+ m, s = divmod(int(seconds), 60)
+ h, m = divmod(m, 60)
+ if h:
+ self.info(f" Total elapsed: {h}h {m}m {s}s")
+ elif m:
+ self.info(f" Total elapsed: {m}m {s}s")
+ else:
+ self.info(f" Total elapsed: {s}s")
+
+ 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:
+ self.info(f" Conflicted files ({len(files)}):")
+ for i, f in enumerate(files, 1):
+ self.info(f" {i}. {f}")
+
+
+# ── Module-level 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 quick 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..b7dd6fe
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/submodule.py
@@ -0,0 +1,24 @@
+"""Submodule helpers for AscendNPU-IR."""
+
+from __future__ import annotations
+
+from pathlib import Path
+
+from TA_main2main_workflow.utils.git import submodule_has_changes, commit_submodule, run_git
+
+__all__ = ["submodule_has_changes", "commit_submodule", "push_submodule"]
+
+
+def push_submodule(repo: Path, branch: str) -> bool:
+ """Push AscendNPU-IR submodule changes to its remote.
+
+ Returns True on success, False on failure.
+ """
+ ascend_npu_ir = repo / "third_party" / "ascend" / "AscendNPU-IR"
+ if not ascend_npu_ir.exists():
+ return False
+ try:
+ run_git(ascend_npu_ir, "push", "-u", "origin", branch)
+ return True
+ except Exception:
+ 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..f3c966e
--- /dev/null
+++ b/src/TA_main2main_workflow/utils/tracker.py
@@ -0,0 +1,33 @@
+"""Timing utilities — context manager for phase-level timing."""
+
+from __future__ import annotations
+
+import time
+from contextlib import contextmanager
+
+from TA_main2main_workflow.utils.logging import get_logger
+
+log = get_logger("tracker")
+
+_flow_start_time: float = 0.0
+
+
+@contextmanager
+def timed(name: str):
+ """Context manager that records elapsed time for a named phase."""
+ global _flow_start_time
+ if _flow_start_time == 0.0:
+ _flow_start_time = time.time()
+ start = time.time()
+ try:
+ yield
+ finally:
+ elapsed = time.time() - start
+ log.info(f" ⏱ {name} took {elapsed:.1f}s")
+
+
+def total_elapsed() -> float:
+ """Return total seconds since the first timed() call."""
+ if _flow_start_time == 0.0:
+ return 0.0
+ return time.time() - _flow_start_time