Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 25 additions & 0 deletions src/TA_main2main_workflow/external_test/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
"""External operator repository pluggable test cases.

Provides config loading and a runner for executing pytest suites
in external (non-vendored) operator repositories, with AI-driven
fix retry loops on failure.

Core entry points:
- :func:`load_external_test_config` — parse a YAML config file
- :func:`run_external_tests` — clone, install, test, fix (per repo)
- :class:`ExternalTestConfig` / :class:`ExternalTestRepoConfig` — dataclasses
"""

from TA_main2main_workflow.external_test.config_loader import (
ExternalTestConfig,
ExternalTestRepoConfig,
load_external_test_config,
)
from TA_main2main_workflow.external_test.runner import run_external_tests

__all__ = [
"ExternalTestConfig",
"ExternalTestRepoConfig",
"load_external_test_config",
"run_external_tests",
]
204 changes: 204 additions & 0 deletions src/TA_main2main_workflow/external_test/config_loader.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,204 @@
"""External test configuration loader and validation.

Parses ``external_test_config.yaml`` into typed dataclasses and validates
URL / path correctness before the runner consumes them.
"""

from __future__ import annotations

import os
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any

import yaml # type: ignore[import-untyped]

from TA_main2main_workflow.utils.logging import get_logger

log = get_logger(__name__)

# ── Default config shipped with the package ───────────────────────────────
_DEFAULT_CONFIG_DIR = Path(__file__).resolve().parent
_DEFAULT_CONFIG_PATH = _DEFAULT_CONFIG_DIR / "external_test_config.yaml"


@dataclass
class ExternalTestRepoConfig:
"""Configuration for a single external operator repository."""

name: str # display name
url: str # git clone URL
branch: str = "main" # branch / tag to checkout
test_cases: list[str] = field(default_factory=list) # test file relative paths
install_cmd: str = "" # optional dep install (empty = skip)


@dataclass
class ExternalTestConfig:
"""Top-level external test configuration."""

enabled: bool = False # master on/off switch
repos: list[ExternalTestRepoConfig] = field(default_factory=list)
test_procs: int = 8 # pytest -n <N>
mode: str = "inline" # inline | standalone | off
max_retries: int = 5 # AI fix retries per repo
timeout: int = 7200 # per-repo test timeout (seconds)


# ═══════════════════════════════════════════════════════════════════════════
# Public API
# ═══════════════════════════════════════════════════════════════════════════


def load_external_test_config(path: str = "") -> ExternalTestConfig | None:
"""Load external test config from *path* (YAML).

Resolution order:
1. Explicit *path* argument
2. ``TA_EXTERNAL_TEST_CONFIG`` environment variable
3. Default ``external_test/external_test_config.yaml`` shipped with package

Returns ``None`` when the config file does not exist (not an error — the
caller treats missing config as "no external tests configured").
"""
resolved = _resolve_config_path(path)
if resolved is None:
return None

log.info(f"Loading external test config: {resolved}")
try:
raw = yaml.safe_load(resolved.read_text(encoding="utf-8")) or {}
except yaml.YAMLError as exc:
log.error(f"Failed to parse external test config: {exc}")
return None

cfg = _dict_to_config(raw)

# Merge env-var overrides (env takes precedence over YAML)
_apply_env_overrides(cfg)

if not validate_config(cfg):
return None

log.key_value("External test enabled", str(cfg.enabled))
log.key_value("External test mode", cfg.mode)
log.key_value("External test repos", str(len(cfg.repos)))
return cfg


def validate_config(cfg: ExternalTestConfig) -> bool:
"""Validate the loaded config. Returns True if usable."""
if cfg.mode not in ("inline", "standalone", "off"):
log.error(f"Invalid external test mode: {cfg.mode}")
return False

if cfg.test_procs < 1:
log.error("test_procs must be >= 1")
return False

if cfg.max_retries < 0:
log.error("max_retries must be >= 0")
return False

for repo in cfg.repos:
if not repo.name:
log.error("External test repo missing 'name'")
return False
if not repo.url or not (
repo.url.startswith("http") or repo.url.startswith("git@")
):
log.error(f"Invalid repo URL for '{repo.name}': {repo.url}")
return False
if not repo.test_cases:
log.warning(f"External repo '{repo.name}' has no test_cases configured")

return True


# ═══════════════════════════════════════════════════════════════════════════
# Internal helpers
# ═══════════════════════════════════════════════════════════════════════════


def _resolve_config_path(explicit: str) -> Path | None:
"""Determine which config file to read."""
if explicit:
p = Path(explicit)
if p.exists():
return p
log.warning(f"External test config not found: {explicit}")
return None

env_path = os.getenv("TA_EXTERNAL_TEST_CONFIG", "")
if env_path:
p = Path(env_path)
if p.exists():
return p
log.warning(f"TA_EXTERNAL_TEST_CONFIG points to missing file: {env_path}")

if _DEFAULT_CONFIG_PATH.exists():
return _DEFAULT_CONFIG_PATH

return None


def _dict_to_config(raw: dict[str, Any]) -> ExternalTestConfig:
"""Convert raw YAML dict to ExternalTestConfig."""
repos: list[ExternalTestRepoConfig] = []
for item in raw.get("external_test_repos", []) or []:
repos.append(
ExternalTestRepoConfig(
name=item.get("name", ""),
url=item.get("url", ""),
branch=item.get("branch", "main"),
test_cases=item.get("test_cases", []),
install_cmd=item.get("install_cmd", ""),
)
)

return ExternalTestConfig(
enabled=bool(raw.get("enabled", False)),
repos=repos,
test_procs=int(raw.get("test_procs", 8)),
mode=str(raw.get("mode", "inline")),
max_retries=int(raw.get("max_retries", 5)),
timeout=int(raw.get("timeout", 7200)),
)


def _apply_env_overrides(cfg: ExternalTestConfig) -> None:
"""Apply environment variable overrides on top of YAML config.

Environment variables always take precedence over the YAML file.
"""
# Master switch
env_enabled = os.getenv("TA_EXTERNAL_TEST_ENABLED", "").lower()
if env_enabled in ("true", "1", "yes"):
cfg.enabled = True
elif env_enabled in ("false", "0", "no"):
cfg.enabled = False

# Mode
env_mode = os.getenv("TA_EXTERNAL_TEST_MODE", "").lower()
if env_mode in ("inline", "standalone", "off"):
cfg.mode = env_mode

# Parallelism
try:
cfg.test_procs = int(os.getenv("TA_EXTERNAL_TEST_PROCS", str(cfg.test_procs)))
except ValueError:
pass

# Max retries
try:
cfg.max_retries = int(
os.getenv("TA_EXTERNAL_TEST_MAX_RETRIES", str(cfg.max_retries))
)
except ValueError:
pass

# Timeout
try:
cfg.timeout = int(os.getenv("TA_EXTERNAL_TEST_TIMEOUT", str(cfg.timeout)))
except ValueError:
pass
74 changes: 74 additions & 0 deletions src/TA_main2main_workflow/external_test/external_test_config.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
# ── 全局开关:是否启用外部测试(默认 false,需用户主动开启)──
enabled: false # 设为 true 启用外部测试

# ── 运行模式:inline(阻塞主流程)| standalone(不阻塞)| off(禁用)──
mode: inline # 仅在 enabled=true 时生效

# ── 全局默认值 ────────────────────────────────────────────────────
test_procs: 8 # 每个外部算子仓 pytest 并行进程数
max_retries: 5 # 每仓最大 AI 修复重试次数
timeout: 7200 # 单仓测试超时(秒)

# ── 外部算子仓列表(按配置顺序串行执行)───────────────────────────
external_test_repos:
- name: "Liger-Kernel"
url: "https://github.com/linkedin/Liger-Kernel.git"
branch: "main"
# install_cmd: "pip install -e ." # 可选,为空则跳过安装
test_cases:
- "test/transformers/test_attn_res.py"
- "test/transformers/test_auto_model.py"
- "test/transformers/test_cute_moe_autograd.py"
- "test/transformers/test_cutedsl_rms_norm.py"
- "test/transformers/test_cutedsl_rms_norm_fastpath.py"
- "test/transformers/test_cutedsl_rope.py"
- "test/transformers/test_cutile_backend.py"
- "test/transformers/test_dyt.py"
- "test/transformers/test_embedding.py"
- "test/transformers/test_flex_attention.py"
- "test/transformers/test_fused_add_rms_norm.py"
- "test/transformers/test_fused_linear_cross_entropy.py"
- "test/transformers/test_fused_linear_scaled_cross_entropy.py"
- "test/transformers/test_group_norm.py"
- "test/transformers/test_kl_div.py"
- "test/transformers/test_layer_norm.py"
- "test/transformers/test_llama4_rope.py"
- "test/transformers/test_mhc.py"
- "test/transformers/test_mm_int8int2.py"
- "test/transformers/test_modulated_rms_norm.py"
- "test/transformers/test_moe.py"
- "test/transformers/test_monkey_patch.py"
- "test/transformers/test_swiglu_cutedsl.py"

- name: "flash-linear-attention"
url: "https://github.com/fla-org/flash-linear-attention.git"
branch: "main"
test_cases:
- "tests/context_parallel/test_cp_conv.py"
- "tests/context_parallel/test_cp_dplr.py"
- "tests/context_parallel/test_cp_gdn.py"
- "tests/context_parallel/test_cp_kda.py"
- "tests/context_parallel/test_cp_rwkv7.py"
- "tests/context_parallel/test_cp_token_shift.py"
- "tests/layers/test_layer_cache_layer_idx.py"
- "tests/models/test_cache.py"
- "tests/models/test_generation_utils.py"
- "tests/models/test_hybrid_attention.py"
- "tests/models/test_modeling_bitnet.py"
- "tests/models/test_modeling_deltaformer.py"
- "tests/models/test_modeling_mla.py"
- "tests/models/test_modeling_moba.py"
- "tests/models/test_modeling_mom.py"
- "tests/models/test_modeling_nsa.py"
- "tests/models/test_modeling_rodimus.py"
- "tests/models/test_modeling_samba.py"
- "tests/models/test_modeling_transformer.py"
- "tests/modules/test_grpo.py"
- "tests/modules/test_l2norm.py"
- "tests/ops/test_cache.py"
- "tests/ops/test_forgetting_attn.py"
- "tests/ops/test_moba.py"
- "tests/ops/test_titans.py"
- "tests/test_public_api.py"
- "tests/test_split_package_release.py"
- "tests/utils/test_ascend_ub_manager.py"
Loading