Download tests/test_coach_api.py from PYTHAI/mindXtrain: direct link, hf CLI and curl.
- Browser
- Download file 7.33 kB
-
https://huggingface.co/PYTHAI/mindXtrain/resolve/main/tests/test_coach_api.py
- Command line
-
hf download hf://PYTHAI/mindXtrain/tests/test_coach_api.py
-
curl -L -o test_coach_api.py https://huggingface.co/PYTHAI/mindXtrain/resolve/main/tests/test_coach_api.py
7.33 kB
| """Coach API: recipes / bench / compile / cost / health.""" | |
| from __future__ import annotations | |
| from fastapi.testclient import TestClient | |
| from mindxtrain.operator.app import app | |
| client = TestClient(app) | |
| def test_root_redirects_to_coach(): | |
| r = client.get("/", follow_redirects=False) | |
| assert r.status_code in (302, 307) | |
| assert r.headers["location"] == "/coach/" | |
| def test_coach_index_serves_html(): | |
| r = client.get("/coach/") | |
| assert r.status_code == 200 | |
| assert "<title>mindXtrain" in r.text | |
| assert "/coach/static/coach.js" in r.text | |
| def test_coach_static_files_served(): | |
| r_css = client.get("/coach/static/style.css") | |
| r_js = client.get("/coach/static/coach.js") | |
| assert r_css.status_code == 200 | |
| assert r_js.status_code == 200 | |
| assert "AMD orange" in r_css.text or "--accent" in r_css.text | |
| assert "loadRecipes" in r_js.text | |
| def test_coach_index_includes_receipt_card(): | |
| r = client.get("/coach/") | |
| assert r.status_code == 200 | |
| assert 'id="step-receipt"' in r.text | |
| assert 'id="receipt-badge"' in r.text | |
| # Chat re-probe affordance. | |
| assert 'id="chat-recheck"' in r.text | |
| def test_coach_js_wires_receipt_loader(): | |
| r = client.get("/coach/static/coach.js") | |
| assert r.status_code == 200 | |
| assert "loadReceiptForRun" in r.text | |
| assert "/coach/api/receipt/" in r.text | |
| def test_coach_train_diagnostics_accordions(): | |
| r = client.get("/coach/") | |
| assert r.status_code == 200 | |
| # Per-step metrics table is now a collapsible accordion with a live count. | |
| assert 'id="metrics-table-wrap"' in r.text | |
| assert 'id="metrics-table-count"' in r.text | |
| assert 'id="train-log-count"' in r.text | |
| assert 'id="chart-window-note"' in r.text | |
| def test_coach_js_caps_and_counts(): | |
| r = client.get("/coach/static/coach.js") | |
| assert r.status_code == 200 | |
| # Honest compression: rolling chart window + per-step/line counters. | |
| assert "MAX_CHART_POINTS" in r.text | |
| assert "_updateMetricsTableCount" in r.text | |
| assert "_updateLogCount" in r.text | |
| def test_coach_index_has_create_script_card(): | |
| r = client.get("/coach/") | |
| assert r.status_code == 200 | |
| assert 'id="step-create-dataset"' in r.text | |
| assert 'id="ds-save"' in r.text | |
| assert 'id="ds-exchanges"' in r.text | |
| def test_coach_js_wires_create_dataset(): | |
| r = client.get("/coach/static/coach.js") | |
| assert r.status_code == 200 | |
| assert "wireCreateDataset" in r.text | |
| assert "/coach/api/datasets" in r.text | |
| def test_recipes_list_includes_known_recipes(): | |
| r = client.get("/coach/api/recipes") | |
| assert r.status_code == 200 | |
| items = r.json() | |
| assert len(items) >= 12 | |
| names = {item["name"] for item in items} | |
| assert "qwen3_8b_sft_lora" in names | |
| assert "instella_3b_lora" in names | |
| assert "mindx_fallback_qwen3_1_5b_sft_lora" in names | |
| assert "mindx_fallback_qwen3_1_5b_cpu_smoke" in names | |
| for item in items: | |
| assert "base_model" in item | |
| assert "method" in item | |
| assert "gpus" in item | |
| def test_recipe_detail_returns_yaml_and_summary(): | |
| r = client.get("/coach/api/recipes/qwen3_8b_sft_lora") | |
| assert r.status_code == 200 | |
| data = r.json() | |
| assert "yaml" in data | |
| assert "Qwen/Qwen3-8B" in data["yaml"] | |
| assert data["summary"]["base_model"] == "Qwen/Qwen3-8B" | |
| assert data["summary"]["method"] == "lora" | |
| def test_recipe_detail_404_for_unknown(): | |
| r = client.get("/coach/api/recipes/does_not_exist") | |
| assert r.status_code == 404 | |
| def test_bench_returns_autotune_plan(): | |
| r = client.post("/coach/api/bench") | |
| assert r.status_code == 200 | |
| plan = r.json() | |
| assert plan["schema_version"] == "1" | |
| assert plan["gpu_arch"] == "gfx942" | |
| assert plan["attention_backend"] in ("ck", "triton") | |
| def test_compile_returns_axolotl_yaml_and_overrides(): | |
| r = client.post( | |
| "/coach/api/compile", | |
| json={"recipe": "qwen3_8b_sft_lora"}, | |
| ) | |
| assert r.status_code == 200 | |
| data = r.json() | |
| assert data["recipe"] == "qwen3_8b_sft_lora" | |
| assert data["axolotl_yaml"]["base_model"] == "Qwen/Qwen3-8B" | |
| assert data["axolotl_yaml"]["adapter"] == "lora" | |
| assert any("attention_backend" in o for o in data["overrides"]) | |
| def test_compile_404_for_unknown_recipe(): | |
| r = client.post("/coach/api/compile", json={"recipe": "ghost"}) | |
| assert r.status_code == 404 | |
| def test_cost_returns_three_breakdowns(): | |
| r = client.post("/coach/api/cost", json={"gpus": 1, "hours": 1.5}) | |
| assert r.status_code == 200 | |
| data = r.json() | |
| for key in ("mi300x", "h100", "h200"): | |
| assert key in data | |
| assert data[key]["cost_usdc"] > 0 | |
| # MI300X must come out cheapest in this configuration. | |
| assert data["mi300x"]["cost_usdc"] < data["h100"]["cost_usdc"] | |
| assert data["speedup_vs_h100_x"] > 1.0 | |
| assert data["mi300x"]["fits_qwen3_8b_bf16_bs8_seq4096"] is True | |
| assert data["h100"]["fits_qwen3_8b_bf16_bs8_seq4096"] is False | |
| def test_cost_validates_input(): | |
| r = client.post("/coach/api/cost", json={"gpus": 0, "hours": 1.5}) | |
| assert r.status_code == 422 | |
| def test_cost_calculator_generalizes_and_includes_a100(): | |
| # Default 8B full FT: needs ~137 GB → MI300X/H200 fit, H100/A100 don't. | |
| r = client.post("/coach/api/cost", json={"gpus": 1, "hours": 1.0}) | |
| assert r.status_code == 200, r.text | |
| d = r.json() | |
| assert "a100" in d | |
| assert d["needed_vram_gb"] > 100 | |
| assert d["mi300x"]["fits_qwen3_8b_bf16_bs8_seq4096"] is True | |
| assert d["a100"]["fits_qwen3_8b_bf16_bs8_seq4096"] is False | |
| assert len(d["comparisons"]) == 4 | |
| assert d["cheapest_that_fits"] | |
| # A tiny LoRA workload fits everywhere. | |
| r2 = client.post("/coach/api/cost", json={ | |
| "gpus": 1, "hours": 1.0, "params_b": 0.135, "method": "lora", "seq_len": 256, "batch": 1, | |
| }) | |
| d2 = r2.json() | |
| assert d2["a100"]["fits_qwen3_8b_bf16_bs8_seq4096"] is True | |
| assert d2["needed_vram_gb"] < d["needed_vram_gb"] | |
| def test_cost_card_hidden_recipe_default_present(): | |
| html = client.get("/coach/").text | |
| # Cost card is kept in the background but not displayed. | |
| assert 'id="step-cost" class="card" data-step-id="step-cost" hidden' in html | |
| # Recipe picker shows a default + an accordion of the rest. | |
| assert 'id="recipe-default"' in html | |
| assert 'id="recipe-more"' in html | |
| js = client.get("/coach/static/coach.js").text | |
| assert "renderDefaultRecipe" in js | |
| assert "DEFAULT_RECIPE" in js | |
| def test_health_endpoint_reports_recipes_count(monkeypatch): | |
| # Force the auto-detect probe off so the legacy "no live backend" shape | |
| # holds regardless of whether ollama happens to be running on the host | |
| # executing the suite. Backend-ready specifics are covered by | |
| # tests/test_ollama_backend.py. | |
| from mindxtrain.operator import app as operator_app | |
| monkeypatch.delenv("MINDXTRAIN_BACKEND", raising=False) | |
| monkeypatch.delenv("AUTOMINDX_BACKEND", raising=False) | |
| monkeypatch.setattr(operator_app, "_ollama_reachable", lambda: False) | |
| r = client.get("/coach/api/health") | |
| assert r.status_code == 200 | |
| data = r.json() | |
| assert data["recipes_available"] >= 12 | |
| assert data["chat_backend_ready"] is False | |
| def test_app_health_mentions_coach_url(): | |
| r = client.get("/health") | |
| assert r.status_code == 200 | |
| assert r.json()["coach_url"] == "/coach/" | |