| |
| """Run independent, resumable stages for one registered model pair.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import os |
| import subprocess |
| import sys |
| import traceback |
| from copy import deepcopy |
| from pathlib import Path |
| from typing import Any |
|
|
| from pipeline_common import ( |
| REPO_ROOT, |
| atomic_write_json, |
| basic_validate_config, |
| canonical_json_sha256, |
| expand, |
| file_record, |
| load_json, |
| resolve_path, |
| run_id, |
| safe_slug, |
| shell_join, |
| tool_versions, |
| utc_now, |
| ) |
|
|
|
|
| def parse_args(argv: list[str] | None = None) -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--config", required=True, type=Path) |
| parser.add_argument("--stage", action="append", dest="stages", help="stage id to run; repeatable") |
| parser.add_argument("--dry-run", action="store_true", help="record resolved commands without executing them") |
| parser.add_argument("--force", action="store_true", help="rerun a matching successful stage") |
| parser.add_argument( |
| "--resume-failed", |
| action="store_true", |
| help="run failed/blocked/missing stages and reuse matching successful stages", |
| ) |
| parser.add_argument("--result-out", type=Path, help="override aggregate result path") |
| return parser.parse_args(argv) |
|
|
|
|
| def load_and_validate(config_path: Path) -> tuple[dict[str, Any], str]: |
| config = load_json(config_path) |
| errors = basic_validate_config(config) |
| try: |
| import jsonschema |
|
|
| schema = load_json(REPO_ROOT / "schemas" / "model_config.schema.json") |
| validator = jsonschema.Draft202012Validator(schema) |
| errors.extend(error.message for error in validator.iter_errors(config)) |
| except ImportError: |
| pass |
| if errors: |
| raise ValueError("invalid model configuration:\n- " + "\n- ".join(sorted(set(errors)))) |
| return config, canonical_json_sha256(config) |
|
|
|
|
| def read_previous(stage_root: Path) -> dict[str, Any] | None: |
| latest = stage_root / "latest_result.json" |
| if not latest.is_file(): |
| return None |
| try: |
| value = load_json(latest) |
| return value if isinstance(value, dict) else None |
| except (OSError, json.JSONDecodeError): |
| return None |
|
|
|
|
| def same_inputs(previous: dict[str, Any], current: list[dict[str, Any]]) -> bool: |
| def signature(records: list[dict[str, Any]]) -> list[tuple[str, bool, str | None, int | None]]: |
| return [(r["path"], r["exists"], r["sha256"], r["bytes"]) for r in records] |
|
|
| return signature(previous.get("inputs", [])) == signature(current) |
|
|
|
|
| def stage_config_sha256(config: dict[str, Any], stage: dict[str, Any]) -> str: |
| variant = stage["variant"] |
| artifact = config["artifacts"].get(variant) if variant in {"fp32", "public_quantized"} else None |
| return canonical_json_sha256({ |
| "schema_version": config["schema_version"], |
| "model_id": config["model"]["model_id"], |
| "random_seed": config["model"].get("random_seed", 0), |
| "artifact": artifact, |
| "stage": stage, |
| }) |
|
|
|
|
| def previous_stage_matches( |
| previous: dict[str, Any] | None, stage_sha: str, argv: list[str], stage: dict[str, Any] |
| ) -> bool: |
| if previous is None: |
| return False |
| if previous.get("stage_config_sha256") == stage_sha: |
| return True |
| |
| return ( |
| "stage_config_sha256" not in previous |
| and previous.get("stage_id") == stage["id"] |
| and previous.get("stage") == stage["stage"] |
| and previous.get("variant") == stage["variant"] |
| and previous.get("command_argv") == argv |
| and previous.get("options", {}) == stage.get("options", {}) |
| and previous.get("patch") == stage.get("patch") |
| ) |
|
|
|
|
| def summarize_error(stderr_path: Path, exception_text: str = "") -> str: |
| text = exception_text |
| if stderr_path.is_file(): |
| try: |
| lines = stderr_path.read_text(encoding="utf-8", errors="replace").splitlines() |
| if lines: |
| text = "\n".join(lines[-20:]) |
| except OSError: |
| pass |
| return text[-4000:] |
|
|
|
|
| def load_validation_report(path: Path | None) -> dict[str, Any]: |
| if path is None: |
| return {} |
| if not path.is_file(): |
| return {"validation_report_exists": False, "validation_report_path": str(path)} |
| try: |
| value = load_json(path) |
| except (OSError, json.JSONDecodeError) as error: |
| return { |
| "validation_report_exists": True, |
| "validation_report_path": str(path), |
| "validation_report_parse_error": str(error), |
| } |
| return { |
| "validation_report_exists": True, |
| "validation_report_path": str(path), |
| "stage_report": value, |
| } |
|
|
|
|
| def make_result_base( |
| config: dict[str, Any], stage: dict[str, Any], stage_id: str, config_sha: str, |
| stage_sha: str, |
| argv: list[str], cwd: Path, inputs: list[dict[str, Any]], stdout_path: Path, |
| stderr_path: Path, |
| ) -> dict[str, Any]: |
| variant = stage["variant"] |
| artifact = config["artifacts"].get(variant, {}) if variant in {"fp32", "public_quantized"} else {} |
| return { |
| "model_id": config["model"]["model_id"], |
| "artifact_id": stage.get("artifact_id", artifact.get("artifact_id", f"{config['model']['model_id']}-{variant}")), |
| "variant": variant, |
| "stage": stage["stage"], |
| "stage_id": stage_id, |
| "status": "RUNNING", |
| "failure_code": None, |
| "source_artifact": artifact.get("local_path"), |
| "source_checksum": artifact.get("sha256"), |
| "artifact": None, |
| "command": shell_join(argv), |
| "command_argv": argv, |
| "working_directory": str(cwd), |
| "inputs": inputs, |
| "outputs": [], |
| "tool_versions": tool_versions(), |
| "options": stage.get("options", {}), |
| "random_seed": config["model"].get("random_seed", 0), |
| "started_at": utc_now(), |
| "ended_at": utc_now(), |
| "duration_sec": 0.0, |
| "exit_code": None, |
| "error_summary": "", |
| "stdout_log": str(stdout_path), |
| "stderr_log": str(stderr_path), |
| "patch": stage.get("patch"), |
| "validation": {}, |
| "config_sha256": config_sha, |
| "stage_config_sha256": stage_sha, |
| "reused_from_previous_run": False, |
| } |
|
|
|
|
| def run_stage( |
| config: dict[str, Any], stage: dict[str, Any], variables: dict[str, str], |
| config_sha: str, current_run_id: str, dry_run: bool, force: bool, |
| selected_ids: set[str] | None, statuses: dict[str, str], |
| ) -> dict[str, Any]: |
| stage_id = stage["id"] |
| stage_root = Path(variables["model_dir"]) / "logs" / stage_id |
| invocation_dir = stage_root / current_run_id |
| invocation_dir.mkdir(parents=True, exist_ok=True) |
| stdout_path = invocation_dir / "stdout.log" |
| stderr_path = invocation_dir / "stderr.log" |
| cwd = resolve_path(stage.get("working_directory", "{repo_root}"), variables) |
| argv = [expand(str(part), variables) for part in stage["command"]] |
| inputs = [file_record(resolve_path(value, variables, cwd)) for value in stage.get("inputs", [])] |
| outputs_paths = [resolve_path(value, variables, cwd) for value in stage.get("outputs", [])] |
| previous = read_previous(stage_root) |
| stage_sha = stage_config_sha256(config, stage) |
| result = make_result_base(config, stage, stage_id, config_sha, stage_sha, argv, cwd, inputs, stdout_path, stderr_path) |
|
|
| if selected_ids is not None and stage_id not in selected_ids: |
| if previous is not None: |
| reused = deepcopy(previous) |
| reused["reused_result_config_sha256"] = previous.get("config_sha256") |
| reused["config_sha256"] = config_sha |
| reused["stage_config_sha256"] = stage_sha |
| reused["reused_from_previous_run"] = True |
| atomic_write_json(invocation_dir / "run_result.json", reused) |
| atomic_write_json(stage_root / "latest_result.json", reused) |
| return reused |
| result.update(status="SKIPPED", ended_at=utc_now(), error_summary="stage not selected") |
| return result |
|
|
| dependencies = stage.get("requires", []) |
| dependency_success = {"PASS", "PASS_WITH_PATCH", "QUEUED"} if dry_run else {"PASS", "PASS_WITH_PATCH"} |
| unavailable = [dependency for dependency in dependencies if statuses.get(dependency) not in dependency_success] |
| if unavailable: |
| result.update( |
| status="BLOCKED", |
| failure_code=stage["failure_code_on_error"], |
| ended_at=utc_now(), |
| error_summary=f"required stages not successful: {', '.join(unavailable)}", |
| validation={"dependencies_satisfied": False, "unavailable_dependencies": unavailable}, |
| ) |
| atomic_write_json(invocation_dir / "run_result.json", result) |
| atomic_write_json(stage_root / "latest_result.json", result) |
| return result |
|
|
| if dry_run: |
| result.update( |
| status="QUEUED", |
| ended_at=utc_now(), |
| validation={"dry_run": True, "resolved_outputs": [str(path) for path in outputs_paths]}, |
| ) |
| atomic_write_json(invocation_dir / "run_result.json", result) |
| return result |
|
|
| if ( |
| not force and previous is not None |
| and previous.get("status") in {"PASS", "PASS_WITH_PATCH"} |
| and previous_stage_matches(previous, stage_sha, argv, stage) |
| and same_inputs(previous, inputs) |
| and all(record["exists"] for record in previous.get("outputs", [])) |
| ): |
| reused = deepcopy(previous) |
| reused["reused_result_config_sha256"] = previous.get("config_sha256") |
| reused["config_sha256"] = config_sha |
| reused["stage_config_sha256"] = stage_sha |
| reused["reused_from_previous_run"] = True |
| atomic_write_json(invocation_dir / "run_result.json", reused) |
| atomic_write_json(stage_root / "latest_result.json", reused) |
| return reused |
|
|
| conflicting_outputs = [path for path in outputs_paths if path.exists()] |
| previous_matches = previous_stage_matches(previous, stage_sha, argv, stage) |
| if conflicting_outputs and not previous_matches: |
| result.update( |
| status="FAIL", |
| failure_code="FAIL_ENVIRONMENT", |
| ended_at=utc_now(), |
| error_summary="refusing to overwrite output from a different or unknown configuration", |
| outputs=[file_record(path) for path in outputs_paths], |
| validation={"overwrite_guard": False, "conflicting_outputs": [str(path) for path in conflicting_outputs]}, |
| ) |
| atomic_write_json(invocation_dir / "run_result.json", result) |
| atomic_write_json(stage_root / "latest_result.json", result) |
| return result |
|
|
| import time |
|
|
| started_monotonic = time.monotonic() |
| exit_code: int | None = None |
| exception_text = "" |
| cwd.mkdir(parents=True, exist_ok=True) |
| for path in outputs_paths: |
| path.parent.mkdir(parents=True, exist_ok=True) |
| try: |
| with stdout_path.open("wb") as stdout_handle, stderr_path.open("wb") as stderr_handle: |
| completed = subprocess.run( |
| argv, |
| cwd=cwd, |
| stdout=stdout_handle, |
| stderr=stderr_handle, |
| timeout=stage["timeout_sec"], |
| check=False, |
| env={**os.environ, "PYTHONHASHSEED": str(config["model"].get("random_seed", 0))}, |
| ) |
| exit_code = completed.returncode |
| except subprocess.TimeoutExpired as error: |
| exit_code = 124 |
| exception_text = f"timeout after {stage['timeout_sec']} seconds: {error}" |
| except (OSError, ValueError) as error: |
| exit_code = 127 |
| exception_text = f"execution failed: {error}" |
| duration = time.monotonic() - started_monotonic |
| output_records = [file_record(path) for path in outputs_paths] |
| outputs_exist = all(record["exists"] for record in output_records) |
| validation_path = resolve_path(stage["validation_report"], variables, cwd) if stage.get("validation_report") else None |
| validation = { |
| "outputs_exist": outputs_exist, |
| "expected_output_count": len(output_records), |
| "quantization_preserved": None, |
| "runtime_output_match": None, |
| **load_validation_report(validation_path), |
| } |
| passed = exit_code == 0 and outputs_exist |
| result.update( |
| status="PASS_WITH_PATCH" if passed and stage.get("patch") else ("PASS" if passed else "FAIL"), |
| failure_code=None if passed else stage["failure_code_on_error"], |
| artifact=str(outputs_paths[0]) if passed and outputs_paths else None, |
| outputs=output_records, |
| ended_at=utc_now(), |
| duration_sec=round(duration, 6), |
| exit_code=exit_code, |
| error_summary="" if passed else summarize_error(stderr_path, exception_text or "stage failed or output missing"), |
| validation=validation, |
| ) |
| atomic_write_json(invocation_dir / "run_result.json", result) |
| atomic_write_json(stage_root / "latest_result.json", result) |
| return result |
|
|
|
|
| def overall_status(results: list[dict[str, Any]]) -> str: |
| relevant = [result["status"] for result in results if result["status"] != "SKIPPED"] |
| if relevant and all(status in {"PASS", "PASS_WITH_PATCH"} for status in relevant): |
| return "PASS" |
| if relevant and all(status == "QUEUED" for status in relevant): |
| return "QUEUED" |
| if any(status in {"PASS", "PASS_WITH_PATCH"} for status in relevant): |
| return "PARTIAL" |
| if any(status == "BLOCKED" for status in relevant) and not any(status == "FAIL" for status in relevant): |
| return "BLOCKED" |
| return "FAIL" |
|
|
|
|
| def execute(args: argparse.Namespace) -> tuple[dict[str, Any], int]: |
| config_path = args.config.resolve() |
| config, config_sha = load_and_validate(config_path) |
| model = config["model"] |
| model_dir = REPO_ROOT / "models" / safe_slug(model["task"]) / model["model_id"] |
| for relative in ( |
| "config", "source/fp32", "source/quantized", "baseline", "onnx/fp32", |
| "onnx/quantized", "tflite/fp32", "tflite/quantized", "mlir/fp32", |
| "mlir/quantized", "graphs", "analysis", "logs", |
| ): |
| (model_dir / relative).mkdir(parents=True, exist_ok=True) |
| variables = { |
| "repo_root": str(REPO_ROOT), |
| "model_dir": str(model_dir), |
| "python": sys.executable, |
| "config_dir": str(config_path.parent), |
| } |
| current_run_id = run_id() |
| started_at = utc_now() |
| selected = set(args.stages) if args.stages else None |
| known_ids = {stage["id"] for stage in config["stages"]} |
| if selected and (unknown := selected - known_ids): |
| raise ValueError(f"unknown stage ids: {', '.join(sorted(unknown))}") |
| statuses: dict[str, str] = {} |
| results: list[dict[str, Any]] = [] |
| for stage in config["stages"]: |
| result = run_stage( |
| config, stage, variables, config_sha, current_run_id, args.dry_run, |
| args.force, selected, statuses, |
| ) |
| statuses[stage["id"]] = result["status"] |
| results.append(result) |
| aggregate = { |
| "schema_version": "1.0", |
| "model_id": model["model_id"], |
| "config_path": str(config_path), |
| "config_sha256": config_sha, |
| "run_id": current_run_id, |
| "started_at": started_at, |
| "ended_at": utc_now(), |
| "overall_status": overall_status(results), |
| "stages": results, |
| } |
| output_path = args.result_out.resolve() if args.result_out else model_dir / ("dry_run_result.json" if args.dry_run else "run_result.json") |
| atomic_write_json(output_path, aggregate) |
| return aggregate, 0 if aggregate["overall_status"] in {"PASS", "QUEUED"} else 1 |
|
|
|
|
| def main(argv: list[str] | None = None) -> int: |
| args = parse_args(argv) |
| try: |
| aggregate, exit_code = execute(args) |
| print(json.dumps({ |
| "model_id": aggregate["model_id"], |
| "run_id": aggregate["run_id"], |
| "overall_status": aggregate["overall_status"], |
| }, sort_keys=True)) |
| return exit_code |
| except Exception as error: |
| traceback.print_exc(file=sys.stderr) |
| print(f"run_model failed: {error}", file=sys.stderr) |
| return 2 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|