From 05864d12cb2c9388b8bc6a8d868091edee87bfed Mon Sep 17 00:00:00 2001 From: MichaelChung Date: Thu, 24 Sep 2026 14:20:09 +0000 Subject: [PATCH] rl: add elastic resource benchmark contract, gating, evidence and paired work Implements the GPU-free half of the rl-elastic-resource-benchmark change: versioned study manifest with canonical hash and calibration/formal modes, config/edge/pool validation gated by runtime capability attestation, evidence index with tamper-rejecting resume, dry-run plan CLI, layered results and report, and paired splits/schedules/scenarios that reuse the legacy benchmark_rl pairing without changing its CLI. Co-Authored-By: Claude Fable 5.1 --- docs/RL_ELASTIC_BENCHMARK.md | 119 ++++++ scripts/benchmark_rl_elastic.py | 18 + tests/test_rl_elastic_benchmark.py | 349 +++++++++++++++++ yeto/rl/elastic_benchmark/__init__.py | 7 + yeto/rl/elastic_benchmark/capabilities.py | 306 +++++++++++++++ yeto/rl/elastic_benchmark/cli.py | 82 ++++ yeto/rl/elastic_benchmark/evidence.py | 186 +++++++++ yeto/rl/elastic_benchmark/legacy.py | 90 +++++ yeto/rl/elastic_benchmark/manifest.py | 456 ++++++++++++++++++++++ yeto/rl/elastic_benchmark/plan.py | 133 +++++++ yeto/rl/elastic_benchmark/results.py | 147 +++++++ yeto/rl/elastic_benchmark/work.py | 246 ++++++++++++ 12 files changed, 2139 insertions(+) create mode 100644 docs/RL_ELASTIC_BENCHMARK.md create mode 100755 scripts/benchmark_rl_elastic.py create mode 100644 tests/test_rl_elastic_benchmark.py create mode 100644 yeto/rl/elastic_benchmark/__init__.py create mode 100644 yeto/rl/elastic_benchmark/capabilities.py create mode 100644 yeto/rl/elastic_benchmark/cli.py create mode 100644 yeto/rl/elastic_benchmark/evidence.py create mode 100644 yeto/rl/elastic_benchmark/legacy.py create mode 100644 yeto/rl/elastic_benchmark/manifest.py create mode 100644 yeto/rl/elastic_benchmark/plan.py create mode 100644 yeto/rl/elastic_benchmark/results.py create mode 100644 yeto/rl/elastic_benchmark/work.py diff --git a/docs/RL_ELASTIC_BENCHMARK.md b/docs/RL_ELASTIC_BENCHMARK.md new file mode 100644 index 00000000..80415213 --- /dev/null +++ b/docs/RL_ELASTIC_BENCHMARK.md @@ -0,0 +1,119 @@ +# RL elastic resource benchmark + +`scripts/benchmark_rl_elastic.py` plans and reports studies that compare fixed +GPU partitions (trainer/rollout/standby) against scheduled and automatic +resizing inside one RL island. It is separate from `scripts/benchmark_rl.py`; +that harness, its CLI and its result fields are unchanged. The suite reuses +its prompt pairing and file format through `yeto/rl/elastic_benchmark/legacy.py`. + +The planning contract lives in the miles repository under +`openspec/changes/rl-elastic-resource-benchmark/`. This document covers what is +implemented today: the study contract, capability gating, evidence, paired work +and reporting. No GPU runner ships yet; matrix items are executed by runners +that register against a runtime capability attestation. + +## Commands + +```bash +python scripts/benchmark_rl_elastic.py example --mode calibration --output study.json +python scripts/benchmark_rl_elastic.py validate --study study.json +python scripts/benchmark_rl_elastic.py plan --study study.json --capabilities caps.json +python scripts/benchmark_rl_elastic.py report --study study.json --capabilities caps.json --study-dir out/ +``` + +None of these load a model, import torch/ray/miles, or create cloud resources. +`plan` prints the config table (legal configs with gradient accumulation, +illegal ones with the rejection reason), every arm × config × scenario × seed +item with its status, and the runnable budget. `report` exits non-zero while +the matrix is incomplete. + +## Study manifest + +A manifest is JSON with `schema_version`, `mode` (`calibration` or `formal`) +and seven field groups: `identity`, `profile`, `resources`, `matrix`, `work`, +`evaluation`, `timing`. `dump_manifest` stamps `study_hash`, a sha256 over the +canonical JSON; `load_manifest` refuses a file whose content no longer matches. + +Calibration manifests may leave fields as the string `"unresolved"` and may +have an empty GPU list and no quality rules. Formal manifests must have +immutable model and data revisions, reward/source/runtime fingerprints, +physical GPU UUIDs, per-metric `delta_quality` and `catastrophic` bounds, +`delta_stable`, `min_speedup` and a statistics method. Contradictions are +rejected instead of repaired: phase updates must sum to `update_budget`, +`groups_per_update × samples_per_group` must equal `global_batch`, warmup plus +measured updates cannot exceed the budget. + +## Configs, edges and capability gating + +Configs are named partitions such as `P62` (6 trainer, 2 rollout). The first +round certifies TP=PP=CP=EP=1, so trainer count is the DP size. A config is +illegal when it has no trainer or rollout GPU, does not match the pool size, +uses DP>1 with a dense full-parameter profile, or `global_batch` is not +divisible by DP × micro batch. Illegal configs stay in the table with their +reason; they are never swapped for a neighbour. + +Edges are directed and typed: `rollout-only`, `same-shape-restore`, +`trainer-dp`, `role-transfer`, `standby-scale`. Shape checks reject, for +example, a `rollout-only` edge whose trainer count changes. + +A capability attestation (`caps.json`) says what the runtime has certified: + +```json +{ + "runtime_fingerprint": "sha256:...", + "execution_modes": ["partitioned-serial"], + "partitioned_driver": true, + "certified_edges": [{"source": "P422", "target": "P44", "kind": "rollout-only"}], + "optimized_paths": [], + "auto_controller": false +} +``` + +Each matrix item is `supported`, `unsupported` (illegal config, unknown edge) +or `blocked_dependency` (declared but uncertified edge, unattested execution +mode, missing optimized path or controller). Without an attestation only the +`legacy-fixed` arm is runnable. + +## Evidence and resume + +Every attempt writes to +`runs/[@config]//seed-/attempt-/` and ends with +`evidence-index.json` (sha256 and size of every raw file) and `result.json`. +Reuse on resume requires the same `study_hash`, the same matrix key and every +digest to match; a modified, missing or unindexed file is an error rather than +a silent rerun. Results are never overwritten; a new attempt gets a new number +and failed attempts stay on disk. + +`result.json` has four layers: `execution` (completed, failed, unsupported, +blocked_dependency, pending), `correctness`, `quality` and `benefit`. Only +completed attempts may carry verdicts. The study summary is `incomplete` while +any supported item lacks a verified completion, and the study-level benefit is +`demonstrated` only when the matrix is complete and every dynamic arm passes +correctness, quality and benefit. Negative results are kept as `negative`. + +## Paired work and scenarios + +`work.split_rows` draws disjoint train/calibration/test/held-out row sets from +a seeded permutation. `legacy.materialize_paired_inputs` writes the training +stream in the legacy `combined.jsonl` format and refuses any held-out row in +the training assignment. `work.group_seeds` gives every (group, sample) a +stable logical seed so arms can reconcile their inputs. + +Phases bind to logical update indexes. `work.phase_schedule` expands them to +one slot per update, and `work.assign_groups` draws prompt ids per update from +calibration-measured length buckets, so a slow and a fast arm see the same +phase sequence. Scenario recipes (`stable`, `phased`, `tail`, `tool-wait`, +`oscillating`) only shape the bucket mix and the environment contract; they +take measured mixes, the measured payback window and the frozen environment +from a calibration record and refuse to run without them. They never change +global batch, group size, optimizer steps or truncation rules. +`work.length_diagnostics` reports real lengths, cap-hit ratio and +zero-advantage ratio from captured samples. + +## What is not implemented + +Runners for the fixed partition, scheduled switching, auto control, the +switching timeline and cost aggregation, the amortization sweep and the +statistical acceptance are pending on the upstream reconfiguration change and +on real calibration runs. See the tasks file in the miles repository for +status. diff --git a/scripts/benchmark_rl_elastic.py b/scripts/benchmark_rl_elastic.py new file mode 100755 index 00000000..7b2d9141 --- /dev/null +++ b/scripts/benchmark_rl_elastic.py @@ -0,0 +1,18 @@ +#!/usr/bin/env python3 +"""Thin entry for the RL elastic resource benchmark suite. + +See docs/RL_ELASTIC_BENCHMARK.md. Logic lives in yeto/rl/elastic_benchmark/. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parent.parent +sys.path.insert(0, str(REPO_ROOT)) + +from yeto.rl.elastic_benchmark.cli import main # noqa: E402 + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/test_rl_elastic_benchmark.py b/tests/test_rl_elastic_benchmark.py new file mode 100644 index 00000000..3cf29f5e --- /dev/null +++ b/tests/test_rl_elastic_benchmark.py @@ -0,0 +1,349 @@ +"""Tests for the RL elastic resource benchmark suite (no GPU, no runtime).""" + +from __future__ import annotations + +import copy +import json +import subprocess +import sys +from pathlib import Path + +import pytest + +from yeto.rl.elastic_benchmark import capabilities as caps +from yeto.rl.elastic_benchmark import evidence, legacy, results, work +from yeto.rl.elastic_benchmark.cli import main +from yeto.rl.elastic_benchmark.manifest import ( + ManifestError, + dump_manifest, + example_manifest, + load_manifest, + manifest_hash, + validate_manifest, +) +from yeto.rl.elastic_benchmark.plan import build_plan + +REPO_ROOT = Path(__file__).resolve().parents[1] +SCRIPT = REPO_ROOT / "scripts" / "benchmark_rl_elastic.py" + + +def _attested(**overrides) -> caps.Attestation: + payload = { + "runtime_fingerprint": "sha256:runtime", + "execution_modes": ["partitioned-serial"], + "partitioned_driver": True, + "certified_edges": [{"source": "P422", "target": "P44", "kind": "rollout-only"}], + } + payload.update(overrides) + return caps.attestation_from_dict(payload) + + +# --- 1.1 manifest --------------------------------------------------------- + + +def test_calibration_manifest_round_trips_with_stable_hash(tmp_path): + manifest = example_manifest() + digest = dump_manifest(manifest, tmp_path / "study.json") + loaded = load_manifest(tmp_path / "study.json") + assert loaded["study_hash"] == digest == manifest_hash(loaded) + reordered = json.loads(json.dumps(manifest, sort_keys=False)) + assert manifest_hash(reordered) == digest + + +def test_formal_manifest_rejects_mutable_revision_and_missing_quality_rules(): + formal = example_manifest("formal") + validate_manifest(formal) + mutable = copy.deepcopy(formal) + mutable["identity"]["model"]["revision"] = "main" + with pytest.raises(ManifestError, match="immutable identity.model.revision"): + validate_manifest(mutable) + no_rules = copy.deepcopy(formal) + no_rules["evaluation"]["metrics"] = [] + with pytest.raises(ManifestError, match="predeclared bounds"): + validate_manifest(no_rules) + unresolved = copy.deepcopy(formal) + unresolved["resources"]["pool_id"] = "unresolved" + with pytest.raises(ManifestError, match="unresolved fields: resources.pool_id"): + validate_manifest(unresolved) + + +def test_contradictory_budget_is_rejected(): + manifest = example_manifest() + manifest["work"]["update_budget"] = 11 + with pytest.raises(ManifestError, match="contradicts the phase sum"): + validate_manifest(manifest) + manifest = example_manifest() + manifest["timing"]["measured_updates"] = 99 + with pytest.raises(ManifestError, match="exceed work.update_budget"): + validate_manifest(manifest) + manifest = example_manifest() + manifest["work"]["samples_per_group"] = 5 + with pytest.raises(ManifestError, match="must equal profile.global_batch"): + validate_manifest(manifest) + + +def test_edited_manifest_is_rejected_by_recorded_hash(tmp_path): + path = tmp_path / "study.json" + dump_manifest(example_manifest(), path) + payload = json.loads(path.read_text()) + payload["matrix"]["seeds"] = [1] + path.write_text(json.dumps(payload)) + with pytest.raises(ManifestError, match="recorded study_hash"): + load_manifest(path) + + +# --- 1.2 configs, edges, pool, attestation ---------------------------------- + + +def test_config_rejections_keep_reasons_instead_of_substituting(): + manifest = example_manifest() + manifest["resources"]["configs"]["P71"] = {"trainer": 7, "rollout": 1} + manifest["resources"]["configs"]["P80"] = {"trainer": 8, "rollout": 0} + configs = caps.parse_configs(manifest["resources"]) + table = {row["config"]: row for row in caps.config_table(configs, profile=manifest["profile"], pool_size=8)} + assert table["P62"]["legal"] and table["P62"]["gradient_accumulation"] == 8 + assert "not divisible" in table["P71"]["reason"] + assert "rollout" in table["P80"]["reason"] + dense = dict(manifest["profile"], parameter_mode="full") + assert "DP=1" in caps.config_rejection(configs["P62"], profile=dense, pool_size=8) + assert caps.config_rejection(configs["P62"], profile=manifest["profile"], pool_size=4).startswith("uses 8") + + +def test_duplicate_gpu_uuid_and_wrong_edge_shape_are_rejected(): + resources = example_manifest("formal")["resources"] + resources["gpus"].append(dict(resources["gpus"][0])) + with pytest.raises(ManifestError, match="duplicate GPU uuid"): + caps.validate_pool(resources) + resources = example_manifest()["resources"] + resources["edges"].append({"source": "P62", "target": "P26", "kind": "rollout-only"}) + with pytest.raises(ManifestError, match="changes trainer count"): + caps.validate_edges(resources, caps.parse_configs(resources)) + + +def test_uncertified_edges_block_without_downgrade(): + manifest = example_manifest() + manifest["matrix"]["arms"].append( + {"name": "b1", "kind": "scheduled-rebuild", "config": "P422", "switch_plan": [{"at_update": 3, "target": "P44"}]} + ) + plan = build_plan(manifest, _attested(), study_hash="h") + status = {(i.key.arm, i.key.config): (i.status, i.reason) for i in plan.items} + assert status[("b1", None)] == ("supported", None) + assert status[("rebuild", None)][0] == "blocked_dependency" + assert "role-transfer" in status[("rebuild", None)][1] + assert status[("auto", None)][0] == "blocked_dependency" + assert status[("legacy", None)] == ("supported", None) + assert status[("sweep", "P44")] == ("supported", None) + none = build_plan(manifest, caps.Attestation.none(), study_hash="h") + assert all(i.status == "blocked_dependency" for i in none.items if i.key.arm != "legacy") + + +# --- 1.3 evidence and resume -------------------------------------------------- + + +def _completed_attempt(study_dir: Path, key: evidence.MatrixKey, study_hash: str, attempt: int = 1) -> Path: + attempt_dir = study_dir / key.relative_dir(attempt) + (attempt_dir / "rollout").mkdir(parents=True) + (attempt_dir / "rollout" / "1.jsonl").write_text('{"reward":1}\n') + (attempt_dir / "metrics.jsonl").write_text('{"step":1}\n') + evidence.build_evidence_index(attempt_dir, study_hash=study_hash, key=key, attempt=attempt) + evidence.write_result(attempt_dir, results.make_result(execution="completed", correctness="passed")) + return attempt_dir + + +def test_tampered_or_missing_evidence_cannot_be_reused(tmp_path): + key = evidence.MatrixKey("default", "stable", 17) + attempt_dir = _completed_attempt(tmp_path, key, "h1") + assert evidence.reusable_completion(tmp_path, study_hash="h1", key=key)["execution"] == "completed" + (attempt_dir / "metrics.jsonl").write_text('{"step":2}\n') + with pytest.raises(evidence.EvidenceError, match="modified: metrics.jsonl"): + evidence.reusable_completion(tmp_path, study_hash="h1", key=key) + (attempt_dir / "metrics.jsonl").unlink() + with pytest.raises(evidence.EvidenceError, match="missing: metrics.jsonl"): + evidence.reusable_completion(tmp_path, study_hash="h1", key=key) + + +def test_changed_study_identity_and_unindexed_files_are_rejected(tmp_path): + key = evidence.MatrixKey("default", "stable", 17) + attempt_dir = _completed_attempt(tmp_path, key, "h1") + with pytest.raises(evidence.EvidenceError, match="different study"): + evidence.reusable_completion(tmp_path, study_hash="h2", key=key) + (attempt_dir / "extra.log").write_text("late") + with pytest.raises(evidence.EvidenceError, match="unindexed"): + evidence.reusable_completion(tmp_path, study_hash="h1", key=key) + + +def test_failed_attempts_are_kept_and_results_never_overwritten(tmp_path): + key = evidence.MatrixKey("default", "stable", 17) + failed_dir = tmp_path / key.relative_dir(1) + failed_dir.mkdir(parents=True) + evidence.write_result(failed_dir, results.make_result(execution="failed", reason="oom")) + with pytest.raises(evidence.EvidenceError, match="refusing to overwrite"): + evidence.write_result(failed_dir, results.make_result(execution="completed")) + assert evidence.next_attempt(tmp_path, key) == 2 + assert evidence.reusable_completion(tmp_path, study_hash="h1", key=key) is None + _completed_attempt(tmp_path, key, "h1", attempt=2) + collected = evidence.collect_results(tmp_path, [key])[key] + assert [r["execution"] for r in collected] == ["failed", "completed"] + + +# --- 1.4 CLI / dry-run ---------------------------------------------------------- + + +def test_dry_run_plan_lists_budget_and_never_imports_runtimes(tmp_path): + study = tmp_path / "study.json" + assert main(["example", "--output", str(study)]) == 0 + caps_path = tmp_path / "caps.json" + caps_path.write_text(json.dumps({"runtime_fingerprint": "x", "execution_modes": ["partitioned-serial"], "partitioned_driver": True})) + proc = subprocess.run( + [sys.executable, str(SCRIPT), "plan", "--study", str(study), "--capabilities", str(caps_path), "--json"], + capture_output=True, text=True, check=False, cwd=tmp_path, + ) + assert proc.returncode == 0, proc.stderr + payload = json.loads(proc.stdout) + assert payload["counts"] == {"supported": 30, "unsupported": 0, "blocked_dependency": 18} + assert payload["budget"]["updates"] == 360 + assert any(i["status"] == "blocked_dependency" and i["arm"] == "auto" for i in payload["items"]) + forbidden = subprocess.run( + [sys.executable, "-c", "import sys;from yeto.rl.elastic_benchmark import cli,plan,results;" + "bad=[m for m in ('torch','ray','miles','sky','transformers') if m in sys.modules];print(bad)"], + capture_output=True, text=True, check=True, cwd=REPO_ROOT, + ) + assert forbidden.stdout.strip() == "[]" + assert not list(tmp_path.glob("runs")) + + +def test_cli_reports_manifest_errors_as_exit_2(tmp_path, capsys): + study = tmp_path / "study.json" + manifest = example_manifest() + manifest["matrix"]["scenarios"] = ["bogus"] + study.write_text(json.dumps(manifest)) + assert main(["validate", "--study", str(study)]) == 2 + assert "matrix.scenarios" in capsys.readouterr().err + + +# --- 1.5 layered results --------------------------------------------------------- + + +def test_summary_is_incomplete_until_every_required_item_has_verified_evidence(tmp_path): + manifest = example_manifest() + manifest["matrix"]["arms"] = [{"name": "default", "kind": "target-fixed-default", "config": "P44"}] + manifest["matrix"]["scenarios"] = ["stable"] + manifest["matrix"]["seeds"] = [17, 29] + plan = build_plan(manifest, _attested(), study_hash="h") + _completed_attempt(tmp_path, evidence.MatrixKey("default", "stable", 17), "h") + summary = results.summarize(plan, tmp_path) + assert summary["status"] == "incomplete" + assert summary["missing"] == [{"arm": "default", "scenario": "stable", "seed": 29, "config": None}] + assert summary["benefit"] == "not_demonstrated" + report = results.render_report(summary) + assert "no survivor mean" in report + _completed_attempt(tmp_path, evidence.MatrixKey("default", "stable", 29), "h") + complete = results.summarize(plan, tmp_path) + assert complete["status"] == "complete" and complete["benefit"] == "not_checked" + + +def test_benefit_needs_every_gate_and_negative_results_survive(tmp_path): + manifest = example_manifest() + manifest["matrix"]["arms"] = [ + {"name": "b1", "kind": "scheduled-rebuild", "config": "P422", "switch_plan": [{"at_update": 3, "target": "P44"}]} + ] + manifest["matrix"]["scenarios"], manifest["matrix"]["seeds"] = ["stable"], [17] + plan = build_plan(manifest, _attested(), study_hash="h") + key = evidence.MatrixKey("b1", "stable", 17) + attempt_dir = tmp_path / key.relative_dir(1) + attempt_dir.mkdir(parents=True) + evidence.build_evidence_index(attempt_dir, study_hash="h", key=key, attempt=1) + evidence.write_result(attempt_dir, results.make_result(execution="completed", correctness="passed", quality="insufficient", benefit="demonstrated")) + assert results.summarize(plan, tmp_path)["benefit"] == "not_demonstrated" + with pytest.raises(ValueError, match="only completed attempts"): + results.make_result(execution="failed", benefit="demonstrated") + + +# --- 2.1 legacy reuse ------------------------------------------------------------- + + +def test_legacy_pairing_is_reused_and_held_out_never_leaks(tmp_path): + rows = [{"prompt": f"q{i}", "label": str(i)} for i in range(10)] + streams = legacy.paired_streams(rows, islands=2, groups=2, rounds=1) + assert streams.combined_ids == (0, 1, 2, 3) + info = legacy.materialize_paired_inputs( + rows, train_ids=[0, 1, 2, 3], held_out_ids=[8, 9], assignment=[[0, 1], [2, 3]], directory=tmp_path + ) + combined = legacy.read_jsonl(Path(info["combined"])) + assert [r["metadata"]["benchmark_prompt_id"] for r in combined] == [0, 1, 2, 3] + assert [r["label"] for r in legacy.read_jsonl(Path(info["eval"]))] == ["8", "9"] + with pytest.raises(ValueError, match="held-out rows leaked"): + legacy.materialize_paired_inputs(rows, train_ids=[0], held_out_ids=[1], assignment=[[1]], directory=tmp_path / "x") + + +# --- 2.2 splits, seeds, schedules ----------------------------------------------------- + + +def test_splits_are_disjoint_deterministic_and_seeds_reconcile(): + sizes = {"train": 6, "calibration": 2, "test": 1, "held_out": 1} + a, b = work.split_rows(12, sizes, seed=17), work.split_rows(12, sizes, seed=17) + assert a == b + parts = [set(getattr(a, n)) for n in work.SPLIT_NAMES] + assert sum(len(p) for p in parts) == 10 and not (parts[0] & parts[3]) + assert work.split_rows(12, sizes, seed=29) != a + with pytest.raises(ManifestError, match="only 5 are available"): + work.split_rows(5, sizes, seed=1) + seeds = work.group_seeds(study_seed=17, groups=2, samples_per_group=2) + assert seeds == work.group_seeds(study_seed=17, groups=2, samples_per_group=2) + assert len({s for g in seeds for s in g}) == 4 + + +def test_phase_schedule_binds_to_logical_updates_not_wall_time(): + phases = example_manifest()["work"]["phases"] + slots = work.phase_schedule(phases) + assert [s.update for s in slots] == list(range(1, 13)) + assert [s.phase for s in slots[3:6]] == ["A", "B", "B"] + pools = {"short": [0, 1, 2], "long": [7, 8]} + fast = work.assign_groups(slots, pools=pools, groups_per_update=5, seed=17) + slow = work.assign_groups(slots, pools=pools, groups_per_update=5, seed=17) + assert fast == slow and all(len(u) == 5 for u in fast) + assert work.bucket_counts({"short": 0.8, "long": 0.2}, 5) == {"short": 4, "long": 1} + with pytest.raises(ManifestError, match="shares must sum"): + work.phase_schedule([{"name": "x", "updates": 1, "mix": {"short": 0.5}}]) + + +# --- 2.3 scenarios --------------------------------------------------------------------- + + +def test_scenarios_only_shape_the_mix_and_need_calibration_facts(): + calibration = { + "mixes": { + "stable": {"short": 0.5, "long": 0.5}, + "short-heavy": {"short": 0.9, "long": 0.1}, + "long-heavy": {"short": 0.1, "long": 0.9}, + "tail": {"short": 0.9, "long": 0.1}, + "tool-wait": {"short": 1.0}, + }, + "payback_updates": 6, + "environment": {"task_pack": "tp1", "tool_concurrency": 4, "response_contract": "v1"}, + } + for scenario in work.SPLIT_NAMES and ("stable", "phased", "tail", "tool-wait", "oscillating"): + phases = work.scenario_phases(scenario, updates=12, calibration=calibration) + assert sum(p["updates"] for p in phases) == 12 + assert all(set(p) <= {"name", "updates", "mix", "environment"} for p in phases) + oscillating = work.scenario_phases("oscillating", updates=12, calibration=calibration) + assert all(p["updates"] < 6 for p in oscillating) and len(oscillating) == 4 + with pytest.raises(ManifestError, match="payback_updates"): + work.scenario_phases("oscillating", updates=12, calibration={"mixes": calibration["mixes"]}) + with pytest.raises(ManifestError, match="calibration.environment"): + work.scenario_phases("tool-wait", updates=4, calibration={"mixes": calibration["mixes"]}) + bad_tail = {"mixes": {"tail": {"short": 0.2, "long": 0.8}}} + with pytest.raises(ManifestError, match="minority of long"): + work.scenario_phases("tail", updates=4, calibration=bad_tail) + + +def test_length_diagnostics_report_caps_and_zero_advantage(): + samples = [ + {"group": 0, "reward": 1.0, "response_length": 10, "status": "completed"}, + {"group": 0, "reward": 1.0, "response_length": 64, "status": "truncated"}, + {"group": 1, "reward": 0.0, "response_length": 5, "status": "completed"}, + {"group": 1, "reward": 1.0, "response_length": 7, "status": "completed"}, + ] + diag = work.length_diagnostics(samples, max_response_len=64) + assert diag["cap_hit_ratio"] == 0.25 and diag["zero_advantage_ratio"] == 0.5 + assert diag["max_response_length"] == 64 diff --git a/yeto/rl/elastic_benchmark/__init__.py b/yeto/rl/elastic_benchmark/__init__.py new file mode 100644 index 00000000..58a617e3 --- /dev/null +++ b/yeto/rl/elastic_benchmark/__init__.py @@ -0,0 +1,7 @@ +"""RL elastic resource benchmark suite. + +Study contracts, capability gating, evidence indexing, paired workloads and +layered reporting for fixed-versus-elastic GPU partition experiments. Nothing +in this package imports a GPU runtime; runners plug in through the runtime +capability attestation described in ``capabilities.py``. +""" diff --git a/yeto/rl/elastic_benchmark/capabilities.py b/yeto/rl/elastic_benchmark/capabilities.py new file mode 100644 index 00000000..a6848c6c --- /dev/null +++ b/yeto/rl/elastic_benchmark/capabilities.py @@ -0,0 +1,306 @@ +"""Resource configs, directed edges, physical pools and runtime attestation. + +Everything here is pure validation. A config that fails is reported with its +reason and never replaced by a neighbouring legal config; an edge that the +runtime has not attested is ``blocked_dependency``, never silently downgraded. +""" + +from __future__ import annotations + +import json +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +from yeto.rl.elastic_benchmark.manifest import ( + ARM_KINDS, + EDGE_KINDS, + EXECUTION_MODES, + FIXED_ARM_KINDS, + ManifestError, +) + +LEGACY_CONFIG = "colocated" +STATUS_SUPPORTED = "supported" +STATUS_UNSUPPORTED = "unsupported" +STATUS_BLOCKED = "blocked_dependency" +_ARM_EDGE_KINDS = { + "scheduled-rebuild": ("rollout-only", "trainer-dp", "role-transfer", "standby-scale"), + "scheduled-optimized": ("rollout-only", "trainer-dp", "role-transfer", "standby-scale"), + "auto": ("rollout-only", "trainer-dp", "role-transfer", "standby-scale"), +} + + +@dataclass(frozen=True) +class ResourceConfig: + name: str + trainer: int + rollout: int + standby: int = 0 + + @property + def total(self) -> int: + return self.trainer + self.rollout + self.standby + + @property + def data_parallel(self) -> int: + # First round certifies TP=PP=CP=EP=1, so trainer GPUs are the DP size. + return self.trainer + + def gradient_accumulation(self, global_batch: int, micro_batch: int) -> int: + return global_batch // (self.data_parallel * micro_batch) + + +@dataclass(frozen=True) +class Edge: + source: str + target: str + kind: str + + @property + def key(self) -> tuple[str, str, str]: + return (self.source, self.target, self.kind) + + +@dataclass(frozen=True) +class Attestation: + """What the runtime says it can do. Absent attestation means nothing is certified.""" + + runtime_fingerprint: str | None + execution_modes: frozenset[str] + certified_edges: frozenset[tuple[str, str, str]] + optimized_paths: frozenset[str] + auto_controller: bool + partitioned_driver: bool + + @staticmethod + def none() -> "Attestation": + return Attestation(None, frozenset(), frozenset(), frozenset(), False, False) + + +def load_attestation(path: Path | None) -> Attestation: + if path is None: + return Attestation.none() + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise ManifestError(f"cannot read capability attestation {path}: {exc}") from exc + return attestation_from_dict(payload) + + +def attestation_from_dict(payload: dict[str, Any]) -> Attestation: + if not isinstance(payload, dict): + raise ManifestError("capability attestation must be a JSON object") + modes = payload.get("execution_modes", []) + unknown = sorted(set(modes) - set(EXECUTION_MODES)) + if unknown: + raise ManifestError(f"attestation lists unknown execution modes: {unknown}") + edges = [] + for edge in payload.get("certified_edges", []): + parsed = parse_edge(edge) + edges.append(parsed.key) + fingerprint = payload.get("runtime_fingerprint") + if fingerprint is not None and not isinstance(fingerprint, str): + raise ManifestError("attestation runtime_fingerprint must be a string") + return Attestation( + runtime_fingerprint=fingerprint, + execution_modes=frozenset(modes), + certified_edges=frozenset(edges), + optimized_paths=frozenset(payload.get("optimized_paths", [])), + auto_controller=bool(payload.get("auto_controller", False)), + partitioned_driver=bool(payload.get("partitioned_driver", False)), + ) + + +def parse_edge(edge: dict[str, Any]) -> Edge: + if not isinstance(edge, dict): + raise ManifestError("edges must be objects with source, target and kind") + source, target, kind = edge.get("source"), edge.get("target"), edge.get("kind") + if not source or not target or source == target: + raise ManifestError(f"edge needs distinct source and target: {edge}") + if kind not in EDGE_KINDS: + raise ManifestError(f"edge {source}->{target} has unknown kind {kind!r}") + return Edge(str(source), str(target), str(kind)) + + +def parse_configs(resources: dict[str, Any]) -> dict[str, ResourceConfig]: + raw = resources.get("configs") + if not isinstance(raw, dict) or not raw: + raise ManifestError("resources.configs must be a non-empty object") + configs = {} + for name, block in raw.items(): + if name == LEGACY_CONFIG: + raise ManifestError(f"{LEGACY_CONFIG!r} is reserved for the legacy colocated arm") + if not isinstance(block, dict): + raise ManifestError(f"config {name!r} must be an object") + values = {} + for role in ("trainer", "rollout", "standby"): + count = block.get(role, 0) + if not isinstance(count, int) or isinstance(count, bool) or count < 0: + raise ManifestError(f"config {name!r}.{role} must be a non-negative integer") + values[role] = count + configs[name] = ResourceConfig(name, **values) + return configs + + +def validate_pool(resources: dict[str, Any]) -> int | None: + """Return the physical pool size, or None when the pool is still unresolved.""" + gpus = resources.get("gpus") + if not isinstance(gpus, list): + raise ManifestError("resources.gpus must be a list") + if not gpus: + return None + seen = set() + for gpu in gpus: + if not isinstance(gpu, dict) or not isinstance(gpu.get("uuid"), str): + raise ManifestError("each resources.gpus entry needs a string uuid") + if gpu["uuid"] in seen: + raise ManifestError(f"duplicate GPU uuid in resources.gpus: {gpu['uuid']}") + seen.add(gpu["uuid"]) + return len(gpus) + + +def config_rejection( + config: ResourceConfig, + *, + profile: dict[str, Any], + pool_size: int | None, +) -> str | None: + """Why this config is illegal for the profile and pool, or None when legal.""" + if config.trainer < 1: + return "needs at least one trainer GPU" + if config.rollout < 1: + return "needs at least one rollout GPU" + if pool_size is not None and config.total != pool_size: + return f"uses {config.total} GPUs but the pool has {pool_size}" + if profile["parameter_mode"] == "full" and config.data_parallel > 1: + return "dense full-parameter profile requires trainer DP=1" + divisor = config.data_parallel * profile["micro_batch"] + if profile["global_batch"] % divisor: + return ( + f"global_batch {profile['global_batch']} is not divisible by " + f"DP {config.data_parallel} x micro_batch {profile['micro_batch']}" + ) + return None + + +def validate_edges(resources: dict[str, Any], configs: dict[str, ResourceConfig]) -> list[Edge]: + edges = [parse_edge(edge) for edge in resources.get("edges", [])] + keys = [edge.key for edge in edges] + if len(set(keys)) != len(keys): + raise ManifestError("resources.edges contains duplicate edges") + for edge in edges: + for endpoint in (edge.source, edge.target): + if endpoint not in configs: + raise ManifestError(f"edge references unknown config {endpoint!r}") + _check_edge_shape(edge, configs[edge.source], configs[edge.target]) + return edges + + +def _check_edge_shape(edge: Edge, source: ResourceConfig, target: ResourceConfig) -> None: + trainer_changes = source.trainer != target.trainer + if edge.kind in ("rollout-only", "standby-scale") and trainer_changes: + raise ManifestError(f"{edge.kind} edge {edge.source}->{edge.target} changes trainer count") + if edge.kind == "same-shape-restore" and trainer_changes: + raise ManifestError(f"same-shape-restore edge {edge.source}->{edge.target} changes DP") + if edge.kind in ("trainer-dp", "role-transfer") and not trainer_changes: + raise ManifestError(f"{edge.kind} edge {edge.source}->{edge.target} keeps trainer count") + + +def arm_status( + arm: dict[str, Any], + *, + profile: dict[str, Any], + configs: dict[str, ResourceConfig], + edges: list[Edge], + attestation: Attestation, + pool_size: int | None, +) -> tuple[str, str | None]: + """Classify an arm as supported / unsupported / blocked_dependency with a reason.""" + kind = arm["kind"] + if kind not in ARM_KINDS: + return STATUS_UNSUPPORTED, f"unknown arm kind {kind!r}" + if kind == "legacy-fixed": + return _legacy_status(arm) + for name in _arm_config_names(arm): + if name not in configs: + return STATUS_UNSUPPORTED, f"unknown config {name!r}" + reason = config_rejection(configs[name], profile=profile, pool_size=pool_size) + if reason: + return STATUS_UNSUPPORTED, f"config {name}: {reason}" + if profile["execution_mode"] not in attestation.execution_modes: + return STATUS_BLOCKED, f"runtime has not attested {profile['execution_mode']}" + if not attestation.partitioned_driver: + return STATUS_BLOCKED, "runtime has not attested the partitioned driver" + if kind in FIXED_ARM_KINDS: + return STATUS_SUPPORTED, None + return _dynamic_status(arm, edges=edges, attestation=attestation) + + +def _legacy_status(arm: dict[str, Any]) -> tuple[str, str | None]: + if arm.get("config") != LEGACY_CONFIG: + return STATUS_UNSUPPORTED, f"legacy-fixed arm must use config {LEGACY_CONFIG!r}" + return STATUS_SUPPORTED, None + + +def _dynamic_status( + arm: dict[str, Any], *, edges: list[Edge], attestation: Attestation +) -> tuple[str, str | None]: + declared = {(e.source, e.target): e for e in edges} + for source, target in _arm_transitions(arm): + edge = declared.get((source, target)) + if edge is None: + return STATUS_UNSUPPORTED, f"no declared edge {source}->{target}" + if edge.key not in attestation.certified_edges: + return STATUS_BLOCKED, f"edge {source}->{target} ({edge.kind}) is not certified" + if arm["kind"] == "scheduled-optimized" and not attestation.optimized_paths: + return STATUS_BLOCKED, "no optimized migration path is certified" + if arm["kind"] == "auto" and not attestation.auto_controller: + return STATUS_BLOCKED, "auto controller is not certified" + return STATUS_SUPPORTED, None + + +def _arm_config_names(arm: dict[str, Any]) -> list[str]: + names = [] + if arm.get("config"): + names.append(arm["config"]) + names.extend(arm.get("configs", [])) + names.extend(arm.get("candidates", [])) + names.extend(step["target"] for step in arm.get("switch_plan", [])) + return list(dict.fromkeys(names)) + + +def _arm_transitions(arm: dict[str, Any]) -> list[tuple[str, str]]: + if arm["kind"] == "auto": + candidates = list(arm["candidates"]) + return [(a, b) for a in candidates for b in candidates if a != b] + current = arm["config"] + transitions = [] + for step in arm.get("switch_plan", []): + transitions.append((current, step["target"])) + current = step["target"] + return transitions + + +def config_table( + configs: dict[str, ResourceConfig], *, profile: dict[str, Any], pool_size: int | None +) -> list[dict[str, Any]]: + """Legal table rows record accumulation; illegal rows keep their rejection reason.""" + rows = [] + for config in configs.values(): + reason = config_rejection(config, profile=profile, pool_size=pool_size) + row = { + "config": config.name, + "trainer": config.trainer, + "rollout": config.rollout, + "standby": config.standby, + "data_parallel": config.data_parallel, + "legal": reason is None, + "reason": reason, + } + if reason is None: + row["gradient_accumulation"] = config.gradient_accumulation( + profile["global_batch"], profile["micro_batch"] + ) + rows.append(row) + return rows diff --git a/yeto/rl/elastic_benchmark/cli.py b/yeto/rl/elastic_benchmark/cli.py new file mode 100644 index 00000000..0f35596e --- /dev/null +++ b/yeto/rl/elastic_benchmark/cli.py @@ -0,0 +1,82 @@ +"""``benchmark_rl_elastic`` command line: example, validate, plan, report. + +None of these subcommands import torch, ray, miles or a cloud SDK, load a +model, or create resources. Execution of matrix items is delegated to runners +registered against a capability attestation and lives outside this module. +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +from yeto.rl.elastic_benchmark import capabilities as caps +from yeto.rl.elastic_benchmark import results as results_mod +from yeto.rl.elastic_benchmark.manifest import ( + MODES, + ManifestError, + dump_manifest, + example_manifest, + load_manifest, + manifest_hash, +) +from yeto.rl.elastic_benchmark.plan import build_plan, plan_as_dict, render_plan + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + prog="benchmark_rl_elastic", + description="Plan and report RL elastic resource benchmark studies (no GPU work here).", + ) + commands = parser.add_subparsers(dest="command", required=True) + + example = commands.add_parser("example", help="write the documented starting manifest") + example.add_argument("--mode", choices=MODES, default="calibration") + example.add_argument("--output", type=Path, required=True) + + validate = commands.add_parser("validate", help="validate a study manifest and print its hash") + validate.add_argument("--study", type=Path, required=True) + + plan = commands.add_parser("plan", help="dry-run: expand the matrix and budget without running") + plan.add_argument("--study", type=Path, required=True) + plan.add_argument("--capabilities", type=Path, help="runtime capability attestation JSON") + plan.add_argument("--json", action="store_true", help="print the plan as JSON") + + report = commands.add_parser("report", help="summarize verified evidence under a study directory") + report.add_argument("--study", type=Path, required=True) + report.add_argument("--capabilities", type=Path) + report.add_argument("--study-dir", type=Path, required=True) + return parser + + +def main(argv: list[str] | None = None) -> int: + args = build_parser().parse_args(argv) + try: + return _dispatch(args) + except (ManifestError, caps.ManifestError, ValueError) as exc: + print(f"error: {exc}", file=sys.stderr) + return 2 + + +def _dispatch(args: argparse.Namespace) -> int: + if args.command == "example": + digest = dump_manifest(example_manifest(args.mode), args.output) + print(f"wrote {args.output} (study_hash {digest[:12]})") + return 0 + manifest = load_manifest(args.study) + study_hash = manifest_hash(manifest) + if args.command == "validate": + print(f"valid {manifest['mode']} study {manifest['identity']['study_id']} hash={study_hash}") + return 0 + attestation = caps.load_attestation(args.capabilities) + plan = build_plan(manifest, attestation, study_hash=study_hash) + if args.command == "plan": + print(json.dumps(plan_as_dict(plan), indent=2, sort_keys=True) if args.json else render_plan(plan)) + return 0 + summary = results_mod.summarize(plan, args.study_dir) + results_mod.write_summary(summary, args.study_dir) + path = results_mod.write_report(summary, args.study_dir) + print(f"{summary['status']} benefit={summary['benefit']} report={path}") + return 0 if summary["status"] == "complete" else 1 diff --git a/yeto/rl/elastic_benchmark/evidence.py b/yeto/rl/elastic_benchmark/evidence.py new file mode 100644 index 00000000..b9debb58 --- /dev/null +++ b/yeto/rl/elastic_benchmark/evidence.py @@ -0,0 +1,186 @@ +"""Evidence index, per-attempt results and resume validation. + +Every attempt directory carries ``evidence-index.json`` (content digests of the +raw files it produced) and ``result.json`` (the four result layers). Resume +only reuses attempts whose study hash, matrix key and every digest still +match; failed attempts are kept as evidence and never overwritten. +""" + +from __future__ import annotations + +import hashlib +import json +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +from yeto.benchmark_resume import write_json_atomic + +EVIDENCE_INDEX = "evidence-index.json" +RESULT_FILE = "result.json" +EXECUTION_STATES = ("completed", "failed", "unsupported", "blocked_dependency", "pending") +_CHUNK = 1024 * 1024 + + +class EvidenceError(ValueError): + """Evidence is missing, altered or belongs to a different study.""" + + +@dataclass(frozen=True, order=True) +class MatrixKey: + arm: str + scenario: str + seed: int + config: str | None = None + + def relative_dir(self, attempt: int) -> Path: + arm = self.arm if self.config is None else f"{self.arm}@{self.config}" + return Path("runs") / arm / self.scenario / f"seed-{self.seed}" / f"attempt-{attempt}" + + def as_dict(self) -> dict[str, Any]: + return {"arm": self.arm, "scenario": self.scenario, "seed": self.seed, "config": self.config} + + @staticmethod + def from_dict(payload: dict[str, Any]) -> "MatrixKey": + try: + return MatrixKey( + str(payload["arm"]), str(payload["scenario"]), int(payload["seed"]), payload.get("config") + ) + except (KeyError, TypeError, ValueError) as exc: + raise EvidenceError(f"malformed matrix key: {payload}") from exc + + +def file_digest(path: Path) -> tuple[str, int]: + hasher = hashlib.sha256() + size = 0 + with path.open("rb") as handle: + while chunk := handle.read(_CHUNK): + hasher.update(chunk) + size += len(chunk) + return hasher.hexdigest(), size + + +def build_evidence_index( + attempt_dir: Path, *, study_hash: str, key: MatrixKey, attempt: int +) -> dict[str, Any]: + """Digest every raw file under the attempt directory except the index and result.""" + files = {} + for path in sorted(attempt_dir.rglob("*")): + if path.is_symlink(): + raise EvidenceError(f"evidence directories do not support symlinks: {path}") + if not path.is_file() or path.name in (EVIDENCE_INDEX, RESULT_FILE): + continue + digest, size = file_digest(path) + files[path.relative_to(attempt_dir).as_posix()] = {"sha256": digest, "bytes": size} + index = { + "format_version": 1, + "study_hash": study_hash, + "key": key.as_dict(), + "attempt": attempt, + "files": files, + } + write_json_atomic(attempt_dir / EVIDENCE_INDEX, index) + return index + + +def verify_evidence_index(attempt_dir: Path, *, study_hash: str, key: MatrixKey) -> dict[str, Any]: + index = _read_json(attempt_dir / EVIDENCE_INDEX, "evidence index") + if index.get("study_hash") != study_hash: + raise EvidenceError(f"{attempt_dir}: evidence belongs to a different study") + if MatrixKey.from_dict(index.get("key") or {}) != key: + raise EvidenceError(f"{attempt_dir}: evidence belongs to a different matrix item") + files = index.get("files") + if not isinstance(files, dict): + raise EvidenceError(f"{attempt_dir}: evidence index has no file table") + for relative, expected in files.items(): + path = attempt_dir / relative + if not path.is_file(): + raise EvidenceError(f"{attempt_dir}: evidence file is missing: {relative}") + digest, size = file_digest(path) + if digest != expected.get("sha256") or size != expected.get("bytes"): + raise EvidenceError(f"{attempt_dir}: evidence file was modified: {relative}") + extra = sorted( + p.relative_to(attempt_dir).as_posix() + for p in attempt_dir.rglob("*") + if p.is_file() and p.name not in (EVIDENCE_INDEX, RESULT_FILE) + ) + unexpected = sorted(set(extra) - set(files)) + if unexpected: + raise EvidenceError(f"{attempt_dir}: unindexed evidence files: {unexpected[:4]}") + return index + + +def write_result(attempt_dir: Path, result: dict[str, Any]) -> None: + if result.get("execution") not in EXECUTION_STATES: + raise EvidenceError(f"result.execution must be one of {EXECUTION_STATES}") + for layer in ("correctness", "quality", "benefit"): + if layer not in result: + raise EvidenceError(f"result is missing the {layer!r} layer") + target = attempt_dir / RESULT_FILE + if target.exists(): + raise EvidenceError(f"refusing to overwrite an existing attempt result: {target}") + write_json_atomic(target, result) + + +def _read_json(path: Path, label: str) -> dict[str, Any]: + if not path.is_file(): + raise EvidenceError(f"{label} is missing: {path}") + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise EvidenceError(f"cannot read {label} {path}: {exc}") from exc + if not isinstance(payload, dict): + raise EvidenceError(f"{label} must be a JSON object: {path}") + return payload + + +def list_attempts(study_dir: Path, key: MatrixKey) -> list[tuple[int, Path]]: + base = study_dir / key.relative_dir(0).parent + if not base.is_dir(): + return [] + attempts = [] + for child in base.iterdir(): + if child.is_dir() and child.name.startswith("attempt-"): + try: + attempts.append((int(child.name.split("-", 1)[1]), child)) + except ValueError: + raise EvidenceError(f"malformed attempt directory: {child}") from None + return sorted(attempts) + + +def next_attempt(study_dir: Path, key: MatrixKey) -> int: + attempts = list_attempts(study_dir, key) + return attempts[-1][0] + 1 if attempts else 1 + + +def load_attempt_result(attempt_dir: Path) -> dict[str, Any] | None: + path = attempt_dir / RESULT_FILE + return _read_json(path, "attempt result") if path.is_file() else None + + +def reusable_completion(study_dir: Path, *, study_hash: str, key: MatrixKey) -> dict[str, Any] | None: + """The verified completed result for a key, or None when it must be (re)run. + + A completed attempt whose evidence fails verification is an error, not a + silent rerun: a benchmark must not launder tampered evidence by repeating. + """ + for _attempt, attempt_dir in list_attempts(study_dir, key): + result = load_attempt_result(attempt_dir) + if result is None or result.get("execution") != "completed": + continue + verify_evidence_index(attempt_dir, study_hash=study_hash, key=key) + return result + return None + + +def collect_results(study_dir: Path, keys: list[MatrixKey]) -> dict[MatrixKey, list[dict[str, Any]]]: + """Every attempt result per key, failed ones included, oldest first.""" + collected = {} + for key in keys: + results = [] + for attempt, attempt_dir in list_attempts(study_dir, key): + result = load_attempt_result(attempt_dir) + if result is not None: + results.append({**result, "attempt": attempt}) + collected[key] = results + return collected diff --git a/yeto/rl/elastic_benchmark/legacy.py b/yeto/rl/elastic_benchmark/legacy.py new file mode 100644 index 00000000..6a7c0a0c --- /dev/null +++ b/yeto/rl/elastic_benchmark/legacy.py @@ -0,0 +1,90 @@ +"""Adapter over the existing ``scripts/benchmark_rl.py`` harness. + +The legacy harness is a script, not a package, so it is loaded by path. Only +its prompt pairing, normalisation and JSONL helpers are reused; its CLI and +result contract stay untouched. Loading is lazy so ``plan``/``validate`` never +pay for the script's imports. +""" + +from __future__ import annotations + +import importlib.util +import sys +from functools import lru_cache +from pathlib import Path +from types import ModuleType +from typing import Any + +REPO_ROOT = Path(__file__).resolve().parents[3] +LEGACY_SCRIPT = REPO_ROOT / "scripts" / "benchmark_rl.py" +_MODULE_NAME = "yeto_legacy_benchmark_rl" + + +@lru_cache(maxsize=1) +def load_legacy_harness(script: Path = LEGACY_SCRIPT) -> ModuleType: + if not script.is_file(): + raise FileNotFoundError(f"legacy RL benchmark harness is missing: {script}") + spec = importlib.util.spec_from_file_location(_MODULE_NAME, script) + if spec is None or spec.loader is None: + raise ImportError(f"cannot load legacy harness from {script}") + module = importlib.util.module_from_spec(spec) + sys.modules[_MODULE_NAME] = module + spec.loader.exec_module(module) + return module + + +def read_jsonl(path: Path) -> list[dict[str, Any]]: + return load_legacy_harness()._read_jsonl(path) + + +def paired_streams(rows: list[dict[str, Any]], *, islands: int, groups: int, rounds: int): + """Round-major paired prompt streams, identical to the legacy arms' pairing.""" + return load_legacy_harness().paired_prompt_streams(rows, islands=islands, groups=groups, rounds=rounds) + + +def write_prompt_files(streams, evaluation_rows: list[dict[str, Any]], directory: Path): + return load_legacy_harness().write_prompt_files(streams, evaluation_rows, directory) + + +def normalized_prompt(row: dict[str, Any], prompt_id: int | str) -> dict[str, Any]: + return load_legacy_harness()._normalized_prompt(row, prompt_id) + + +def materialize_paired_inputs( + rows: list[dict[str, Any]], + *, + train_ids: list[int], + held_out_ids: list[int], + assignment: list[list[int]], + directory: Path, +) -> dict[str, Any]: + """Write the elastic study's per-update prompt stream in the legacy file format. + + ``assignment`` lists train row indexes per logical update (see + ``work.assign_groups``). Every arm reads the same ``combined.jsonl``; + ``eval.jsonl`` holds held-out rows only. + """ + leaked = sorted(set(held_out_ids) & {i for update in assignment for i in update}) + if leaked: + raise ValueError(f"held-out rows leaked into training assignment: {leaked[:4]}") + stream_rows = [dict(rows[i]) for update in assignment for i in update] + stream_ids = [i for update in assignment for i in update] + harness = load_legacy_harness() + streams = harness.PromptStreams( + combined_rows=tuple(stream_rows), + island_rows=(tuple(stream_rows),), + combined_ids=tuple(stream_ids), + island_ids=(tuple(stream_ids),), + ) + combined, island_paths, evaluation = harness.write_prompt_files( + streams, [dict(rows[i]) for i in held_out_ids], directory + ) + return { + "combined": str(combined), + "islands": [str(p) for p in island_paths], + "eval": str(evaluation), + "train_rows": len(stream_rows), + "held_out_rows": len(held_out_ids), + "updates": len(assignment), + "train_ids_unique": len(set(train_ids)), + } diff --git a/yeto/rl/elastic_benchmark/manifest.py b/yeto/rl/elastic_benchmark/manifest.py new file mode 100644 index 00000000..923dd124 --- /dev/null +++ b/yeto/rl/elastic_benchmark/manifest.py @@ -0,0 +1,456 @@ +"""Versioned study manifest: field groups, canonical hash and mode validation. + +A manifest freezes one study. ``calibration`` manifests may leave resources and +quality rules unresolved; ``formal`` manifests must resolve everything before a +single GPU is touched. Validation never rewrites a manifest: an illegal or +contradictory manifest is rejected with the reason, not silently repaired. +""" + +from __future__ import annotations + +import hashlib +import json +from pathlib import Path +from typing import Any + +from yeto.provenance import is_immutable_commit + +SCHEMA_VERSION = 1 +MODES = ("calibration", "formal") +FIELD_GROUPS = ( + "identity", + "profile", + "resources", + "matrix", + "work", + "evaluation", + "timing", +) +PARAMETER_MODES = ("lora", "full") +EXECUTION_MODES = ("colocated-serial", "partitioned-serial", "partitioned-overlap") +OUTER_PROTOCOLS = ("none", "strict-avg", "decoupled") +ARM_KINDS = ( + "legacy-fixed", + "target-fixed-default", + "target-fixed-sweep", + "scheduled-rebuild", + "scheduled-optimized", + "auto", +) +FIXED_ARM_KINDS = ARM_KINDS[:3] +SCENARIOS = ("stable", "phased", "tail", "tool-wait", "oscillating") +MEASUREMENT_LEVELS = ("component", "single-turn", "agent") +EDGE_KINDS = ("rollout-only", "same-shape-restore", "trainer-dp", "role-transfer", "standby-scale") +DEFAULT_SEEDS = (17, 29, 43) +_UNRESOLVED = "unresolved" + + +class ManifestError(ValueError): + """A manifest is malformed, contradictory or not frozen enough for its mode.""" + + +def canonical_json(payload: Any) -> str: + return json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=False) + + +def manifest_hash(manifest: dict[str, Any]) -> str: + """Content hash over every field group; ``study_hash`` itself is excluded.""" + body = {key: value for key, value in manifest.items() if key != "study_hash"} + return hashlib.sha256(canonical_json(body).encode("utf-8")).hexdigest() + + +def load_manifest(path: Path) -> dict[str, Any]: + try: + manifest = json.loads(path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise ManifestError(f"cannot read study manifest {path}: {exc}") from exc + if not isinstance(manifest, dict): + raise ManifestError("study manifest must be a JSON object") + validate_manifest(manifest) + recorded = manifest.get("study_hash") + if recorded is not None and recorded != manifest_hash(manifest): + raise ManifestError("study manifest content does not match its recorded study_hash") + return manifest + + +def dump_manifest(manifest: dict[str, Any], path: Path) -> str: + """Validate, stamp ``study_hash`` and write atomically. Returns the hash.""" + validate_manifest(manifest) + stamped = {key: value for key, value in manifest.items() if key != "study_hash"} + stamped["study_hash"] = manifest_hash(stamped) + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_name(path.name + ".tmp") + temporary.write_text(json.dumps(stamped, indent=2, sort_keys=True) + "\n", encoding="utf-8") + temporary.replace(path) + return stamped["study_hash"] + + +def validate_manifest(manifest: dict[str, Any]) -> None: + _require_shape(manifest) + mode = manifest["mode"] + _validate_identity(manifest["identity"], mode) + _validate_profile(manifest["profile"]) + _validate_matrix(manifest["matrix"]) + _validate_work(manifest["work"], manifest["profile"]) + _validate_evaluation(manifest["evaluation"], mode) + _validate_timing(manifest["timing"], manifest["work"]) + if mode == "formal": + _reject_unresolved(manifest) + + +def _require_shape(manifest: dict[str, Any]) -> None: + if manifest.get("schema_version") != SCHEMA_VERSION: + raise ManifestError(f"study manifest schema_version must be {SCHEMA_VERSION}") + if manifest.get("mode") not in MODES: + raise ManifestError(f"study manifest mode must be one of {MODES}") + for group in FIELD_GROUPS: + if not isinstance(manifest.get(group), dict): + raise ManifestError(f"study manifest is missing the {group!r} field group") + + +def _validate_identity(identity: dict[str, Any], mode: str) -> None: + if not isinstance(identity.get("study_id"), str) or not identity["study_id"]: + raise ManifestError("identity.study_id must be a non-empty string") + model = identity.get("model") + if not isinstance(model, dict) or not isinstance(model.get("id"), str): + raise ManifestError("identity.model.id is required") + fingerprints = identity.get("fingerprints") + if not isinstance(fingerprints, dict): + raise ManifestError("identity.fingerprints must be an object") + if mode != "formal": + return + for label, block in (("model", model), ("data", identity.get("data") or {})): + revision = block.get("revision") + if not isinstance(revision, str) or not is_immutable_commit(revision): + raise ManifestError(f"formal study requires an immutable identity.{label}.revision") + for key in ("reward", "source", "runtime"): + if not isinstance(fingerprints.get(key), str) or not fingerprints[key]: + raise ManifestError(f"formal study requires identity.fingerprints.{key}") + + +def _validate_profile(profile: dict[str, Any]) -> None: + if profile.get("parameter_mode") not in PARAMETER_MODES: + raise ManifestError(f"profile.parameter_mode must be one of {PARAMETER_MODES}") + if profile.get("execution_mode") not in EXECUTION_MODES: + raise ManifestError(f"profile.execution_mode must be one of {EXECUTION_MODES}") + if profile.get("outer_protocol") not in OUTER_PROTOCOLS: + raise ManifestError(f"profile.outer_protocol must be one of {OUTER_PROTOCOLS}") + for key in ("global_batch", "micro_batch", "optimizer_steps_per_round"): + if not _positive_int(profile.get(key)): + raise ManifestError(f"profile.{key} must be a positive integer") + if profile["global_batch"] % profile["micro_batch"]: + raise ManifestError("profile.micro_batch must divide profile.global_batch") + age = profile.get("max_policy_age", 0) + if not isinstance(age, int) or age < 0: + raise ManifestError("profile.max_policy_age must be a non-negative integer") + if profile["execution_mode"] != "partitioned-overlap" and age != 0: + raise ManifestError("only partitioned-overlap profiles may allow a non-zero policy age") + for key in ("publish_rule", "reset_rule", "loss_normalization"): + if not isinstance(profile.get(key), str) or not profile[key]: + raise ManifestError(f"profile.{key} must be a non-empty string") + + +def _validate_matrix(matrix: dict[str, Any]) -> None: + arms = matrix.get("arms") + if not isinstance(arms, list) or not arms: + raise ManifestError("matrix.arms must be a non-empty list") + names = [arm.get("name") for arm in arms if isinstance(arm, dict)] + if len(names) != len(arms) or len(set(names)) != len(names) or not all(names): + raise ManifestError("matrix.arms entries need unique non-empty names") + for arm in arms: + _validate_arm(arm) + scenarios = matrix.get("scenarios") + if not isinstance(scenarios, list) or not scenarios: + raise ManifestError("matrix.scenarios must be a non-empty list") + unknown = sorted(set(scenarios) - set(SCENARIOS)) + if unknown or len(set(scenarios)) != len(scenarios): + raise ManifestError(f"matrix.scenarios must be distinct entries of {SCENARIOS}") + seeds = matrix.get("seeds") + if not isinstance(seeds, list) or not seeds or not all(_positive_int(s) for s in seeds): + raise ManifestError("matrix.seeds must be a non-empty list of positive integers") + if len(set(seeds)) != len(seeds): + raise ManifestError("matrix.seeds contains duplicates") + if not _positive_int(matrix.get("repeats", 1)): + raise ManifestError("matrix.repeats must be a positive integer") + if matrix.get("measurement_level") not in MEASUREMENT_LEVELS: + raise ManifestError(f"matrix.measurement_level must be one of {MEASUREMENT_LEVELS}") + + +def _validate_arm(arm: dict[str, Any]) -> None: + kind = arm.get("kind") + if kind not in ARM_KINDS: + raise ManifestError(f"arm {arm.get('name')!r} has unknown kind {kind!r}") + if kind == "target-fixed-sweep": + configs = arm.get("configs") + if not isinstance(configs, list) or not configs: + raise ManifestError(f"sweep arm {arm['name']!r} needs a non-empty configs list") + elif kind in ("scheduled-rebuild", "scheduled-optimized"): + plan = arm.get("switch_plan") + if not isinstance(plan, list) or not plan: + raise ManifestError(f"scheduled arm {arm['name']!r} needs a switch_plan") + for step in plan: + if not isinstance(step, dict) or not _positive_int(step.get("at_update")): + raise ManifestError(f"arm {arm['name']!r} switch_plan entries need at_update") + if not step.get("target"): + raise ManifestError(f"arm {arm['name']!r} switch_plan entries need a target") + if not arm.get("config"): + raise ManifestError(f"scheduled arm {arm['name']!r} needs an initial config") + elif kind == "auto": + candidates = arm.get("candidates") + if not isinstance(candidates, list) or len(candidates) < 2: + raise ManifestError(f"auto arm {arm['name']!r} needs at least two candidates") + if not arm.get("config"): + raise ManifestError(f"auto arm {arm['name']!r} needs an initial config") + elif not arm.get("config"): + raise ManifestError(f"fixed arm {arm['name']!r} needs a config") + + +def _validate_work(work: dict[str, Any], profile: dict[str, Any]) -> None: + for key in ("groups_per_update", "samples_per_group", "update_budget", "max_response_len"): + if not _positive_int(work.get(key)): + raise ManifestError(f"work.{key} must be a positive integer") + if work["groups_per_update"] * work["samples_per_group"] != profile["global_batch"]: + raise ManifestError( + "work.groups_per_update * work.samples_per_group must equal profile.global_batch" + ) + phases = work.get("phases") + if not isinstance(phases, list) or not phases: + raise ManifestError("work.phases must be a non-empty list") + total = 0 + for phase in phases: + if not isinstance(phase, dict) or not phase.get("name"): + raise ManifestError("work.phases entries need a name") + if not _positive_int(phase.get("updates")): + raise ManifestError(f"phase {phase.get('name')!r} needs a positive updates count") + total += phase["updates"] + if total != work["update_budget"]: + raise ManifestError( + f"work.update_budget ({work['update_budget']}) contradicts the phase sum ({total})" + ) + for key in ("truncation", "filter", "retry"): + if not isinstance(work.get(key), str) or not work[key]: + raise ManifestError(f"work.{key} must name a rule") + + +def _validate_evaluation(evaluation: dict[str, Any], mode: str) -> None: + splits = evaluation.get("splits") + if not isinstance(splits, dict): + raise ManifestError("evaluation.splits must be an object") + for key in ("train", "calibration", "test", "held_out"): + if not isinstance(splits.get(key), int) or splits[key] < 0: + raise ManifestError(f"evaluation.splits.{key} must be a non-negative integer") + if splits["train"] < 1 or splits["held_out"] < 1: + raise ManifestError("evaluation.splits needs at least one train and one held_out row") + if not _positive_int(evaluation.get("eval_seed")): + raise ManifestError("evaluation.eval_seed must be a positive integer") + points = evaluation.get("eval_points") + if not isinstance(points, list) or not all(_positive_int(p) for p in points): + raise ManifestError("evaluation.eval_points must be a list of positive update indexes") + if points != sorted(set(points)): + raise ManifestError("evaluation.eval_points must be strictly increasing") + if mode == "formal": + _validate_formal_quality_rules(evaluation) + + +def _validate_formal_quality_rules(evaluation: dict[str, Any]) -> None: + metrics = evaluation.get("metrics") + if not isinstance(metrics, list) or not metrics: + raise ManifestError("formal study requires evaluation.metrics with predeclared bounds") + for metric in metrics: + if not isinstance(metric, dict) or metric.get("direction") not in ("higher", "lower"): + raise ManifestError("each evaluation metric needs name and direction higher|lower") + for key in ("delta_quality", "catastrophic"): + if not _non_negative_number(metric.get(key)): + raise ManifestError(f"metric {metric.get('name')!r} needs numeric {key}") + for key in ("delta_stable", "min_speedup"): + if not _non_negative_number(evaluation.get(key)): + raise ManifestError(f"formal study requires numeric evaluation.{key}") + statistics = evaluation.get("statistics") + if not isinstance(statistics, dict) or not statistics.get("method"): + raise ManifestError("formal study requires evaluation.statistics.method") + if not (0 < float(statistics.get("confidence", 0)) < 1): + raise ManifestError("evaluation.statistics.confidence must be in (0, 1)") + + +def _validate_timing(timing: dict[str, Any], work: dict[str, Any]) -> None: + warmup = timing.get("warmup_updates", 0) + if not isinstance(warmup, int) or warmup < 0: + raise ManifestError("timing.warmup_updates must be a non-negative integer") + measured = timing.get("measured_updates") + if not _positive_int(measured): + raise ManifestError("timing.measured_updates must be a positive integer") + if warmup + measured > work["update_budget"]: + raise ManifestError("timing.warmup_updates + measured_updates exceed work.update_budget") + if not _positive_int(timing.get("timeout_s")): + raise ManifestError("timing.timeout_s must be a positive integer") + order = timing.get("arm_order") + if order not in ("declared", "rotated", "seeded-random"): + raise ManifestError("timing.arm_order must be declared|rotated|seeded-random") + cost = timing.get("cost", {}) + if not isinstance(cost, dict): + raise ManifestError("timing.cost must be an object") + price = cost.get("gpu_hour_cost") + if price is not None and not _non_negative_number(price): + raise ManifestError("timing.cost.gpu_hour_cost must be a number or null (unknown)") + + +def _reject_unresolved(manifest: dict[str, Any]) -> None: + unresolved = sorted(_unresolved_paths(manifest)) + if unresolved: + detail = ", ".join(unresolved[:8]) + raise ManifestError(f"formal study still has unresolved fields: {detail}") + resources = manifest["resources"] + gpus = resources.get("gpus") + if not isinstance(gpus, list) or not gpus: + raise ManifestError("formal study requires resources.gpus with physical UUIDs") + if not resources.get("pool_id") or not _positive_int(resources.get("pool_epoch")): + raise ManifestError("formal study requires resources.pool_id and pool_epoch") + + +def _unresolved_paths(value: Any, prefix: str = "") -> list[str]: + if isinstance(value, dict): + found = [] + for key, item in value.items(): + name = f"{prefix}.{key}" if prefix else str(key) + found.extend(_unresolved_paths(item, name)) + return found + if isinstance(value, list): + return [p for i, item in enumerate(value) for p in _unresolved_paths(item, f"{prefix}[{i}]")] + if value == _UNRESOLVED: + return [prefix or "root"] + return [] + + +def _positive_int(value: Any) -> bool: + return isinstance(value, int) and not isinstance(value, bool) and value > 0 + + +def _non_negative_number(value: Any) -> bool: + return isinstance(value, (int, float)) and not isinstance(value, bool) and value >= 0 + + +def example_manifest(mode: str = "calibration") -> dict[str, Any]: + """The documented starting point: Qwen3-4B LoRA, GRPO, strict-avg, P62/P44/P26.""" + formal = mode == "formal" + revision = "0" * 40 if formal else _UNRESOLVED + return { + "schema_version": SCHEMA_VERSION, + "mode": mode, + "identity": { + "study_id": "elastic-calibration-example", + "created_at": "1970-01-01T00:00:00Z", + "model": {"id": "Qwen/Qwen3-4B", "revision": revision}, + "tokenizer": {"id": "Qwen/Qwen3-4B", "revision": revision}, + "data": {"id": "org/prompt-dataset", "revision": revision}, + "initial_policy": {"kind": "base"}, + "fingerprints": { + "reward": "sha256:" + "0" * 64 if formal else _UNRESOLVED, + "source": "sha256:" + "0" * 64 if formal else _UNRESOLVED, + "runtime": "sha256:" + "0" * 64 if formal else _UNRESOLVED, + "image": None, + }, + }, + "profile": { + "name": "lora-grpo-strict-avg-partitioned-serial", + "parameter_mode": "lora", + "algorithm": "grpo", + "outer_protocol": "strict-avg", + "execution_mode": "partitioned-serial", + "max_policy_age": 0, + "publish_rule": "after-every-update", + "reset_rule": "strict-avg-moment-reset", + "global_batch": 48, + "micro_batch": 1, + "optimizer_steps_per_round": 1, + "packing": False, + "loss_normalization": "token-mean", + "max_in_flight_groups": 12, + }, + "resources": { + "pool_id": "example-pool" if formal else _UNRESOLVED, + "pool_epoch": 1, + "gpus": ( + [{"uuid": f"GPU-{i:08d}", "model": "H100", "node": "n0", "index": i} for i in range(8)] + if formal + else [] + ), + "interconnect": "nvlink", + "configs": { + "P62": {"trainer": 6, "rollout": 2, "standby": 0}, + "P44": {"trainer": 4, "rollout": 4, "standby": 0}, + "P26": {"trainer": 2, "rollout": 6, "standby": 0}, + "P422": {"trainer": 4, "rollout": 2, "standby": 2}, + }, + "edges": [ + {"source": "P422", "target": "P44", "kind": "rollout-only"}, + {"source": "P44", "target": "P422", "kind": "rollout-only"}, + {"source": "P62", "target": "P44", "kind": "role-transfer"}, + {"source": "P44", "target": "P62", "kind": "role-transfer"}, + ], + "lease": None, + }, + "matrix": { + "arms": [ + {"name": "legacy", "kind": "legacy-fixed", "config": "colocated"}, + {"name": "default", "kind": "target-fixed-default", "config": "P44"}, + {"name": "sweep", "kind": "target-fixed-sweep", "configs": ["P62", "P44", "P26"]}, + { + "name": "rebuild", + "kind": "scheduled-rebuild", + "config": "P62", + "switch_plan": [ + {"at_update": 5, "target": "P44"}, + {"at_update": 9, "target": "P62"}, + ], + }, + { + "name": "optimized", + "kind": "scheduled-optimized", + "config": "P62", + "switch_plan": [ + {"at_update": 5, "target": "P44"}, + {"at_update": 9, "target": "P62"}, + ], + }, + {"name": "auto", "kind": "auto", "config": "P44", "candidates": ["P62", "P44"]}, + ], + "scenarios": ["stable", "phased"], + "seeds": list(DEFAULT_SEEDS), + "repeats": 1, + "measurement_level": "single-turn", + }, + "work": { + "groups_per_update": 12, + "samples_per_group": 4, + "update_budget": 12, + "max_response_len": 1024, + "phases": [ + {"name": "A", "updates": 4, "mix": {"short": 0.8, "long": 0.2}}, + {"name": "B", "updates": 4, "mix": {"short": 0.2, "long": 0.8}}, + {"name": "A2", "updates": 4, "mix": {"short": 0.8, "long": 0.2}}, + ], + "truncation": "hard-cap", + "filter": "none", + "retry": "none", + }, + "evaluation": { + "splits": {"train": 96, "calibration": 16, "test": 16, "held_out": 32}, + "eval_seed": 7, + "eval_points": [4, 8, 12], + "metrics": [ + {"name": "held_out_reward", "direction": "higher", "delta_quality": 0.02, "catastrophic": 0.1} + ] + if formal + else [], + "delta_stable": 0.05 if formal else _UNRESOLVED, + "min_speedup": 1.1 if formal else _UNRESOLVED, + "statistics": {"method": "paired-bootstrap", "confidence": 0.95}, + }, + "timing": { + "warmup_updates": 2, + "measured_updates": 10, + "timeout_s": 7200, + "arm_order": "rotated", + "cost": {"gpu_hour_cost": None}, + }, + } diff --git a/yeto/rl/elastic_benchmark/plan.py b/yeto/rl/elastic_benchmark/plan.py new file mode 100644 index 00000000..ee076040 --- /dev/null +++ b/yeto/rl/elastic_benchmark/plan.py @@ -0,0 +1,133 @@ +"""Matrix expansion and work budget for a study. Pure; never touches a GPU.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + +from yeto.rl.elastic_benchmark import capabilities as caps +from yeto.rl.elastic_benchmark.evidence import MatrixKey +from yeto.rl.elastic_benchmark.manifest import FIXED_ARM_KINDS + + +@dataclass(frozen=True) +class MatrixItem: + key: MatrixKey + kind: str + status: str + reason: str | None + budget: dict[str, int] = field(default_factory=dict) + + +@dataclass(frozen=True) +class StudyPlan: + study_hash: str + items: tuple[MatrixItem, ...] + config_table: list[dict[str, Any]] + pool_size: int | None + attested: bool + + @property + def runnable(self) -> tuple[MatrixItem, ...]: + return tuple(item for item in self.items if item.status == caps.STATUS_SUPPORTED) + + def counts(self) -> dict[str, int]: + counts = {caps.STATUS_SUPPORTED: 0, caps.STATUS_UNSUPPORTED: 0, caps.STATUS_BLOCKED: 0} + for item in self.items: + counts[item.status] += 1 + return counts + + def total_budget(self) -> dict[str, int]: + total = {"updates": 0, "groups": 0, "trajectories": 0, "gpu_updates": 0} + for item in self.runnable: + for name in total: + total[name] += item.budget.get(name, 0) + return total + + +def build_plan(manifest: dict[str, Any], attestation: caps.Attestation, *, study_hash: str) -> StudyPlan: + resources, profile = manifest["resources"], manifest["profile"] + configs = caps.parse_configs(resources) + pool_size = caps.validate_pool(resources) + edges = caps.validate_edges(resources, configs) + items = [] + for arm in manifest["matrix"]["arms"]: + status, reason = caps.arm_status( + arm, profile=profile, configs=configs, edges=edges, attestation=attestation, pool_size=pool_size + ) + for key in _arm_keys(arm, manifest["matrix"]): + budget = _item_budget(manifest, arm, key, configs) + items.append(MatrixItem(key, arm["kind"], status, reason, budget)) + return StudyPlan( + study_hash=study_hash, + items=tuple(sorted(items, key=lambda item: item.key)), + config_table=caps.config_table(configs, profile=profile, pool_size=pool_size), + pool_size=pool_size, + attested=attestation.runtime_fingerprint is not None, + ) + + +def _arm_keys(arm: dict[str, Any], matrix: dict[str, Any]) -> list[MatrixKey]: + config_axis = arm["configs"] if arm["kind"] == "target-fixed-sweep" else [None] + keys = [] + for config in config_axis: + for scenario in matrix["scenarios"]: + for seed in matrix["seeds"]: + keys.append(MatrixKey(arm["name"], scenario, int(seed), config)) + return keys + + +def _item_budget( + manifest: dict[str, Any], arm: dict[str, Any], key: MatrixKey, configs: dict[str, caps.ResourceConfig] +) -> dict[str, int]: + work = manifest["work"] + updates = int(work["update_budget"]) * int(manifest["matrix"].get("repeats", 1)) + groups = updates * int(work["groups_per_update"]) + config_name = key.config or arm.get("config") + gpus = configs[config_name].total if config_name in configs else 0 + return { + "updates": updates, + "groups": groups, + "trajectories": groups * int(work["samples_per_group"]), + "gpu_updates": updates * gpus, + } + + +def plan_as_dict(plan: StudyPlan) -> dict[str, Any]: + return { + "study_hash": plan.study_hash, + "attested": plan.attested, + "pool_size": plan.pool_size, + "counts": plan.counts(), + "budget": plan.total_budget(), + "configs": plan.config_table, + "items": [ + {**item.key.as_dict(), "kind": item.kind, "status": item.status, "reason": item.reason, "budget": item.budget} + for item in plan.items + ], + } + + +def render_plan(plan: StudyPlan) -> str: + lines = [f"study {plan.study_hash[:12]} pool={plan.pool_size or 'unresolved'} attested={plan.attested}"] + lines.append("configs:") + for row in plan.config_table: + verdict = f"accum={row['gradient_accumulation']}" if row["legal"] else f"REJECTED: {row['reason']}" + lines.append(f" {row['config']:<6} T{row['trainer']} R{row['rollout']} S{row['standby']} {verdict}") + lines.append("matrix:") + for item in plan.items: + label = item.key.arm if item.key.config is None else f"{item.key.arm}@{item.key.config}" + suffix = "" if item.reason is None else f" ({item.reason})" + lines.append(f" {item.status:<18} {label:<18} {item.key.scenario:<12} seed={item.key.seed}{suffix}") + counts, budget = plan.counts(), plan.total_budget() + lines.append( + f"runnable={counts['supported']} unsupported={counts['unsupported']} " + f"blocked_dependency={counts['blocked_dependency']}" + ) + lines.append( + f"runnable budget: {budget['updates']} updates, {budget['groups']} groups, " + f"{budget['trajectories']} trajectories, {budget['gpu_updates']} gpu-updates" + ) + fixed = sum(1 for item in plan.runnable if item.kind in FIXED_ARM_KINDS) + lines.append(f"fixed baseline items runnable now: {fixed}") + return "\n".join(lines) diff --git a/yeto/rl/elastic_benchmark/results.py b/yeto/rl/elastic_benchmark/results.py new file mode 100644 index 00000000..31eb9819 --- /dev/null +++ b/yeto/rl/elastic_benchmark/results.py @@ -0,0 +1,147 @@ +"""Layered results and the minimal study report. + +Each attempt result has four independent layers: execution (did it run), +correctness (ledger reconciles), quality (predeclared held-out rules) and +benefit (predeclared system gain). A study summary is ``incomplete`` whenever a +required matrix item lacks a completed, verified attempt; survivors are never +averaged as if the matrix were whole. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +from yeto.benchmark_resume import write_json_atomic +from yeto.rl.elastic_benchmark import evidence +from yeto.rl.elastic_benchmark.plan import StudyPlan + +CORRECTNESS = ("passed", "failed", "not_checked") +QUALITY = ("passed", "failed", "insufficient", "censored", "not_checked") +BENEFIT = ("demonstrated", "not_demonstrated", "negative", "not_checked") +_LAYER_VALUES = {"correctness": CORRECTNESS, "quality": QUALITY, "benefit": BENEFIT} + + +def make_result( + *, + execution: str, + correctness: str = "not_checked", + quality: str = "not_checked", + benefit: str = "not_checked", + reason: str | None = None, + measurements: dict[str, Any] | None = None, +) -> dict[str, Any]: + if execution not in evidence.EXECUTION_STATES: + raise ValueError(f"execution must be one of {evidence.EXECUTION_STATES}") + layers = {"correctness": correctness, "quality": quality, "benefit": benefit} + for layer, value in layers.items(): + if value not in _LAYER_VALUES[layer]: + raise ValueError(f"{layer} must be one of {_LAYER_VALUES[layer]}") + if execution != "completed" and any(v not in ("not_checked",) for v in layers.values()): + raise ValueError("only completed attempts may carry correctness/quality/benefit verdicts") + return {"execution": execution, **layers, "reason": reason, "measurements": measurements or {}} + + +def summarize(plan: StudyPlan, study_dir: Path) -> dict[str, Any]: + keys = [item.key for item in plan.items] + collected = evidence.collect_results(study_dir, keys) + rows, missing, failed, tampered = [], [], [], [] + for item in plan.items: + attempts = collected.get(item.key, []) + verified = _verified_completion(study_dir, plan.study_hash, item.key, attempts, tampered) + rows.append(_row(item, attempts, verified)) + if item.status != "supported": + continue + if verified is None: + missing.append(item.key.as_dict()) + if any(a.get("execution") == "failed" for a in attempts): + failed.append(item.key.as_dict()) + status = "complete" if not missing and not tampered else "incomplete" + return { + "study_hash": plan.study_hash, + "status": status, + "counts": plan.counts(), + "missing": missing, + "failed_attempts": failed, + "tampered": tampered, + "benefit": _benefit_verdict(rows, status), + "rows": rows, + } + + +def _verified_completion(study_dir, study_hash, key, attempts, tampered) -> dict[str, Any] | None: + for attempt in attempts: + if attempt.get("execution") != "completed": + continue + attempt_dir = study_dir / key.relative_dir(attempt["attempt"]) + try: + evidence.verify_evidence_index(attempt_dir, study_hash=study_hash, key=key) + except evidence.EvidenceError as exc: + tampered.append({**key.as_dict(), "attempt": attempt["attempt"], "error": str(exc)}) + continue + return attempt + return None + + +def _row(item, attempts, verified) -> dict[str, Any]: + latest = attempts[-1] if attempts else None + return { + **item.key.as_dict(), + "kind": item.kind, + "plan_status": item.status, + "plan_reason": item.reason, + "attempts": len(attempts), + "execution": verified["execution"] if verified else (latest["execution"] if latest else "pending"), + "correctness": verified["correctness"] if verified else "not_checked", + "quality": verified["quality"] if verified else "not_checked", + "benefit": verified["benefit"] if verified else "not_checked", + "reason": (verified or latest or {}).get("reason"), + } + + +def _benefit_verdict(rows: list[dict[str, Any]], status: str) -> str: + """Benefit is only claimed when the matrix is whole and every gate passes.""" + if status != "complete": + return "not_demonstrated" + dynamic = [r for r in rows if r["kind"] not in ("legacy-fixed", "target-fixed-default", "target-fixed-sweep")] + runnable = [r for r in dynamic if r["plan_status"] == "supported"] + if not runnable: + return "not_checked" + if any(r["correctness"] == "failed" for r in runnable): + return "negative" + if any(r["benefit"] == "negative" for r in runnable): + return "negative" + if all(r["benefit"] == "demonstrated" and r["quality"] == "passed" and r["correctness"] == "passed" for r in runnable): + return "demonstrated" + return "not_demonstrated" + + +def write_summary(summary: dict[str, Any], study_dir: Path) -> Path: + path = study_dir / "summary.json" + write_json_atomic(path, summary) + return path + + +def render_report(summary: dict[str, Any]) -> str: + lines = [f"# Elastic RL benchmark report", "", f"study: `{summary['study_hash']}`", f"status: **{summary['status']}**", f"benefit: **{summary['benefit']}**", ""] + if summary["missing"]: + lines.append(f"Missing required items: {len(summary['missing'])}. The matrix is incomplete; no survivor mean is reported.") + if summary["tampered"]: + lines.append(f"Evidence verification failed for {len(summary['tampered'])} attempt(s); those results are void.") + if summary["failed_attempts"]: + lines.append(f"Failed attempts retained: {len(summary['failed_attempts'])}.") + lines.extend(["", "| arm | config | scenario | seed | plan | execution | correctness | quality | benefit | reason |", "|---|---|---|---|---|---|---|---|---|---|"]) + for row in summary["rows"]: + lines.append( + f"| {row['arm']} | {row['config'] or '-'} | {row['scenario']} | {row['seed']} | {row['plan_status']} | " + f"{row['execution']} | {row['correctness']} | {row['quality']} | {row['benefit']} | {row['reason'] or row['plan_reason'] or ''} |" + ) + return "\n".join(lines) + "\n" + + +def write_report(summary: dict[str, Any], study_dir: Path) -> Path: + path = study_dir / "report.md" + temporary = path.with_name(path.name + ".tmp") + temporary.write_text(render_report(summary), encoding="utf-8") + temporary.replace(path) + return path diff --git a/yeto/rl/elastic_benchmark/work.py b/yeto/rl/elastic_benchmark/work.py new file mode 100644 index 00000000..4834b2e3 --- /dev/null +++ b/yeto/rl/elastic_benchmark/work.py @@ -0,0 +1,246 @@ +"""Paired work: data splits, logical seeds, phase schedules and scenarios. + +Splits are drawn once from a seeded permutation so every arm sees the same +train rows and the held-out rows are never handed to training. Phases bind to +logical update indexes, not wall time, so a slow arm and a fast arm observe the +same phase sequence. Scenario recipes only shape the prompt mix and external +waits; they never touch GBS, group size, optimizer steps or truncation rules. +""" + +from __future__ import annotations + +import hashlib +import random +from dataclasses import dataclass +from typing import Any + +from yeto.rl.elastic_benchmark.manifest import ManifestError + +SPLIT_NAMES = ("train", "calibration", "test", "held_out") +LENGTH_BUCKETS = ("short", "long") + + +@dataclass(frozen=True) +class Splits: + train: tuple[int, ...] + calibration: tuple[int, ...] + test: tuple[int, ...] + held_out: tuple[int, ...] + + def as_dict(self) -> dict[str, list[int]]: + return {name: list(getattr(self, name)) for name in SPLIT_NAMES} + + +def split_rows(row_count: int, sizes: dict[str, int], *, seed: int) -> Splits: + """Deterministic disjoint splits by row index. Held-out rows never reach train.""" + total = sum(sizes[name] for name in SPLIT_NAMES) + if total > row_count: + raise ManifestError(f"splits need {total} rows but only {row_count} are available") + order = list(range(row_count)) + random.Random(f"elastic-split:{seed}").shuffle(order) + cursor = 0 + parts = {} + for name in SPLIT_NAMES: + parts[name] = tuple(order[cursor : cursor + sizes[name]]) + cursor += sizes[name] + return Splits(**parts) + + +def logical_sample_seed(*, study_seed: int, group_index: int, sample_index: int) -> int: + """Stable per-sample seed so paired arms can reconcile their sampling inputs.""" + material = f"{study_seed}:{group_index}:{sample_index}".encode("utf-8") + return int.from_bytes(hashlib.sha256(material).digest()[:8], "little") + + +def group_seeds(*, study_seed: int, groups: int, samples_per_group: int) -> list[list[int]]: + return [ + [logical_sample_seed(study_seed=study_seed, group_index=g, sample_index=s) for s in range(samples_per_group)] + for g in range(groups) + ] + + +@dataclass(frozen=True) +class PhaseSlot: + update: int # 1-based logical update index + phase: str + mix: dict[str, float] + + +def phase_schedule(phases: list[dict[str, Any]]) -> list[PhaseSlot]: + """Expand phases into one slot per logical update. Independent of wall time.""" + slots = [] + update = 1 + for phase in phases: + mix = _validated_mix(phase.get("mix"), phase["name"]) + for _ in range(int(phase["updates"])): + slots.append(PhaseSlot(update, str(phase["name"]), mix)) + update += 1 + return slots + + +def _validated_mix(mix: Any, phase: str) -> dict[str, float]: + if mix is None: + return {"short": 1.0} + if not isinstance(mix, dict) or not mix: + raise ManifestError(f"phase {phase!r} mix must be a non-empty object") + total = 0.0 + for bucket, share in mix.items(): + if bucket not in LENGTH_BUCKETS: + raise ManifestError(f"phase {phase!r} uses unknown length bucket {bucket!r}") + if not isinstance(share, (int, float)) or share < 0: + raise ManifestError(f"phase {phase!r} bucket {bucket!r} share must be non-negative") + total += float(share) + if abs(total - 1.0) > 1e-6: + raise ManifestError(f"phase {phase!r} mix shares must sum to 1") + return {bucket: float(share) for bucket, share in mix.items()} + + +def bucket_counts(mix: dict[str, float], groups: int) -> dict[str, int]: + """Largest-remainder rounding so every update draws exactly ``groups`` groups.""" + raw = {bucket: share * groups for bucket, share in mix.items()} + counts = {bucket: int(value) for bucket, value in raw.items()} + remainder = groups - sum(counts.values()) + for bucket in sorted(raw, key=lambda b: (raw[b] - counts[b], b), reverse=True)[:remainder]: + counts[bucket] += 1 + return counts + + +def assign_groups( + schedule: list[PhaseSlot], + *, + pools: dict[str, list[int]], + groups_per_update: int, + seed: int, +) -> list[list[int]]: + """Prompt ids per update, drawn round-robin from each bucket's pool. + + Pools are the calibration-measured length buckets over train rows. The + same seed and pools give the same assignment for every arm. + """ + for bucket in LENGTH_BUCKETS: + if any(slot.mix.get(bucket, 0) > 0 for slot in schedule) and not pools.get(bucket): + raise ManifestError(f"schedule needs bucket {bucket!r} but its pool is empty") + cursors = {bucket: 0 for bucket in pools} + rng = random.Random(f"elastic-assign:{seed}") + shuffled = {bucket: rng.sample(ids, len(ids)) for bucket, ids in pools.items()} + assignment = [] + for slot in schedule: + chosen = [] + for bucket, count in bucket_counts(slot.mix, groups_per_update).items(): + for _ in range(count): + pool = shuffled[bucket] + chosen.append(pool[cursors[bucket] % len(pool)]) + cursors[bucket] += 1 + assignment.append(chosen) + return assignment + + +def scenario_phases( + scenario: str, + *, + updates: int, + calibration: dict[str, Any], +) -> list[dict[str, Any]]: + """Phase list for a scenario from calibration-measured facts only. + + ``calibration`` carries ``mixes`` (per-scenario bucket shares measured from + real outputs), ``payback_updates`` (measured amortization window) and, for + tool-wait, the frozen environment contract. Nothing here decides GBS or + optimizer steps. + """ + if updates < 1: + raise ManifestError("scenario needs at least one update") + builders = { + "stable": _stable, + "phased": _phased, + "tail": _tail, + "tool-wait": _tool_wait, + "oscillating": _oscillating, + } + if scenario not in builders: + raise ManifestError(f"unknown scenario {scenario!r}") + return builders[scenario](updates, calibration) + + +def _mix_for(calibration: dict[str, Any], name: str) -> dict[str, float]: + mixes = calibration.get("mixes") or {} + mix = mixes.get(name) + if mix is None: + raise ManifestError(f"calibration provides no measured mix {name!r}") + return _validated_mix(mix, name) + + +def _stable(updates: int, calibration: dict[str, Any]) -> list[dict[str, Any]]: + return [{"name": "stable", "updates": updates, "mix": _mix_for(calibration, "stable")}] + + +def _phased(updates: int, calibration: dict[str, Any]) -> list[dict[str, Any]]: + if updates < 3: + raise ManifestError("phased scenario needs at least three updates for A->B->A") + a, b = _mix_for(calibration, "short-heavy"), _mix_for(calibration, "long-heavy") + third = updates // 3 + return [ + {"name": "A", "updates": third, "mix": a}, + {"name": "B", "updates": updates - 2 * third, "mix": b}, + {"name": "A2", "updates": third, "mix": a}, + ] + + +def _tail(updates: int, calibration: dict[str, Any]) -> list[dict[str, Any]]: + mix = _mix_for(calibration, "tail") + if mix.get("long", 0) <= 0 or mix.get("long", 0) >= 0.5: + raise ManifestError("tail scenario needs a measured minority of long prompts") + return [{"name": "tail", "updates": updates, "mix": mix}] + + +def _tool_wait(updates: int, calibration: dict[str, Any]) -> list[dict[str, Any]]: + environment = calibration.get("environment") + required = ("task_pack", "tool_concurrency", "response_contract") + if not isinstance(environment, dict) or any(not environment.get(k) for k in required): + raise ManifestError(f"tool-wait scenario needs calibration.environment with {required}") + return [ + { + "name": "tool-wait", + "updates": updates, + "mix": _mix_for(calibration, "tool-wait"), + "environment": dict(environment), + } + ] + + +def _oscillating(updates: int, calibration: dict[str, Any]) -> list[dict[str, Any]]: + payback = calibration.get("payback_updates") + if not isinstance(payback, int) or payback < 1: + raise ManifestError("oscillating scenario needs a measured calibration.payback_updates") + length = max(1, payback // 2) + if length >= payback: + raise ManifestError("oscillating phase length must be shorter than the payback window") + if updates < 2 * length: + raise ManifestError("oscillating scenario needs at least two phases") + a, b = _mix_for(calibration, "short-heavy"), _mix_for(calibration, "long-heavy") + phases = [] + remaining, index = updates, 0 + while remaining > 0: + count = min(length, remaining) + phases.append({"name": f"{'A' if index % 2 == 0 else 'B'}{index}", "updates": count, "mix": a if index % 2 == 0 else b}) + remaining -= count + index += 1 + return phases + + +def length_diagnostics(samples: list[dict[str, Any]], *, max_response_len: int) -> dict[str, Any]: + """Real length, cap-hit and zero-advantage facts from captured samples.""" + lengths = [int(s["response_length"]) for s in samples] + capped = sum(1 for s in samples if s.get("status") == "truncated" or int(s["response_length"]) >= max_response_len) + groups: dict[Any, list[float]] = {} + for sample in samples: + groups.setdefault(sample.get("group"), []).append(float(sample.get("reward", 0.0))) + zero_advantage = sum(1 for rewards in groups.values() if len(set(rewards)) <= 1) + return { + "samples": len(samples), + "mean_response_length": (sum(lengths) / len(lengths)) if lengths else None, + "max_response_length": max(lengths) if lengths else None, + "cap_hit_ratio": (capped / len(samples)) if samples else None, + "groups": len(groups), + "zero_advantage_ratio": (zero_advantage / len(groups)) if groups else None, + }