Spaces:
Sleeping
Sleeping
zofiasmolenasana commited on
fix(vertex): synchronous CustomJob.submit; train UI GCP arch sync
Browse files- Replace CustomJob.run(sync=False) with submit() so create_custom_job
finishes before the HTTP response (avoids phantom submitted_count).
- Return job_resource_names for verification; log each created job.
- Include Vertex launch module and related training/dashboard updates.
Made-with: Cursor
- Dockerfile.training +7 -2
- _sync_labels.py +8 -22
- app.py +273 -37
- config.py +45 -0
- embed_text.py +53 -2
- eval_rag.py +210 -10
- gcp_progress_reader.py +3 -0
- metadata_client.py +5 -0
- rag_structure_rag_analysis.py +18 -0
- requirements-app.txt +1 -0
- requirements.txt +1 -0
- run_ablation_sweep.py +3 -1
- scripts/launch_vertex_training.py +69 -68
- static/rag_dashboard.html +12 -3
- static/train.html +347 -12
- train_compare.py +48 -7
- train_gnn.py +39 -2
- vertex_launch.py +153 -0
Dockerfile.training
CHANGED
|
@@ -1,8 +1,13 @@
|
|
| 1 |
# GPU training image for Vertex AI / local CUDA (PyTorch + PyTorch Geometric).
|
| 2 |
-
#
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
# Run (CPU smoke): docker run --rm excel-chunker-train python train_compare.py --help
|
| 4 |
|
| 5 |
-
FROM pytorch/pytorch:2.3.1-cuda12.1-
|
| 6 |
|
| 7 |
WORKDIR /app
|
| 8 |
|
|
|
|
| 1 |
# GPU training image for Vertex AI / local CUDA (PyTorch + PyTorch Geometric).
|
| 2 |
+
#
|
| 3 |
+
# Base tag must exist on Docker Hub (e.g. cudnn8-runtime, not cudnn9 unless published).
|
| 4 |
+
# Push to Artifact Registry (avoids docker-credential-gcloud PATH issues on Mac):
|
| 5 |
+
# ./scripts/push_training_image.sh
|
| 6 |
+
#
|
| 7 |
+
# Local build (no push): docker build -f Dockerfile.training -t excel-chunker-train .
|
| 8 |
# Run (CPU smoke): docker run --rm excel-chunker-train python train_compare.py --help
|
| 9 |
|
| 10 |
+
FROM pytorch/pytorch:2.3.1-cuda12.1-cudnn8-runtime
|
| 11 |
|
| 12 |
WORKDIR /app
|
| 13 |
|
_sync_labels.py
CHANGED
|
@@ -1,24 +1,10 @@
|
|
| 1 |
-
|
| 2 |
-
|
| 3 |
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
| 7 |
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
sid = r["drive_file_id"]
|
| 12 |
-
sn = r["sheet_name"]
|
| 13 |
-
safe = f"{sid}_{sn}".replace("/", "_").replace(" ", "_") + ".json"
|
| 14 |
-
local = config.LABELED_DIR / safe
|
| 15 |
-
try:
|
| 16 |
-
drive_client.download_file(r["labels_file_id"], str(local))
|
| 17 |
-
downloaded += 1
|
| 18 |
-
if downloaded % 20 == 0:
|
| 19 |
-
print(f" ...downloaded {downloaded}/{len(labelled)}")
|
| 20 |
-
except Exception as e:
|
| 21 |
-
print(f" FAILED: {safe}: {e}")
|
| 22 |
-
failed += 1
|
| 23 |
-
|
| 24 |
-
print(f"Done: {downloaded} downloaded, {failed} failed -> {config.LABELED_DIR}")
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""Backward-compatible entrypoint — use ``scripts/sync_labels_from_drive.py``."""
|
| 3 |
|
| 4 |
+
import subprocess
|
| 5 |
+
import sys
|
| 6 |
+
from pathlib import Path
|
| 7 |
|
| 8 |
+
root = Path(__file__).resolve().parent
|
| 9 |
+
script = root / "scripts" / "sync_labels_from_drive.py"
|
| 10 |
+
raise SystemExit(subprocess.call([sys.executable, str(script)] + sys.argv[1:]))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
app.py
CHANGED
|
@@ -23,7 +23,7 @@ from typing import Any, Optional
|
|
| 23 |
from fastapi import FastAPI, HTTPException, Query, Depends, Header, Request, Response
|
| 24 |
from fastapi.staticfiles import StaticFiles
|
| 25 |
from fastapi.responses import FileResponse, JSONResponse
|
| 26 |
-
from pydantic import BaseModel
|
| 27 |
|
| 28 |
import config
|
| 29 |
from training_progress import SuiteProgress
|
|
@@ -123,6 +123,13 @@ def _ensure_xlsx(drive_file_id: str) -> str:
|
|
| 123 |
|
| 124 |
app = FastAPI(title="Spreadsheet Labeler", lifespan=_app_lifespan)
|
| 125 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 126 |
_claim_lock = threading.Lock()
|
| 127 |
_labelled_count_floor = 0
|
| 128 |
_prediction_cache: dict[str, list[dict]] = {}
|
|
@@ -433,7 +440,7 @@ def save_and_next(payload: SavePayload, _=Depends(_require_auth)):
|
|
| 433 |
"cells": [c.model_dump() for c in payload.cells],
|
| 434 |
}
|
| 435 |
|
| 436 |
-
safe_name =
|
| 437 |
json_bytes = json.dumps(labeled_data, indent=2, ensure_ascii=False).encode("utf-8")
|
| 438 |
|
| 439 |
# Save locally
|
|
@@ -601,7 +608,7 @@ def save_review(payload: SavePayload, _=Depends(_require_auth)):
|
|
| 601 |
"cells": [c.model_dump() for c in payload.cells],
|
| 602 |
}
|
| 603 |
|
| 604 |
-
safe_name =
|
| 605 |
json_bytes = json.dumps(labeled_data, indent=2, ensure_ascii=False).encode("utf-8")
|
| 606 |
|
| 607 |
local_label_path = config.LABELED_DIR / f"{safe_name}.json"
|
|
@@ -652,7 +659,7 @@ def backfill_review_labels(_=Depends(_require_auth)):
|
|
| 652 |
for r in rows:
|
| 653 |
sid = r.get("drive_file_id", "")
|
| 654 |
sn = r.get("sheet_name", "")
|
| 655 |
-
safe =
|
| 656 |
row_by_key[safe] = r
|
| 657 |
|
| 658 |
updated = 0
|
|
@@ -755,7 +762,7 @@ def _load_saved_labels(
|
|
| 755 |
labels_file_id: str = "",
|
| 756 |
) -> list[dict] | None:
|
| 757 |
"""Load previously saved labels from local JSON or Drive."""
|
| 758 |
-
safe_name =
|
| 759 |
local_path = config.LABELED_DIR / f"{safe_name}.json"
|
| 760 |
|
| 761 |
data = None
|
|
@@ -1063,25 +1070,55 @@ _MAX_TRAIN_LOG_LINES = 200
|
|
| 1063 |
_GCP_SUITE_CACHE: dict[str, Any] = {
|
| 1064 |
"uri": "",
|
| 1065 |
"ts": 0.0,
|
|
|
|
| 1066 |
"data": None,
|
| 1067 |
"err": None,
|
| 1068 |
}
|
| 1069 |
_GCP_SUITE_LOCK = threading.Lock()
|
| 1070 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1071 |
|
| 1072 |
|
| 1073 |
-
def _get_gcp_suite_cached() -> tuple[dict[str, Any] | None, str | None]:
|
| 1074 |
-
"""Merge remote Vertex progress from ``…/experiments/*/progress.json`` on GCS (cached).
|
| 1075 |
-
|
|
|
|
|
|
|
|
|
|
| 1076 |
if not uri:
|
| 1077 |
-
return None, None
|
| 1078 |
now = time.monotonic()
|
| 1079 |
with _GCP_SUITE_LOCK:
|
| 1080 |
if (
|
| 1081 |
_GCP_SUITE_CACHE["uri"] == uri
|
| 1082 |
and now - float(_GCP_SUITE_CACHE["ts"]) < _GCP_SUITE_TTL_SEC
|
| 1083 |
):
|
| 1084 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1085 |
err: str | None = None
|
| 1086 |
data: dict[str, Any] | None = None
|
| 1087 |
try:
|
|
@@ -1091,12 +1128,18 @@ def _get_gcp_suite_cached() -> tuple[dict[str, Any] | None, str | None]:
|
|
| 1091 |
except Exception as e:
|
| 1092 |
err = str(e)[:800]
|
| 1093 |
logger.warning("GCS training progress fetch failed: %s", e)
|
|
|
|
| 1094 |
with _GCP_SUITE_LOCK:
|
| 1095 |
_GCP_SUITE_CACHE["uri"] = uri
|
| 1096 |
_GCP_SUITE_CACHE["ts"] = now
|
| 1097 |
_GCP_SUITE_CACHE["data"] = data
|
| 1098 |
_GCP_SUITE_CACHE["err"] = err
|
| 1099 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1100 |
|
| 1101 |
|
| 1102 |
@app.get("/api/train/dataset_stats")
|
|
@@ -1478,15 +1521,173 @@ def experiment_detail(run: str, _=Depends(_require_auth)):
|
|
| 1478 |
@app.get("/api/train/status")
|
| 1479 |
def train_status():
|
| 1480 |
"""Return current training status, log, and per-unit progress (local and/or GCS)."""
|
| 1481 |
-
gcp_uri =
|
| 1482 |
-
gcp_suite, gcp_err = _get_gcp_suite_cached()
|
| 1483 |
-
|
|
|
|
|
|
|
| 1484 |
**_train_status,
|
| 1485 |
"suite": _train_suite.to_api_dict(),
|
| 1486 |
"gcp_progress_enabled": bool(gcp_uri),
|
| 1487 |
"gcp_suite": gcp_suite,
|
| 1488 |
"gcp_suite_error": gcp_err,
|
| 1489 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1490 |
|
| 1491 |
|
| 1492 |
class TrainRequest(BaseModel):
|
|
@@ -1502,6 +1703,13 @@ class TrainRequest(BaseModel):
|
|
| 1502 |
clear_previous: bool = False
|
| 1503 |
no_va_margin_loss: bool = False
|
| 1504 |
no_col_consistency_loss: bool = False
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1505 |
|
| 1506 |
|
| 1507 |
@app.post("/api/train/start")
|
|
@@ -1521,9 +1729,10 @@ def start_training(req: TrainRequest, _=Depends(_require_auth)):
|
|
| 1521 |
try:
|
| 1522 |
if req.mode == "compare":
|
| 1523 |
_train_suite = SuiteProgress()
|
|
|
|
| 1524 |
_train_status["log"].append(
|
| 1525 |
f"Architectures: {req.architectures}, Seeds: {req.seeds}, "
|
| 1526 |
-
f"Folds: {req.folds}, Epochs: {req.epochs}"
|
| 1527 |
)
|
| 1528 |
|
| 1529 |
if req.clear_previous:
|
|
@@ -1563,22 +1772,42 @@ def start_training(req: TrainRequest, _=Depends(_require_auth)):
|
|
| 1563 |
old_stdout = sys.stdout
|
| 1564 |
sys.stdout = LogCapture()
|
| 1565 |
try:
|
| 1566 |
-
|
|
|
|
| 1567 |
architectures=archs, seeds=seeds,
|
| 1568 |
n_folds=req.folds, num_epochs=req.epochs,
|
| 1569 |
do_ablations=req.ablations,
|
| 1570 |
use_va_margin_loss=not req.no_va_margin_loss,
|
| 1571 |
use_col_consistency_loss=not req.no_col_consistency_loss,
|
|
|
|
|
|
|
| 1572 |
progress=_train_suite,
|
| 1573 |
)
|
| 1574 |
finally:
|
| 1575 |
sys.stdout = old_stdout
|
| 1576 |
|
| 1577 |
-
|
| 1578 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1579 |
|
| 1580 |
elif req.mode == "production":
|
| 1581 |
-
|
|
|
|
|
|
|
|
|
|
| 1582 |
from train_gnn import train_model
|
| 1583 |
|
| 1584 |
class LogCapture(io.StringIO):
|
|
@@ -1590,7 +1819,13 @@ def start_training(req: TrainRequest, _=Depends(_require_auth)):
|
|
| 1590 |
old_stdout = sys.stdout
|
| 1591 |
sys.stdout = LogCapture()
|
| 1592 |
try:
|
| 1593 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1594 |
finally:
|
| 1595 |
sys.stdout = old_stdout
|
| 1596 |
|
|
@@ -1625,25 +1860,21 @@ def run_evaluation(_=Depends(_require_auth)):
|
|
| 1625 |
|
| 1626 |
@app.post("/api/train/embed")
|
| 1627 |
def run_embeddings(_=Depends(_require_auth)):
|
| 1628 |
-
"""
|
| 1629 |
-
|
| 1630 |
-
|
| 1631 |
-
|
| 1632 |
try:
|
| 1633 |
-
from embed_text import
|
| 1634 |
-
|
| 1635 |
-
|
| 1636 |
-
if embed_path.exists():
|
| 1637 |
-
skipped += 1
|
| 1638 |
-
continue
|
| 1639 |
-
try:
|
| 1640 |
-
embed_file(fp)
|
| 1641 |
-
count += 1
|
| 1642 |
-
except Exception as e:
|
| 1643 |
-
errors.append(f"{fp.name}: {e}")
|
| 1644 |
except ImportError:
|
| 1645 |
return {"error": "embed_text module not available"}
|
| 1646 |
-
return {
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1647 |
|
| 1648 |
|
| 1649 |
# ── RAG Evaluation ─────────────────────────────────────────────────────────
|
|
@@ -2066,6 +2297,10 @@ def rag_eval_run(
|
|
| 2066 |
resume_structure: bool = Query(True),
|
| 2067 |
force_structure_metrics: bool = Query(False),
|
| 2068 |
save_structure_details: bool = Query(False),
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2069 |
_=Depends(_require_auth),
|
| 2070 |
):
|
| 2071 |
"""Trigger a RAG evaluation run in the background.
|
|
@@ -2100,6 +2335,7 @@ def rag_eval_run(
|
|
| 2100 |
resume_structure=resume_structure,
|
| 2101 |
force_structure_metrics=force_structure_metrics,
|
| 2102 |
save_structure_details=save_structure_details,
|
|
|
|
| 2103 |
)
|
| 2104 |
_rag_eval_status["result"] = result
|
| 2105 |
_rag_eval_status["log"].append("RAG evaluation completed.")
|
|
|
|
| 23 |
from fastapi import FastAPI, HTTPException, Query, Depends, Header, Request, Response
|
| 24 |
from fastapi.staticfiles import StaticFiles
|
| 25 |
from fastapi.responses import FileResponse, JSONResponse
|
| 26 |
+
from pydantic import BaseModel, Field
|
| 27 |
|
| 28 |
import config
|
| 29 |
from training_progress import SuiteProgress
|
|
|
|
| 123 |
|
| 124 |
app = FastAPI(title="Spreadsheet Labeler", lifespan=_app_lifespan)
|
| 125 |
|
| 126 |
+
|
| 127 |
+
@app.get("/favicon.ico", include_in_schema=False)
|
| 128 |
+
def favicon_ico():
|
| 129 |
+
"""Browsers request this by default; avoid noisy 404 in devtools."""
|
| 130 |
+
return Response(status_code=204)
|
| 131 |
+
|
| 132 |
+
|
| 133 |
_claim_lock = threading.Lock()
|
| 134 |
_labelled_count_floor = 0
|
| 135 |
_prediction_cache: dict[str, list[dict]] = {}
|
|
|
|
| 440 |
"cells": [c.model_dump() for c in payload.cells],
|
| 441 |
}
|
| 442 |
|
| 443 |
+
safe_name = metadata_client.label_file_stem(payload.spreadsheet_id, payload.sheet_name)
|
| 444 |
json_bytes = json.dumps(labeled_data, indent=2, ensure_ascii=False).encode("utf-8")
|
| 445 |
|
| 446 |
# Save locally
|
|
|
|
| 608 |
"cells": [c.model_dump() for c in payload.cells],
|
| 609 |
}
|
| 610 |
|
| 611 |
+
safe_name = metadata_client.label_file_stem(payload.spreadsheet_id, payload.sheet_name)
|
| 612 |
json_bytes = json.dumps(labeled_data, indent=2, ensure_ascii=False).encode("utf-8")
|
| 613 |
|
| 614 |
local_label_path = config.LABELED_DIR / f"{safe_name}.json"
|
|
|
|
| 659 |
for r in rows:
|
| 660 |
sid = r.get("drive_file_id", "")
|
| 661 |
sn = r.get("sheet_name", "")
|
| 662 |
+
safe = metadata_client.label_file_stem(sid, sn) + ".json"
|
| 663 |
row_by_key[safe] = r
|
| 664 |
|
| 665 |
updated = 0
|
|
|
|
| 762 |
labels_file_id: str = "",
|
| 763 |
) -> list[dict] | None:
|
| 764 |
"""Load previously saved labels from local JSON or Drive."""
|
| 765 |
+
safe_name = metadata_client.label_file_stem(spreadsheet_id, sheet_name)
|
| 766 |
local_path = config.LABELED_DIR / f"{safe_name}.json"
|
| 767 |
|
| 768 |
data = None
|
|
|
|
| 1070 |
_GCP_SUITE_CACHE: dict[str, Any] = {
|
| 1071 |
"uri": "",
|
| 1072 |
"ts": 0.0,
|
| 1073 |
+
"wall_ts": 0.0,
|
| 1074 |
"data": None,
|
| 1075 |
"err": None,
|
| 1076 |
}
|
| 1077 |
_GCP_SUITE_LOCK = threading.Lock()
|
| 1078 |
+
# Short TTL: dashboard polls ~2s; keeps GCS list fresh without hammering the API.
|
| 1079 |
+
_GCP_SUITE_TTL_SEC = 2.0
|
| 1080 |
+
|
| 1081 |
+
|
| 1082 |
+
def _invalidate_gcp_suite_cache() -> None:
|
| 1083 |
+
"""Force the next /api/train/status to refetch GCS (e.g. after launching new Vertex jobs)."""
|
| 1084 |
+
with _GCP_SUITE_LOCK:
|
| 1085 |
+
_GCP_SUITE_CACHE["ts"] = 0.0
|
| 1086 |
+
|
| 1087 |
+
|
| 1088 |
+
def _gcp_wall_ts_iso(wall_ts: float) -> str | None:
|
| 1089 |
+
if not wall_ts or wall_ts <= 0:
|
| 1090 |
+
return None
|
| 1091 |
+
return datetime.fromtimestamp(wall_ts, tz=timezone.utc).isoformat(timespec="seconds")
|
| 1092 |
+
|
| 1093 |
+
_docker_image_lock = threading.Lock()
|
| 1094 |
+
_docker_image_status: dict[str, Any] = {
|
| 1095 |
+
"running": False,
|
| 1096 |
+
"log": [],
|
| 1097 |
+
"ok": None,
|
| 1098 |
+
"error": None,
|
| 1099 |
+
}
|
| 1100 |
|
| 1101 |
|
| 1102 |
+
def _get_gcp_suite_cached() -> tuple[dict[str, Any] | None, str | None, str | None]:
|
| 1103 |
+
"""Merge remote Vertex progress from ``…/experiments/*/progress.json`` on GCS (cached).
|
| 1104 |
+
|
| 1105 |
+
Returns ``(suite_dict, error, refreshed_at_iso)`` for the train dashboard.
|
| 1106 |
+
"""
|
| 1107 |
+
uri = config.gcp_experiments_uri_for_progress()
|
| 1108 |
if not uri:
|
| 1109 |
+
return None, None, None
|
| 1110 |
now = time.monotonic()
|
| 1111 |
with _GCP_SUITE_LOCK:
|
| 1112 |
if (
|
| 1113 |
_GCP_SUITE_CACHE["uri"] == uri
|
| 1114 |
and now - float(_GCP_SUITE_CACHE["ts"]) < _GCP_SUITE_TTL_SEC
|
| 1115 |
):
|
| 1116 |
+
wt = float(_GCP_SUITE_CACHE.get("wall_ts") or 0)
|
| 1117 |
+
return (
|
| 1118 |
+
_GCP_SUITE_CACHE["data"],
|
| 1119 |
+
_GCP_SUITE_CACHE["err"],
|
| 1120 |
+
_gcp_wall_ts_iso(wt),
|
| 1121 |
+
)
|
| 1122 |
err: str | None = None
|
| 1123 |
data: dict[str, Any] | None = None
|
| 1124 |
try:
|
|
|
|
| 1128 |
except Exception as e:
|
| 1129 |
err = str(e)[:800]
|
| 1130 |
logger.warning("GCS training progress fetch failed: %s", e)
|
| 1131 |
+
refreshed_iso: str | None = None
|
| 1132 |
with _GCP_SUITE_LOCK:
|
| 1133 |
_GCP_SUITE_CACHE["uri"] = uri
|
| 1134 |
_GCP_SUITE_CACHE["ts"] = now
|
| 1135 |
_GCP_SUITE_CACHE["data"] = data
|
| 1136 |
_GCP_SUITE_CACHE["err"] = err
|
| 1137 |
+
if data is not None:
|
| 1138 |
+
_GCP_SUITE_CACHE["wall_ts"] = time.time()
|
| 1139 |
+
refreshed_iso = _gcp_wall_ts_iso(float(_GCP_SUITE_CACHE["wall_ts"]))
|
| 1140 |
+
else:
|
| 1141 |
+
refreshed_iso = _gcp_wall_ts_iso(float(_GCP_SUITE_CACHE.get("wall_ts") or 0))
|
| 1142 |
+
return data, err, refreshed_iso
|
| 1143 |
|
| 1144 |
|
| 1145 |
@app.get("/api/train/dataset_stats")
|
|
|
|
| 1521 |
@app.get("/api/train/status")
|
| 1522 |
def train_status():
|
| 1523 |
"""Return current training status, log, and per-unit progress (local and/or GCS)."""
|
| 1524 |
+
gcp_uri = config.gcp_experiments_uri_for_progress()
|
| 1525 |
+
gcp_suite, gcp_err, gcp_refreshed_at = _get_gcp_suite_cached()
|
| 1526 |
+
if gcp_suite is not None and gcp_refreshed_at:
|
| 1527 |
+
gcp_suite = {**gcp_suite, "refreshed_at": gcp_refreshed_at}
|
| 1528 |
+
body = {
|
| 1529 |
**_train_status,
|
| 1530 |
"suite": _train_suite.to_api_dict(),
|
| 1531 |
"gcp_progress_enabled": bool(gcp_uri),
|
| 1532 |
"gcp_suite": gcp_suite,
|
| 1533 |
"gcp_suite_error": gcp_err,
|
| 1534 |
}
|
| 1535 |
+
return JSONResponse(
|
| 1536 |
+
content=body,
|
| 1537 |
+
headers={
|
| 1538 |
+
"Cache-Control": "no-store, no-cache",
|
| 1539 |
+
"Pragma": "no-cache",
|
| 1540 |
+
},
|
| 1541 |
+
)
|
| 1542 |
+
|
| 1543 |
+
|
| 1544 |
+
@app.get("/api/train/gcp-defaults")
|
| 1545 |
+
def train_gcp_defaults(_=Depends(_require_auth)):
|
| 1546 |
+
"""Default GCP/Vertex values for the dashboard (from config / .env)."""
|
| 1547 |
+
return {
|
| 1548 |
+
"project_id": config.GCP_PROJECT_ID,
|
| 1549 |
+
"region": config.GCP_REGION,
|
| 1550 |
+
"gcs_data_uri": config.GCP_GCS_DATA_URI,
|
| 1551 |
+
"gcs_experiments_uri": config.gcp_experiments_uri_for_progress(),
|
| 1552 |
+
"artifact_registry_repo": config.GCP_ARTIFACT_REGISTRY_REPO,
|
| 1553 |
+
"image_uri": config.gcp_training_image_uri(),
|
| 1554 |
+
"training_image_name": config.GCP_TRAINING_IMAGE_NAME,
|
| 1555 |
+
"training_image_tag": config.GCP_TRAINING_IMAGE_TAG,
|
| 1556 |
+
"machine_type": config.GCP_VERTEX_MACHINE_TYPE,
|
| 1557 |
+
"accelerator_type": config.GCP_VERTEX_ACCELERATOR_TYPE,
|
| 1558 |
+
"accelerator_count": config.GCP_VERTEX_ACCELERATOR_COUNT,
|
| 1559 |
+
"vertex_spot": False,
|
| 1560 |
+
"docker_build_from_dashboard": config.docker_build_from_dashboard_enabled(),
|
| 1561 |
+
}
|
| 1562 |
+
|
| 1563 |
+
|
| 1564 |
+
@app.post("/api/train/gcp-image-build")
|
| 1565 |
+
def start_gcp_image_build(_=Depends(_require_auth)):
|
| 1566 |
+
"""Run ``docker build`` + ``docker push`` for the training image (local machine only)."""
|
| 1567 |
+
if not config.docker_build_from_dashboard_enabled():
|
| 1568 |
+
raise HTTPException(
|
| 1569 |
+
status_code=403,
|
| 1570 |
+
detail="Włącz ENABLE_DOCKER_BUILD_FROM_DASHBOARD=1 tylko lokalnie (wymaga Dockera).",
|
| 1571 |
+
)
|
| 1572 |
+
with _docker_image_lock:
|
| 1573 |
+
if _docker_image_status["running"]:
|
| 1574 |
+
raise HTTPException(status_code=409, detail="Build już trwa.")
|
| 1575 |
+
_docker_image_status["running"] = True
|
| 1576 |
+
_docker_image_status["log"] = []
|
| 1577 |
+
_docker_image_status["ok"] = None
|
| 1578 |
+
_docker_image_status["error"] = None
|
| 1579 |
+
|
| 1580 |
+
def _run() -> None:
|
| 1581 |
+
def append(line: str) -> None:
|
| 1582 |
+
line = line.rstrip()
|
| 1583 |
+
if not line:
|
| 1584 |
+
return
|
| 1585 |
+
with _docker_image_lock:
|
| 1586 |
+
log = _docker_image_status["log"]
|
| 1587 |
+
log.append(line)
|
| 1588 |
+
if len(log) > 500:
|
| 1589 |
+
del log[: len(log) - 500]
|
| 1590 |
+
|
| 1591 |
+
try:
|
| 1592 |
+
from docker_image_build import run_docker_build_and_push
|
| 1593 |
+
|
| 1594 |
+
ok, msg = run_docker_build_and_push(
|
| 1595 |
+
image_uri=config.gcp_training_image_uri(),
|
| 1596 |
+
repo_root=config.BASE_DIR,
|
| 1597 |
+
log_append=append,
|
| 1598 |
+
)
|
| 1599 |
+
with _docker_image_lock:
|
| 1600 |
+
_docker_image_status["ok"] = ok
|
| 1601 |
+
_docker_image_status["error"] = None if ok else msg
|
| 1602 |
+
except Exception as e:
|
| 1603 |
+
logger.exception("docker build/push failed")
|
| 1604 |
+
with _docker_image_lock:
|
| 1605 |
+
_docker_image_status["ok"] = False
|
| 1606 |
+
_docker_image_status["error"] = str(e)
|
| 1607 |
+
finally:
|
| 1608 |
+
with _docker_image_lock:
|
| 1609 |
+
_docker_image_status["running"] = False
|
| 1610 |
+
|
| 1611 |
+
threading.Thread(target=_run, daemon=True).start()
|
| 1612 |
+
return {"started": True}
|
| 1613 |
+
|
| 1614 |
+
|
| 1615 |
+
@app.get("/api/train/gcp-image-build/status")
|
| 1616 |
+
def gcp_image_build_status(_=Depends(_require_auth)):
|
| 1617 |
+
with _docker_image_lock:
|
| 1618 |
+
return {
|
| 1619 |
+
"running": _docker_image_status["running"],
|
| 1620 |
+
"log": list(_docker_image_status["log"]),
|
| 1621 |
+
"ok": _docker_image_status["ok"],
|
| 1622 |
+
"error": _docker_image_status["error"],
|
| 1623 |
+
}
|
| 1624 |
+
|
| 1625 |
+
|
| 1626 |
+
class GcpLaunchRequest(BaseModel):
|
| 1627 |
+
architectures: str = "gat"
|
| 1628 |
+
seeds: str = "42"
|
| 1629 |
+
folds: int = 5
|
| 1630 |
+
epochs: int = 100
|
| 1631 |
+
subset_percent: float = Field(100.0, ge=1.0, le=100.0)
|
| 1632 |
+
project_id: Optional[str] = None
|
| 1633 |
+
region: Optional[str] = None
|
| 1634 |
+
image_uri: Optional[str] = None
|
| 1635 |
+
gcs_data_uri: Optional[str] = None
|
| 1636 |
+
machine_type: str = Field(default_factory=lambda: config.GCP_VERTEX_MACHINE_TYPE)
|
| 1637 |
+
accelerator_type: str = Field(default_factory=lambda: config.GCP_VERTEX_ACCELERATOR_TYPE)
|
| 1638 |
+
accelerator_count: int = Field(default_factory=lambda: config.GCP_VERTEX_ACCELERATOR_COUNT)
|
| 1639 |
+
vertex_spot: bool = False
|
| 1640 |
+
# Optional per-architecture accelerator (Vertex enum). Omitted archs use accelerator_type.
|
| 1641 |
+
arch_accelerators: Optional[dict[str, str]] = None
|
| 1642 |
+
|
| 1643 |
+
|
| 1644 |
+
@app.post("/api/train/gcp-launch")
|
| 1645 |
+
def train_gcp_launch(req: GcpLaunchRequest, _=Depends(_require_auth)):
|
| 1646 |
+
"""Submit Vertex Custom Training jobs (same logic as ``scripts/launch_vertex_training.py``)."""
|
| 1647 |
+
try:
|
| 1648 |
+
from vertex_launch import submit_vertex_training_jobs
|
| 1649 |
+
except ImportError as e:
|
| 1650 |
+
raise HTTPException(
|
| 1651 |
+
status_code=500,
|
| 1652 |
+
detail="Install Vertex deps: pip install google-cloud-aiplatform",
|
| 1653 |
+
) from e
|
| 1654 |
+
|
| 1655 |
+
project = (req.project_id or config.GCP_PROJECT_ID).strip()
|
| 1656 |
+
region = (req.region or config.GCP_REGION).strip()
|
| 1657 |
+
image = (req.image_uri or config.gcp_training_image_uri()).strip()
|
| 1658 |
+
gcs_uri = (req.gcs_data_uri or config.GCP_GCS_DATA_URI).strip()
|
| 1659 |
+
archs = [a.strip() for a in req.architectures.split(",") if a.strip()]
|
| 1660 |
+
seeds: list[int] = []
|
| 1661 |
+
for s in req.seeds.split(","):
|
| 1662 |
+
s = s.strip()
|
| 1663 |
+
if s:
|
| 1664 |
+
seeds.append(int(s))
|
| 1665 |
+
if not archs or not seeds:
|
| 1666 |
+
raise HTTPException(status_code=400, detail="architectures and seeds must be non-empty")
|
| 1667 |
+
|
| 1668 |
+
try:
|
| 1669 |
+
out = submit_vertex_training_jobs(
|
| 1670 |
+
project=project,
|
| 1671 |
+
region=region,
|
| 1672 |
+
image_uri=image,
|
| 1673 |
+
gcs_data_uri=gcs_uri,
|
| 1674 |
+
architectures=archs,
|
| 1675 |
+
seeds=seeds,
|
| 1676 |
+
folds=req.folds,
|
| 1677 |
+
epochs=req.epochs,
|
| 1678 |
+
subset_percent=req.subset_percent,
|
| 1679 |
+
machine_type=req.machine_type,
|
| 1680 |
+
accelerator_type=req.accelerator_type,
|
| 1681 |
+
accelerator_count=req.accelerator_count,
|
| 1682 |
+
use_spot=req.vertex_spot,
|
| 1683 |
+
arch_accelerator=req.arch_accelerators,
|
| 1684 |
+
)
|
| 1685 |
+
except Exception as e:
|
| 1686 |
+
logger.exception("Vertex launch failed")
|
| 1687 |
+
raise HTTPException(status_code=500, detail=str(e)) from e
|
| 1688 |
+
|
| 1689 |
+
_invalidate_gcp_suite_cache()
|
| 1690 |
+
return {"status": "ok", **out}
|
| 1691 |
|
| 1692 |
|
| 1693 |
class TrainRequest(BaseModel):
|
|
|
|
| 1703 |
clear_previous: bool = False
|
| 1704 |
no_va_margin_loss: bool = False
|
| 1705 |
no_col_consistency_loss: bool = False
|
| 1706 |
+
subset_percent: float = Field(
|
| 1707 |
+
100.0,
|
| 1708 |
+
ge=1.0,
|
| 1709 |
+
le=100.0,
|
| 1710 |
+
description="Random subset of labeled sheets before training (100 = all).",
|
| 1711 |
+
)
|
| 1712 |
+
subset_seed: int = Field(0, ge=0, description="RNG seed for subset sampling when < 100%.")
|
| 1713 |
|
| 1714 |
|
| 1715 |
@app.post("/api/train/start")
|
|
|
|
| 1729 |
try:
|
| 1730 |
if req.mode == "compare":
|
| 1731 |
_train_suite = SuiteProgress()
|
| 1732 |
+
sub = f"{req.subset_percent:g}% of sheets" if req.subset_percent < 100 else "100% of sheets"
|
| 1733 |
_train_status["log"].append(
|
| 1734 |
f"Architectures: {req.architectures}, Seeds: {req.seeds}, "
|
| 1735 |
+
f"Folds: {req.folds}, Epochs: {req.epochs}, Data: {sub} (seed {req.subset_seed})"
|
| 1736 |
)
|
| 1737 |
|
| 1738 |
if req.clear_previous:
|
|
|
|
| 1772 |
old_stdout = sys.stdout
|
| 1773 |
sys.stdout = LogCapture()
|
| 1774 |
try:
|
| 1775 |
+
sp = None if req.subset_percent >= 100.0 else req.subset_percent
|
| 1776 |
+
result, failed_units = run_experiments(
|
| 1777 |
architectures=archs, seeds=seeds,
|
| 1778 |
n_folds=req.folds, num_epochs=req.epochs,
|
| 1779 |
do_ablations=req.ablations,
|
| 1780 |
use_va_margin_loss=not req.no_va_margin_loss,
|
| 1781 |
use_col_consistency_loss=not req.no_col_consistency_loss,
|
| 1782 |
+
subset_percent=sp,
|
| 1783 |
+
subset_seed=req.subset_seed,
|
| 1784 |
progress=_train_suite,
|
| 1785 |
)
|
| 1786 |
finally:
|
| 1787 |
sys.stdout = old_stdout
|
| 1788 |
|
| 1789 |
+
if failed_units:
|
| 1790 |
+
_train_status["result"] = {
|
| 1791 |
+
"status": "failed",
|
| 1792 |
+
"failed_units": failed_units,
|
| 1793 |
+
"architectures": list(result.keys()),
|
| 1794 |
+
}
|
| 1795 |
+
_train_status["log"].append(
|
| 1796 |
+
f"Training finished with {failed_units} failed unit(s) "
|
| 1797 |
+
f"(e.g. CUDA OOM). Check log lines above."
|
| 1798 |
+
)
|
| 1799 |
+
else:
|
| 1800 |
+
_train_status["result"] = {
|
| 1801 |
+
"status": "completed",
|
| 1802 |
+
"architectures": list(result.keys()),
|
| 1803 |
+
}
|
| 1804 |
+
_train_status["log"].append("Training completed successfully!")
|
| 1805 |
|
| 1806 |
elif req.mode == "production":
|
| 1807 |
+
sub = f"{req.subset_percent:g}% of sheets" if req.subset_percent < 100 else "100% of sheets"
|
| 1808 |
+
_train_status["log"].append(
|
| 1809 |
+
f"Training production model for {req.epochs} epochs ({sub}, seed {req.subset_seed})..."
|
| 1810 |
+
)
|
| 1811 |
from train_gnn import train_model
|
| 1812 |
|
| 1813 |
class LogCapture(io.StringIO):
|
|
|
|
| 1819 |
old_stdout = sys.stdout
|
| 1820 |
sys.stdout = LogCapture()
|
| 1821 |
try:
|
| 1822 |
+
sp = None if req.subset_percent >= 100.0 else req.subset_percent
|
| 1823 |
+
train_model(
|
| 1824 |
+
num_epochs=req.epochs,
|
| 1825 |
+
skip_predict=True,
|
| 1826 |
+
subset_percent=sp,
|
| 1827 |
+
subset_seed=req.subset_seed,
|
| 1828 |
+
)
|
| 1829 |
finally:
|
| 1830 |
sys.stdout = old_stdout
|
| 1831 |
|
|
|
|
| 1860 |
|
| 1861 |
@app.post("/api/train/embed")
|
| 1862 |
def run_embeddings(_=Depends(_require_auth)):
|
| 1863 |
+
"""Optional: embed every labeled JSON (same as ``embed_text.embed_all``).
|
| 1864 |
+
|
| 1865 |
+
Training paths already call ``ensure_embeddings_for_labeled`` before loading graphs.
|
| 1866 |
+
"""
|
| 1867 |
try:
|
| 1868 |
+
from embed_text import ensure_embeddings_for_labeled
|
| 1869 |
+
|
| 1870 |
+
r = ensure_embeddings_for_labeled()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1871 |
except ImportError:
|
| 1872 |
return {"error": "embed_text module not available"}
|
| 1873 |
+
return {
|
| 1874 |
+
"embedded": r["generated"],
|
| 1875 |
+
"skipped": r["skipped"],
|
| 1876 |
+
"errors": r["errors"],
|
| 1877 |
+
}
|
| 1878 |
|
| 1879 |
|
| 1880 |
# ── RAG Evaluation ─────────────────────────────────────────────────────────
|
|
|
|
| 2297 |
resume_structure: bool = Query(True),
|
| 2298 |
force_structure_metrics: bool = Query(False),
|
| 2299 |
save_structure_details: bool = Query(False),
|
| 2300 |
+
timing_exclude_first_sheet: bool = Query(
|
| 2301 |
+
False,
|
| 2302 |
+
description="Exclude first predict_sheet from inference_timing mean/p95 (warmup)",
|
| 2303 |
+
),
|
| 2304 |
_=Depends(_require_auth),
|
| 2305 |
):
|
| 2306 |
"""Trigger a RAG evaluation run in the background.
|
|
|
|
| 2335 |
resume_structure=resume_structure,
|
| 2336 |
force_structure_metrics=force_structure_metrics,
|
| 2337 |
save_structure_details=save_structure_details,
|
| 2338 |
+
timing_exclude_first_sheet=timing_exclude_first_sheet,
|
| 2339 |
)
|
| 2340 |
_rag_eval_status["result"] = result
|
| 2341 |
_rag_eval_status["log"].append("RAG evaluation completed.")
|
config.py
CHANGED
|
@@ -81,6 +81,51 @@ EXPERIMENTS_DIR.mkdir(parents=True, exist_ok=True)
|
|
| 81 |
RAG_EVAL_CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
| 82 |
RAG_EVAL_LABELED_DIR.mkdir(parents=True, exist_ok=True)
|
| 83 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 84 |
|
| 85 |
def get_google_credentials():
|
| 86 |
"""Return service account credentials (works for Sheets; read-only Drive)."""
|
|
|
|
| 81 |
RAG_EVAL_CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
| 82 |
RAG_EVAL_LABELED_DIR.mkdir(parents=True, exist_ok=True)
|
| 83 |
|
| 84 |
+
# ── Vertex AI / GCS (defaults for dashboard; override via .env) ─────────────
|
| 85 |
+
GCP_PROJECT_ID = (os.environ.get("GCP_PROJECT_ID") or "acoustic-atom-386613").strip()
|
| 86 |
+
GCP_REGION = (os.environ.get("GCP_REGION") or "europe-west1").strip()
|
| 87 |
+
GCP_GCS_DATA_URI = (os.environ.get("GCP_GCS_DATA_URI") or "gs://excel-chunker/data").strip()
|
| 88 |
+
GCP_ARTIFACT_REGISTRY_REPO = (os.environ.get("GCP_ARTIFACT_REGISTRY_REPO") or "excel-chunker").strip()
|
| 89 |
+
GCP_TRAINING_IMAGE_NAME = (os.environ.get("GCP_TRAINING_IMAGE_NAME") or "excel-chunker-train").strip()
|
| 90 |
+
GCP_TRAINING_IMAGE_TAG = (os.environ.get("GCP_TRAINING_IMAGE_TAG") or "latest").strip()
|
| 91 |
+
GCP_TRAINING_IMAGE_URI = (os.environ.get("GCP_TRAINING_IMAGE_URI") or "").strip()
|
| 92 |
+
|
| 93 |
+
# Vertex Custom Training: VM + GPU must match (e.g. L4 → G2, not N1).
|
| 94 |
+
# Override if you only have T4 quota: GCP_VERTEX_ACCELERATOR_TYPE=NVIDIA_TESLA_T4,
|
| 95 |
+
# GCP_VERTEX_MACHINE_TYPE=n1-standard-8
|
| 96 |
+
GCP_VERTEX_MACHINE_TYPE = (os.environ.get("GCP_VERTEX_MACHINE_TYPE") or "g2-standard-8").strip()
|
| 97 |
+
GCP_VERTEX_ACCELERATOR_TYPE = (os.environ.get("GCP_VERTEX_ACCELERATOR_TYPE") or "NVIDIA_L4").strip()
|
| 98 |
+
GCP_VERTEX_ACCELERATOR_COUNT = max(
|
| 99 |
+
1,
|
| 100 |
+
int((os.environ.get("GCP_VERTEX_ACCELERATOR_COUNT") or "1").strip() or "1"),
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
def gcp_training_image_uri() -> str:
|
| 105 |
+
"""Full Artifact Registry Docker image URI for Vertex training workers."""
|
| 106 |
+
if GCP_TRAINING_IMAGE_URI:
|
| 107 |
+
return GCP_TRAINING_IMAGE_URI
|
| 108 |
+
return (
|
| 109 |
+
f"{GCP_REGION}-docker.pkg.dev/{GCP_PROJECT_ID}/"
|
| 110 |
+
f"{GCP_ARTIFACT_REGISTRY_REPO}/{GCP_TRAINING_IMAGE_NAME}:{GCP_TRAINING_IMAGE_TAG}"
|
| 111 |
+
)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def gcp_experiments_uri_for_progress() -> str:
|
| 115 |
+
"""GCS prefix for live progress in the dashboard (explicit env or ``.../data/experiments``)."""
|
| 116 |
+
explicit = (os.environ.get("EXCEL_CHUNKER_GCS_EXPERIMENTS_URI") or "").strip()
|
| 117 |
+
if explicit:
|
| 118 |
+
return explicit
|
| 119 |
+
if GCP_GCS_DATA_URI:
|
| 120 |
+
return GCP_GCS_DATA_URI.rstrip("/") + "/experiments"
|
| 121 |
+
return ""
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
def docker_build_from_dashboard_enabled() -> bool:
|
| 125 |
+
"""Allow POST /api/train/gcp-image-build only on trusted machines with Docker (local dev)."""
|
| 126 |
+
v = (os.environ.get("ENABLE_DOCKER_BUILD_FROM_DASHBOARD") or "").strip().lower()
|
| 127 |
+
return v in ("1", "true", "yes", "on")
|
| 128 |
+
|
| 129 |
|
| 130 |
def get_google_credentials():
|
| 131 |
"""Return service account credentials (works for Sheets; read-only Drive)."""
|
embed_text.py
CHANGED
|
@@ -12,6 +12,7 @@ Usage:
|
|
| 12 |
from __future__ import annotations
|
| 13 |
|
| 14 |
import json
|
|
|
|
| 15 |
from pathlib import Path
|
| 16 |
|
| 17 |
import numpy as np
|
|
@@ -80,11 +81,61 @@ def embed_file(filepath: Path) -> Path:
|
|
| 80 |
return out_path
|
| 81 |
|
| 82 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 83 |
def embed_all() -> None:
|
| 84 |
"""Process all labeled JSON files in the labeled directory."""
|
| 85 |
labeled_dir = config.LABELED_DIR
|
| 86 |
-
json_files =
|
| 87 |
-
json_files = [f for f in json_files if not f.stem.endswith("_embeddings")]
|
| 88 |
|
| 89 |
if not json_files:
|
| 90 |
print(f"No labeled files found in {labeled_dir}")
|
|
|
|
| 12 |
from __future__ import annotations
|
| 13 |
|
| 14 |
import json
|
| 15 |
+
from collections.abc import Callable
|
| 16 |
from pathlib import Path
|
| 17 |
|
| 18 |
import numpy as np
|
|
|
|
| 81 |
return out_path
|
| 82 |
|
| 83 |
|
| 84 |
+
def _labeled_json_files(labeled_dir: Path) -> list[Path]:
|
| 85 |
+
files = sorted(labeled_dir.glob("*.json"))
|
| 86 |
+
return [f for f in files if not f.stem.endswith("_embeddings")]
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def needs_embedding_file(json_path: Path) -> bool:
|
| 90 |
+
"""True if *_embeddings.npz is missing or older than the labeled JSON."""
|
| 91 |
+
out = json_path.with_name(json_path.stem + "_embeddings.npz")
|
| 92 |
+
if not out.exists():
|
| 93 |
+
return True
|
| 94 |
+
try:
|
| 95 |
+
return json_path.stat().st_mtime > out.stat().st_mtime
|
| 96 |
+
except OSError:
|
| 97 |
+
return True
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
def ensure_embeddings_for_labeled(
|
| 101 |
+
*,
|
| 102 |
+
labeled_dir: Path | None = None,
|
| 103 |
+
log: Callable[[str], None] = print,
|
| 104 |
+
) -> dict[str, int | list[str]]:
|
| 105 |
+
"""Create or refresh ``*_embeddings.npz`` for every labeled JSON that needs it.
|
| 106 |
+
|
| 107 |
+
Call this before ``load_all_graphs`` / training so E5 text features exist.
|
| 108 |
+
Returns ``{"generated": n, "skipped": n, "errors": [...]}``.
|
| 109 |
+
"""
|
| 110 |
+
root = labeled_dir or config.LABELED_DIR
|
| 111 |
+
json_files = _labeled_json_files(root)
|
| 112 |
+
generated = 0
|
| 113 |
+
skipped = 0
|
| 114 |
+
errors: list[str] = []
|
| 115 |
+
|
| 116 |
+
for fp in json_files:
|
| 117 |
+
if not needs_embedding_file(fp):
|
| 118 |
+
skipped += 1
|
| 119 |
+
continue
|
| 120 |
+
try:
|
| 121 |
+
embed_file(fp)
|
| 122 |
+
generated += 1
|
| 123 |
+
except Exception as e:
|
| 124 |
+
errors.append(f"{fp.name}: {e}")
|
| 125 |
+
|
| 126 |
+
if generated:
|
| 127 |
+
log(f" [embed] Done: {generated} file(s) written, {skipped} already up to date.")
|
| 128 |
+
elif json_files:
|
| 129 |
+
log(f" [embed] All {len(json_files)} labeled file(s) already have up-to-date embeddings.")
|
| 130 |
+
if errors:
|
| 131 |
+
log(f" [embed] {len(errors)} error(s) (training may use zeros for those sheets).")
|
| 132 |
+
return {"generated": generated, "skipped": skipped, "errors": errors}
|
| 133 |
+
|
| 134 |
+
|
| 135 |
def embed_all() -> None:
|
| 136 |
"""Process all labeled JSON files in the labeled directory."""
|
| 137 |
labeled_dir = config.LABELED_DIR
|
| 138 |
+
json_files = _labeled_json_files(labeled_dir)
|
|
|
|
| 139 |
|
| 140 |
if not json_files:
|
| 141 |
print(f"No labeled files found in {labeled_dir}")
|
eval_rag.py
CHANGED
|
@@ -29,6 +29,12 @@ Evaluation dimensions:
|
|
| 29 |
3. LLM-as-a-judge: answer generation + correctness scoring
|
| 30 |
4. (Learned graph methods) Structure vs human labels: levels 1–4 via ``evaluate_sheet``,
|
| 31 |
saved to ``data/rag_eval_structure_metrics.json`` (see ``rag_structure_rag_analysis.py`` for joins).
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 32 |
|
| 33 |
Usage:
|
| 34 |
python eval_rag.py # run on rag_eval_dataset.json
|
|
@@ -41,6 +47,7 @@ from __future__ import annotations
|
|
| 41 |
import json
|
| 42 |
import logging
|
| 43 |
import os
|
|
|
|
| 44 |
from datetime import datetime, timezone
|
| 45 |
from pathlib import Path
|
| 46 |
from typing import Any, Callable, Optional
|
|
@@ -81,6 +88,69 @@ def _embed_texts(model, texts: list[str], batch_size: int = 64) -> np.ndarray:
|
|
| 81 |
GNN_CELL_LIMIT = 3000
|
| 82 |
|
| 83 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 84 |
def _predict_and_label(
|
| 85 |
grid_data: dict,
|
| 86 |
model=None,
|
|
@@ -137,26 +207,88 @@ def _fill_prediction_cache(
|
|
| 137 |
sheet_data: dict[str, dict],
|
| 138 |
model,
|
| 139 |
log: Callable[[str], None] | None = None,
|
| 140 |
-
|
| 141 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 142 |
from predict import predict_sheet
|
| 143 |
|
| 144 |
cache: dict[str, list[dict] | None] = {}
|
|
|
|
|
|
|
|
|
|
| 145 |
for sheet_id, data in sheet_data.items():
|
|
|
|
| 146 |
n_cells = len(data.get("cells", []))
|
| 147 |
if n_cells > GNN_CELL_LIMIT:
|
| 148 |
cache[sheet_id] = None
|
|
|
|
| 149 |
continue
|
|
|
|
| 150 |
try:
|
| 151 |
p = predict_sheet(data, model_override=model)
|
|
|
|
| 152 |
cache[sheet_id] = p if p else None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 153 |
except (RuntimeError, MemoryError) as e:
|
|
|
|
| 154 |
if log:
|
| 155 |
log(f" predict_sheet failed {sheet_id[:40]}... ({n_cells} cells): {e}")
|
| 156 |
logger.warning("predict_sheet failed on %s: %s", sheet_id, e)
|
| 157 |
import gc; gc.collect()
|
| 158 |
cache[sheet_id] = None
|
| 159 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 160 |
|
| 161 |
|
| 162 |
def _aggregate_structure_sheet_results(rows: list[dict]) -> dict[str, float | None]:
|
|
@@ -1009,8 +1141,15 @@ def _should_skip_rag_resume(
|
|
| 1009 |
resume: bool,
|
| 1010 |
existing_agg: dict[str, Any],
|
| 1011 |
log_fn: Callable[[str], None] | None,
|
|
|
|
|
|
|
| 1012 |
) -> bool:
|
| 1013 |
-
"""Skip re-eval if this key already has a successful checkpoint (no ``error``).
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1014 |
if not resume:
|
| 1015 |
return False
|
| 1016 |
val = existing_agg.get(key)
|
|
@@ -1018,7 +1157,15 @@ def _should_skip_rag_resume(
|
|
| 1018 |
return False
|
| 1019 |
if val.get("error"):
|
| 1020 |
return False
|
| 1021 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1022 |
if log_fn:
|
| 1023 |
log_fn(f" SKIP {key} (resume — already in rag_eval_results.json)")
|
| 1024 |
return True
|
|
@@ -1086,6 +1233,7 @@ def run_rag_evaluation_v2(
|
|
| 1086 |
resume_structure: bool = True,
|
| 1087 |
force_structure_metrics: bool = False,
|
| 1088 |
save_structure_details: bool = False,
|
|
|
|
| 1089 |
) -> dict:
|
| 1090 |
"""Run full RAG evaluation using the rag_eval_dataset.json.
|
| 1091 |
|
|
@@ -1108,6 +1256,9 @@ def run_rag_evaluation_v2(
|
|
| 1108 |
resume_structure: skip rewriting ``rag_eval_structure_metrics.json`` when an entry exists.
|
| 1109 |
force_structure_metrics: recompute structure-vs-gold even if a checkpoint exists.
|
| 1110 |
save_structure_details: include per-sheet ``_per_sheet`` in the structure metrics JSON.
|
|
|
|
|
|
|
|
|
|
| 1111 |
"""
|
| 1112 |
def log(msg: str):
|
| 1113 |
print(msg)
|
|
@@ -1297,7 +1448,9 @@ def run_rag_evaluation_v2(
|
|
| 1297 |
|
| 1298 |
for strat in strategies:
|
| 1299 |
key = f"{method_key_base}/{strat}"
|
| 1300 |
-
if _should_skip_rag_resume(
|
|
|
|
|
|
|
| 1301 |
results[key] = dict(existing_agg[key])
|
| 1302 |
continue
|
| 1303 |
log(f" Evaluating {key}...")
|
|
@@ -1348,7 +1501,20 @@ def run_rag_evaluation_v2(
|
|
| 1348 |
continue
|
| 1349 |
|
| 1350 |
log(" predict_sheet cache (one inference per sheet)...")
|
| 1351 |
-
pred_cache = _fill_prediction_cache(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1352 |
|
| 1353 |
if not _should_skip_structure_resume(
|
| 1354 |
method_key_base, resume_structure, structure_metrics_doc, force_structure_metrics,
|
|
@@ -1357,7 +1523,10 @@ def run_rag_evaluation_v2(
|
|
| 1357 |
smeta = _compute_structure_vs_gold_for_arch(
|
| 1358 |
sheet_data, pred_cache, save_per_sheet=save_structure_details,
|
| 1359 |
)
|
|
|
|
|
|
|
| 1360 |
entry = {
|
|
|
|
| 1361 |
"method_base": method_key_base,
|
| 1362 |
"chunking": gm,
|
| 1363 |
"computed_at": datetime.now(timezone.utc).isoformat(),
|
|
@@ -1365,7 +1534,8 @@ def run_rag_evaluation_v2(
|
|
| 1365 |
"held_out_sheet_count": len(sheet_data),
|
| 1366 |
**smeta,
|
| 1367 |
}
|
| 1368 |
-
|
|
|
|
| 1369 |
_write_rag_structure_metrics_doc(structure_metrics_doc)
|
| 1370 |
log(
|
| 1371 |
f" Structure vs gold: n_sheets={smeta['n_sheets_with_gold']} → "
|
|
@@ -1374,8 +1544,28 @@ def run_rag_evaluation_v2(
|
|
| 1374 |
except Exception as e:
|
| 1375 |
log(f" WARNING: structure metrics failed for {method_key_base}: {e}")
|
| 1376 |
logger.exception("structure metrics")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1377 |
else:
|
| 1378 |
log(f" SKIP structure metrics for {method_key_base} (resume_structure)")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1379 |
|
| 1380 |
all_chunks = []
|
| 1381 |
try:
|
|
@@ -1450,7 +1640,9 @@ def run_rag_evaluation_v2(
|
|
| 1450 |
|
| 1451 |
for strat in strategies:
|
| 1452 |
key = f"{method_key_base}/{strat}"
|
| 1453 |
-
if _should_skip_rag_resume(
|
|
|
|
|
|
|
| 1454 |
results[key] = dict(existing_agg[key])
|
| 1455 |
continue
|
| 1456 |
log(f" Evaluating {key}...")
|
|
@@ -1545,7 +1737,9 @@ def run_rag_evaluation_v2(
|
|
| 1545 |
|
| 1546 |
for strat in strategies:
|
| 1547 |
key = f"{method}/{strat}"
|
| 1548 |
-
if _should_skip_rag_resume(
|
|
|
|
|
|
|
| 1549 |
results[key] = dict(existing_agg[key])
|
| 1550 |
continue
|
| 1551 |
log(f" Evaluating {key}...")
|
|
@@ -1782,6 +1976,11 @@ if __name__ == "__main__":
|
|
| 1782 |
action="store_true",
|
| 1783 |
help="Store per-sheet _per_sheet in rag_eval_structure_metrics.json",
|
| 1784 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1785 |
args = parser.parse_args()
|
| 1786 |
|
| 1787 |
if args.pull_drive:
|
|
@@ -1805,4 +2004,5 @@ if __name__ == "__main__":
|
|
| 1805 |
resume_structure=not args.no_resume_structure,
|
| 1806 |
force_structure_metrics=args.force_structure_metrics,
|
| 1807 |
save_structure_details=args.save_structure_details,
|
|
|
|
| 1808 |
)
|
|
|
|
| 29 |
3. LLM-as-a-judge: answer generation + correctness scoring
|
| 30 |
4. (Learned graph methods) Structure vs human labels: levels 1–4 via ``evaluate_sheet``,
|
| 31 |
saved to ``data/rag_eval_structure_metrics.json`` (see ``rag_structure_rag_analysis.py`` for joins).
|
| 32 |
+
5. (Learned graph methods) ``predict_sheet`` wall time per held-out sheet; aggregates
|
| 33 |
+
(mean/p95 ms per graph node, throughput) in ``inference_timing`` on the same JSON entry.
|
| 34 |
+
First sheet may include one-off CPU/GPU warmup — interpret percentiles accordingly.
|
| 35 |
+
When structure metrics are recomputed, the new entry is merged with the previous one so
|
| 36 |
+
``inference_timing_per_sheet`` (and other extra keys) are not dropped unless this run
|
| 37 |
+
saves a fresh per-sheet list with ``--save-structure-details``.
|
| 38 |
|
| 39 |
Usage:
|
| 40 |
python eval_rag.py # run on rag_eval_dataset.json
|
|
|
|
| 47 |
import json
|
| 48 |
import logging
|
| 49 |
import os
|
| 50 |
+
import time
|
| 51 |
from datetime import datetime, timezone
|
| 52 |
from pathlib import Path
|
| 53 |
from typing import Any, Callable, Optional
|
|
|
|
| 88 |
GNN_CELL_LIMIT = 3000
|
| 89 |
|
| 90 |
|
| 91 |
+
def _rag_grid_cell_counts(grid: dict) -> tuple[int, int]:
|
| 92 |
+
"""Return (n_json_cells, n_graph_nodes).
|
| 93 |
+
|
| 94 |
+
``n_graph_nodes`` matches ``predict_sheet`` nodes: cells with non-empty ``value``.
|
| 95 |
+
``n_json_cells`` is ``len(cells)`` on the grid dict.
|
| 96 |
+
"""
|
| 97 |
+
cells = grid.get("cells") or []
|
| 98 |
+
n_json = len(cells)
|
| 99 |
+
n_graph = sum(1 for c in cells if (c.get("value") or "") != "")
|
| 100 |
+
return n_json, n_graph
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
def _aggregate_predict_timing(
|
| 104 |
+
per_sheet: list[dict[str, Any]],
|
| 105 |
+
*,
|
| 106 |
+
n_skipped_limit: int = 0,
|
| 107 |
+
stats_rows: list[dict[str, Any]] | None = None,
|
| 108 |
+
) -> dict[str, Any]:
|
| 109 |
+
"""Aggregate wall-clock ``predict_sheet`` timings into JSON-serialisable metrics.
|
| 110 |
+
|
| 111 |
+
``per_sheet`` is the full list of attempts (for counts). ``stats_rows``, if set,
|
| 112 |
+
is the subset used for mean / percentiles (e.g. exclude the first sheet for warmup).
|
| 113 |
+
"""
|
| 114 |
+
rows_stats = stats_rows if stats_rows is not None else per_sheet
|
| 115 |
+
ok_rows = [r for r in rows_stats if r.get("success")]
|
| 116 |
+
n_failed = sum(1 for r in per_sheet if not r.get("skipped_limit") and not r.get("success"))
|
| 117 |
+
|
| 118 |
+
def _ms_per_node(r: dict) -> float:
|
| 119 |
+
n = max(int(r.get("n_graph_nodes") or 0), 1)
|
| 120 |
+
return float(r["elapsed_s"]) * 1000.0 / n
|
| 121 |
+
|
| 122 |
+
def _ms_per_json(r: dict) -> float:
|
| 123 |
+
n = max(int(r.get("n_json_cells") or 0), 1)
|
| 124 |
+
return float(r["elapsed_s"]) * 1000.0 / n
|
| 125 |
+
|
| 126 |
+
ms_nodes = [_ms_per_node(r) for r in ok_rows]
|
| 127 |
+
ms_json = [_ms_per_json(r) for r in ok_rows]
|
| 128 |
+
total_s = sum(float(r["elapsed_s"]) for r in ok_rows)
|
| 129 |
+
total_nodes = sum(int(r.get("n_graph_nodes") or 0) for r in ok_rows)
|
| 130 |
+
|
| 131 |
+
def _pctl(arr: list[float], q: float) -> float | None:
|
| 132 |
+
if not arr:
|
| 133 |
+
return None
|
| 134 |
+
return float(np.percentile(np.array(arr, dtype=np.float64), q))
|
| 135 |
+
|
| 136 |
+
out: dict[str, Any] = {
|
| 137 |
+
"n_sheets_attempted": len(per_sheet),
|
| 138 |
+
"n_sheets_used_for_percentiles": len(rows_stats),
|
| 139 |
+
"n_sheets_predict_ok": len(ok_rows),
|
| 140 |
+
"n_sheets_predict_failed": n_failed,
|
| 141 |
+
"n_sheets_skipped_gnn_cell_limit": int(n_skipped_limit),
|
| 142 |
+
"predict_total_s": round(total_s, 4) if ok_rows else None,
|
| 143 |
+
"total_graph_nodes": int(total_nodes) if ok_rows else None,
|
| 144 |
+
"throughput_nodes_per_s": round(total_nodes / total_s, 2) if ok_rows and total_s > 0 else None,
|
| 145 |
+
"mean_ms_per_graph_node": round(float(np.mean(ms_nodes)), 3) if ms_nodes else None,
|
| 146 |
+
"p50_ms_per_graph_node": (round(v, 3) if (v := _pctl(ms_nodes, 50)) is not None else None),
|
| 147 |
+
"p95_ms_per_graph_node": (round(v, 3) if (v := _pctl(ms_nodes, 95)) is not None else None),
|
| 148 |
+
"mean_ms_per_json_cell": round(float(np.mean(ms_json)), 3) if ms_json else None,
|
| 149 |
+
"p95_ms_per_json_cell": (round(v, 3) if (v := _pctl(ms_json, 95)) is not None else None),
|
| 150 |
+
}
|
| 151 |
+
return out
|
| 152 |
+
|
| 153 |
+
|
| 154 |
def _predict_and_label(
|
| 155 |
grid_data: dict,
|
| 156 |
model=None,
|
|
|
|
| 207 |
sheet_data: dict[str, dict],
|
| 208 |
model,
|
| 209 |
log: Callable[[str], None] | None = None,
|
| 210 |
+
*,
|
| 211 |
+
timing_exclude_first_sheet: bool = False,
|
| 212 |
+
) -> tuple[dict[str, list[dict] | None], dict[str, Any]]:
|
| 213 |
+
"""One ``predict_sheet`` per held-out sheet (under GNN_CELL_LIMIT).
|
| 214 |
+
|
| 215 |
+
Returns ``(pred_cache, timing_bundle)`` where ``timing_bundle`` has keys
|
| 216 |
+
``aggregate`` (from :func:`_aggregate_predict_timing`) and ``per_sheet`` (list of
|
| 217 |
+
per-attempt records for optional persistence / debugging).
|
| 218 |
+
"""
|
| 219 |
from predict import predict_sheet
|
| 220 |
|
| 221 |
cache: dict[str, list[dict] | None] = {}
|
| 222 |
+
per_sheet: list[dict[str, Any]] = []
|
| 223 |
+
n_skipped_limit = 0
|
| 224 |
+
|
| 225 |
for sheet_id, data in sheet_data.items():
|
| 226 |
+
n_json, n_graph = _rag_grid_cell_counts(data)
|
| 227 |
n_cells = len(data.get("cells", []))
|
| 228 |
if n_cells > GNN_CELL_LIMIT:
|
| 229 |
cache[sheet_id] = None
|
| 230 |
+
n_skipped_limit += 1
|
| 231 |
continue
|
| 232 |
+
t0 = time.perf_counter()
|
| 233 |
try:
|
| 234 |
p = predict_sheet(data, model_override=model)
|
| 235 |
+
elapsed = time.perf_counter() - t0
|
| 236 |
cache[sheet_id] = p if p else None
|
| 237 |
+
success = bool(p)
|
| 238 |
+
per_sheet.append({
|
| 239 |
+
"sheet_id": sheet_id,
|
| 240 |
+
"elapsed_s": round(elapsed, 6),
|
| 241 |
+
"n_json_cells": n_json,
|
| 242 |
+
"n_graph_nodes": n_graph,
|
| 243 |
+
"success": success,
|
| 244 |
+
"skipped_limit": False,
|
| 245 |
+
"ms_per_graph_node": round(elapsed * 1000.0 / max(n_graph, 1), 4),
|
| 246 |
+
"ms_per_json_cell": round(elapsed * 1000.0 / max(n_json, 1), 4),
|
| 247 |
+
})
|
| 248 |
except (RuntimeError, MemoryError) as e:
|
| 249 |
+
elapsed = time.perf_counter() - t0
|
| 250 |
if log:
|
| 251 |
log(f" predict_sheet failed {sheet_id[:40]}... ({n_cells} cells): {e}")
|
| 252 |
logger.warning("predict_sheet failed on %s: %s", sheet_id, e)
|
| 253 |
import gc; gc.collect()
|
| 254 |
cache[sheet_id] = None
|
| 255 |
+
per_sheet.append({
|
| 256 |
+
"sheet_id": sheet_id,
|
| 257 |
+
"elapsed_s": round(elapsed, 6),
|
| 258 |
+
"n_json_cells": n_json,
|
| 259 |
+
"n_graph_nodes": n_graph,
|
| 260 |
+
"success": False,
|
| 261 |
+
"skipped_limit": False,
|
| 262 |
+
"error": str(e),
|
| 263 |
+
"ms_per_graph_node": round(elapsed * 1000.0 / max(n_graph, 1), 4),
|
| 264 |
+
"ms_per_json_cell": round(elapsed * 1000.0 / max(n_json, 1), 4),
|
| 265 |
+
})
|
| 266 |
+
|
| 267 |
+
stats_rows = per_sheet
|
| 268 |
+
if timing_exclude_first_sheet and len(per_sheet) > 1:
|
| 269 |
+
stats_rows = per_sheet[1:]
|
| 270 |
+
agg = _aggregate_predict_timing(
|
| 271 |
+
per_sheet, n_skipped_limit=n_skipped_limit, stats_rows=stats_rows,
|
| 272 |
+
)
|
| 273 |
+
bundle = {"aggregate": agg, "per_sheet": per_sheet}
|
| 274 |
+
if log:
|
| 275 |
+
if per_sheet:
|
| 276 |
+
a = agg
|
| 277 |
+
log(
|
| 278 |
+
f" predict_sheet timing: total={a.get('predict_total_s')}s "
|
| 279 |
+
f"ok={a.get('n_sheets_predict_ok')}/{a.get('n_sheets_attempted')} "
|
| 280 |
+
f"mean_ms/node={a.get('mean_ms_per_graph_node')} "
|
| 281 |
+
f"p95_ms/node={a.get('p95_ms_per_graph_node')} "
|
| 282 |
+
f"nodes/s={a.get('throughput_nodes_per_s')} "
|
| 283 |
+
f"skipped_limit={a.get('n_sheets_skipped_gnn_cell_limit')}"
|
| 284 |
+
)
|
| 285 |
+
else:
|
| 286 |
+
log(
|
| 287 |
+
f" predict_sheet timing: no predict attempts "
|
| 288 |
+
f"(skipped GNN_CELL_LIMIT={n_skipped_limit} / {len(sheet_data)} sheets)"
|
| 289 |
+
)
|
| 290 |
+
|
| 291 |
+
return cache, bundle
|
| 292 |
|
| 293 |
|
| 294 |
def _aggregate_structure_sheet_results(rows: list[dict]) -> dict[str, float | None]:
|
|
|
|
| 1141 |
resume: bool,
|
| 1142 |
existing_agg: dict[str, Any],
|
| 1143 |
log_fn: Callable[[str], None] | None,
|
| 1144 |
+
*,
|
| 1145 |
+
need_judge: bool = False,
|
| 1146 |
) -> bool:
|
| 1147 |
+
"""Skip re-eval if this key already has a successful checkpoint (no ``error``).
|
| 1148 |
+
|
| 1149 |
+
If ``need_judge`` is True (LLM-as-a-judge requested), only skip when retrieval
|
| 1150 |
+
metrics *and* judge fields are present — otherwise a prior ``--no-judge`` run
|
| 1151 |
+
would block adding judge scores.
|
| 1152 |
+
"""
|
| 1153 |
if not resume:
|
| 1154 |
return False
|
| 1155 |
val = existing_agg.get(key)
|
|
|
|
| 1157 |
return False
|
| 1158 |
if val.get("error"):
|
| 1159 |
return False
|
| 1160 |
+
has_retrieval = any(k in val for k in ("recall@1", "mrr"))
|
| 1161 |
+
has_judge = any(k in val for k in ("judge_score", "judge_binary_acc"))
|
| 1162 |
+
if need_judge:
|
| 1163 |
+
if has_retrieval and has_judge:
|
| 1164 |
+
if log_fn:
|
| 1165 |
+
log_fn(f" SKIP {key} (resume — already in rag_eval_results.json)")
|
| 1166 |
+
return True
|
| 1167 |
+
return False
|
| 1168 |
+
if has_retrieval or has_judge:
|
| 1169 |
if log_fn:
|
| 1170 |
log_fn(f" SKIP {key} (resume — already in rag_eval_results.json)")
|
| 1171 |
return True
|
|
|
|
| 1233 |
resume_structure: bool = True,
|
| 1234 |
force_structure_metrics: bool = False,
|
| 1235 |
save_structure_details: bool = False,
|
| 1236 |
+
timing_exclude_first_sheet: bool = False,
|
| 1237 |
) -> dict:
|
| 1238 |
"""Run full RAG evaluation using the rag_eval_dataset.json.
|
| 1239 |
|
|
|
|
| 1256 |
resume_structure: skip rewriting ``rag_eval_structure_metrics.json`` when an entry exists.
|
| 1257 |
force_structure_metrics: recompute structure-vs-gold even if a checkpoint exists.
|
| 1258 |
save_structure_details: include per-sheet ``_per_sheet`` in the structure metrics JSON.
|
| 1259 |
+
timing_exclude_first_sheet: if True, mean / percentiles for ``inference_timing`` exclude
|
| 1260 |
+
the first attempted ``predict_sheet`` call (reduces one-off warmup skew); counts
|
| 1261 |
+
still reflect all sheets.
|
| 1262 |
"""
|
| 1263 |
def log(msg: str):
|
| 1264 |
print(msg)
|
|
|
|
| 1448 |
|
| 1449 |
for strat in strategies:
|
| 1450 |
key = f"{method_key_base}/{strat}"
|
| 1451 |
+
if _should_skip_rag_resume(
|
| 1452 |
+
key, resume, existing_agg, log, need_judge=use_llm_judge,
|
| 1453 |
+
):
|
| 1454 |
results[key] = dict(existing_agg[key])
|
| 1455 |
continue
|
| 1456 |
log(f" Evaluating {key}...")
|
|
|
|
| 1501 |
continue
|
| 1502 |
|
| 1503 |
log(" predict_sheet cache (one inference per sheet)...")
|
| 1504 |
+
pred_cache, timing_bundle = _fill_prediction_cache(
|
| 1505 |
+
sheet_data, model, log,
|
| 1506 |
+
timing_exclude_first_sheet=timing_exclude_first_sheet,
|
| 1507 |
+
)
|
| 1508 |
+
inf_agg = timing_bundle.get("aggregate") or {}
|
| 1509 |
+
inf_per = timing_bundle.get("per_sheet") or []
|
| 1510 |
+
|
| 1511 |
+
def _attach_inference_timing(entry: dict[str, Any]) -> dict[str, Any]:
|
| 1512 |
+
out = dict(entry)
|
| 1513 |
+
out["inference_timing"] = inf_agg
|
| 1514 |
+
if save_structure_details and inf_per:
|
| 1515 |
+
out["inference_timing_per_sheet"] = list(inf_per)
|
| 1516 |
+
# If this run did not emit per-sheet timing, keep any existing list from ``entry``.
|
| 1517 |
+
return out
|
| 1518 |
|
| 1519 |
if not _should_skip_structure_resume(
|
| 1520 |
method_key_base, resume_structure, structure_metrics_doc, force_structure_metrics,
|
|
|
|
| 1523 |
smeta = _compute_structure_vs_gold_for_arch(
|
| 1524 |
sheet_data, pred_cache, save_per_sheet=save_structure_details,
|
| 1525 |
)
|
| 1526 |
+
prev_entry = structure_metrics_doc.get(method_key_base)
|
| 1527 |
+
prev_dict = prev_entry if isinstance(prev_entry, dict) else {}
|
| 1528 |
entry = {
|
| 1529 |
+
**prev_dict,
|
| 1530 |
"method_base": method_key_base,
|
| 1531 |
"chunking": gm,
|
| 1532 |
"computed_at": datetime.now(timezone.utc).isoformat(),
|
|
|
|
| 1534 |
"held_out_sheet_count": len(sheet_data),
|
| 1535 |
**smeta,
|
| 1536 |
}
|
| 1537 |
+
entry.pop("structure_metrics_error", None)
|
| 1538 |
+
structure_metrics_doc[method_key_base] = _attach_inference_timing(entry)
|
| 1539 |
_write_rag_structure_metrics_doc(structure_metrics_doc)
|
| 1540 |
log(
|
| 1541 |
f" Structure vs gold: n_sheets={smeta['n_sheets_with_gold']} → "
|
|
|
|
| 1544 |
except Exception as e:
|
| 1545 |
log(f" WARNING: structure metrics failed for {method_key_base}: {e}")
|
| 1546 |
logger.exception("structure metrics")
|
| 1547 |
+
prev = structure_metrics_doc.get(method_key_base)
|
| 1548 |
+
if not isinstance(prev, dict):
|
| 1549 |
+
prev = {}
|
| 1550 |
+
fallback = {
|
| 1551 |
+
**prev,
|
| 1552 |
+
"method_base": method_key_base,
|
| 1553 |
+
"chunking": gm,
|
| 1554 |
+
"computed_at": datetime.now(timezone.utc).isoformat(),
|
| 1555 |
+
"structure_metrics_error": str(e),
|
| 1556 |
+
}
|
| 1557 |
+
structure_metrics_doc[method_key_base] = _attach_inference_timing(fallback)
|
| 1558 |
+
_write_rag_structure_metrics_doc(structure_metrics_doc)
|
| 1559 |
else:
|
| 1560 |
log(f" SKIP structure metrics for {method_key_base} (resume_structure)")
|
| 1561 |
+
prev = structure_metrics_doc.get(method_key_base)
|
| 1562 |
+
if not isinstance(prev, dict):
|
| 1563 |
+
prev = {}
|
| 1564 |
+
merged = {**prev, "method_base": method_key_base, "chunking": gm}
|
| 1565 |
+
merged["inference_timing_updated_at"] = datetime.now(timezone.utc).isoformat()
|
| 1566 |
+
structure_metrics_doc[method_key_base] = _attach_inference_timing(merged)
|
| 1567 |
+
_write_rag_structure_metrics_doc(structure_metrics_doc)
|
| 1568 |
+
log(f" Updated inference_timing only → {config.RAG_EVAL_STRUCTURE_METRICS.name}")
|
| 1569 |
|
| 1570 |
all_chunks = []
|
| 1571 |
try:
|
|
|
|
| 1640 |
|
| 1641 |
for strat in strategies:
|
| 1642 |
key = f"{method_key_base}/{strat}"
|
| 1643 |
+
if _should_skip_rag_resume(
|
| 1644 |
+
key, resume, existing_agg, log, need_judge=use_llm_judge,
|
| 1645 |
+
):
|
| 1646 |
results[key] = dict(existing_agg[key])
|
| 1647 |
continue
|
| 1648 |
log(f" Evaluating {key}...")
|
|
|
|
| 1737 |
|
| 1738 |
for strat in strategies:
|
| 1739 |
key = f"{method}/{strat}"
|
| 1740 |
+
if _should_skip_rag_resume(
|
| 1741 |
+
key, resume, existing_agg, log, need_judge=use_llm_judge,
|
| 1742 |
+
):
|
| 1743 |
results[key] = dict(existing_agg[key])
|
| 1744 |
continue
|
| 1745 |
log(f" Evaluating {key}...")
|
|
|
|
| 1976 |
action="store_true",
|
| 1977 |
help="Store per-sheet _per_sheet in rag_eval_structure_metrics.json",
|
| 1978 |
)
|
| 1979 |
+
parser.add_argument(
|
| 1980 |
+
"--timing-exclude-first-sheet",
|
| 1981 |
+
action="store_true",
|
| 1982 |
+
help="For inference_timing aggregates only: drop the first predict_sheet call from mean/p95 (warmup)",
|
| 1983 |
+
)
|
| 1984 |
args = parser.parse_args()
|
| 1985 |
|
| 1986 |
if args.pull_drive:
|
|
|
|
| 2004 |
resume_structure=not args.no_resume_structure,
|
| 2005 |
force_structure_metrics=args.force_structure_metrics,
|
| 2006 |
save_structure_details=args.save_structure_details,
|
| 2007 |
+
timing_exclude_first_sheet=args.timing_exclude_first_sheet,
|
| 2008 |
)
|
gcp_progress_reader.py
CHANGED
|
@@ -77,7 +77,10 @@ def fetch_suite_from_gcs(uri: str) -> dict[str, Any]:
|
|
| 77 |
def _sort_key(u: dict) -> tuple:
|
| 78 |
return (u.get("seed", 0), u.get("fold", 0), u.get("arch", ""), u.get("run_tag", ""))
|
| 79 |
|
|
|
|
| 80 |
units.sort(key=_sort_key)
|
|
|
|
|
|
|
| 81 |
done_like = {"done", "skipped", "error"}
|
| 82 |
completed = sum(1 for u in units if u.get("status") in done_like)
|
| 83 |
running = any(u.get("status") == "running" for u in units)
|
|
|
|
| 77 |
def _sort_key(u: dict) -> tuple:
|
| 78 |
return (u.get("seed", 0), u.get("fold", 0), u.get("arch", ""), u.get("run_tag", ""))
|
| 79 |
|
| 80 |
+
# Stable fallback order (seed / fold / arch).
|
| 81 |
units.sort(key=_sort_key)
|
| 82 |
+
# Newest worker activity first so a fresh Vertex run is visible at the top.
|
| 83 |
+
units.sort(key=lambda u: u.get("updated_at") or "", reverse=True)
|
| 84 |
done_like = {"done", "skipped", "error"}
|
| 85 |
completed = sum(1 for u in units if u.get("status") in done_like)
|
| 86 |
running = any(u.get("status") == "running" for u in units)
|
metadata_client.py
CHANGED
|
@@ -51,6 +51,11 @@ COL_TIMESTAMP = 7
|
|
| 51 |
COL_NUM_LABELED_CELLS = 8
|
| 52 |
|
| 53 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 54 |
def _get_service():
|
| 55 |
global _service
|
| 56 |
if _service is None:
|
|
|
|
| 51 |
COL_NUM_LABELED_CELLS = 8
|
| 52 |
|
| 53 |
|
| 54 |
+
def label_file_stem(drive_file_id: str, sheet_name: str) -> str:
|
| 55 |
+
"""Stem for ``data/labeled/{stem}.json``; must match ``app.py`` save/load."""
|
| 56 |
+
return f"{drive_file_id}_{sheet_name}".replace("/", "_").replace(" ", "_")
|
| 57 |
+
|
| 58 |
+
|
| 59 |
def _get_service():
|
| 60 |
global _service
|
| 61 |
if _service is None:
|
rag_structure_rag_analysis.py
CHANGED
|
@@ -41,6 +41,21 @@ RAG_KEYS = (
|
|
| 41 |
"source_mrr",
|
| 42 |
)
|
| 43 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 44 |
|
| 45 |
def _load_json(path: Path) -> dict[str, Any]:
|
| 46 |
if not path.exists():
|
|
@@ -79,6 +94,9 @@ def build_joined_rows(
|
|
| 79 |
row[k] = agg.get(k)
|
| 80 |
for k in RAG_KEYS:
|
| 81 |
row[f"rag_{k}"] = rrag.get(k)
|
|
|
|
|
|
|
|
|
|
| 82 |
rows.append(row)
|
| 83 |
|
| 84 |
meta = {
|
|
|
|
| 41 |
"source_mrr",
|
| 42 |
)
|
| 43 |
|
| 44 |
+
# Flattened from ``rag_eval_structure_metrics.json`` ``inference_timing`` (eval_rag /predict_sheet).
|
| 45 |
+
INFERENCE_TIMING_KEYS = (
|
| 46 |
+
"mean_ms_per_graph_node",
|
| 47 |
+
"p50_ms_per_graph_node",
|
| 48 |
+
"p95_ms_per_graph_node",
|
| 49 |
+
"mean_ms_per_json_cell",
|
| 50 |
+
"p95_ms_per_json_cell",
|
| 51 |
+
"predict_total_s",
|
| 52 |
+
"throughput_nodes_per_s",
|
| 53 |
+
"n_sheets_predict_ok",
|
| 54 |
+
"n_sheets_attempted",
|
| 55 |
+
"n_sheets_skipped_gnn_cell_limit",
|
| 56 |
+
"n_sheets_predict_failed",
|
| 57 |
+
)
|
| 58 |
+
|
| 59 |
|
| 60 |
def _load_json(path: Path) -> dict[str, Any]:
|
| 61 |
if not path.exists():
|
|
|
|
| 94 |
row[k] = agg.get(k)
|
| 95 |
for k in RAG_KEYS:
|
| 96 |
row[f"rag_{k}"] = rrag.get(k)
|
| 97 |
+
it = sentry.get("inference_timing") if isinstance(sentry.get("inference_timing"), dict) else {}
|
| 98 |
+
for k in INFERENCE_TIMING_KEYS:
|
| 99 |
+
row[f"inf_{k}"] = it.get(k)
|
| 100 |
rows.append(row)
|
| 101 |
|
| 102 |
meta = {
|
requirements-app.txt
CHANGED
|
@@ -4,5 +4,6 @@ google-api-python-client
|
|
| 4 |
google-auth
|
| 5 |
google-auth-oauthlib
|
| 6 |
google-cloud-storage>=2.14.0
|
|
|
|
| 7 |
python-dotenv
|
| 8 |
openpyxl
|
|
|
|
| 4 |
google-auth
|
| 5 |
google-auth-oauthlib
|
| 6 |
google-cloud-storage>=2.14.0
|
| 7 |
+
google-cloud-aiplatform>=1.38.0
|
| 8 |
python-dotenv
|
| 9 |
openpyxl
|
requirements.txt
CHANGED
|
@@ -3,6 +3,7 @@ uvicorn[standard]
|
|
| 3 |
google-api-python-client
|
| 4 |
google-auth
|
| 5 |
google-auth-oauthlib
|
|
|
|
| 6 |
python-dotenv
|
| 7 |
openpyxl
|
| 8 |
datasets
|
|
|
|
| 3 |
google-api-python-client
|
| 4 |
google-auth
|
| 5 |
google-auth-oauthlib
|
| 6 |
+
google-cloud-aiplatform>=1.38.0
|
| 7 |
python-dotenv
|
| 8 |
openpyxl
|
| 9 |
datasets
|
run_ablation_sweep.py
CHANGED
|
@@ -81,7 +81,7 @@ def run_sweep(
|
|
| 81 |
continue
|
| 82 |
|
| 83 |
t0 = time.time()
|
| 84 |
-
exp_results = run_experiments(
|
| 85 |
architectures=[arch],
|
| 86 |
seeds=seeds,
|
| 87 |
n_folds=n_folds,
|
|
@@ -89,6 +89,8 @@ def run_sweep(
|
|
| 89 |
graph_cfg=graph_cfg,
|
| 90 |
max_nodes=max_nodes,
|
| 91 |
)
|
|
|
|
|
|
|
| 92 |
elapsed = time.time() - t0
|
| 93 |
|
| 94 |
summary = {
|
|
|
|
| 81 |
continue
|
| 82 |
|
| 83 |
t0 = time.time()
|
| 84 |
+
exp_results, n_failed = run_experiments(
|
| 85 |
architectures=[arch],
|
| 86 |
seeds=seeds,
|
| 87 |
n_folds=n_folds,
|
|
|
|
| 89 |
graph_cfg=graph_cfg,
|
| 90 |
max_nodes=max_nodes,
|
| 91 |
)
|
| 92 |
+
if n_failed:
|
| 93 |
+
print(f" WARNING: {n_failed} training unit(s) failed in this config")
|
| 94 |
elapsed = time.time() - t0
|
| 95 |
|
| 96 |
summary = {
|
scripts/launch_vertex_training.py
CHANGED
|
@@ -16,19 +16,28 @@ Example::
|
|
| 16 |
--architectures gat,gcn \\
|
| 17 |
--seeds 42 --folds 5
|
| 18 |
|
| 19 |
-
|
| 20 |
-
|
|
|
|
|
|
|
| 21 |
"""
|
| 22 |
|
| 23 |
from __future__ import annotations
|
| 24 |
|
| 25 |
import argparse
|
| 26 |
import sys
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 27 |
|
| 28 |
|
| 29 |
def main() -> None:
|
| 30 |
try:
|
| 31 |
-
|
| 32 |
except ImportError:
|
| 33 |
print("Install: pip install -r requirements-gcp.txt", file=sys.stderr)
|
| 34 |
raise SystemExit(1)
|
|
@@ -46,13 +55,40 @@ def main() -> None:
|
|
| 46 |
parser.add_argument("--seeds", default="42", help="Comma-separated ints")
|
| 47 |
parser.add_argument("--folds", type=int, default=5)
|
| 48 |
parser.add_argument("--epochs", type=int, default=100)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 49 |
parser.add_argument(
|
| 50 |
"--machine-type",
|
| 51 |
-
default="
|
| 52 |
-
help="
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 53 |
)
|
| 54 |
-
parser.add_argument("--accelerator-type", default="NVIDIA_TESLA_T4")
|
| 55 |
parser.add_argument("--accelerator-count", type=int, default=1)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 56 |
parser.add_argument(
|
| 57 |
"--staging-bucket",
|
| 58 |
default=None,
|
|
@@ -60,71 +96,36 @@ def main() -> None:
|
|
| 60 |
)
|
| 61 |
args = parser.parse_args()
|
| 62 |
|
| 63 |
-
staging = args.staging_bucket
|
| 64 |
-
if not staging:
|
| 65 |
-
uri = args.gcs_data_uri.rstrip("/")
|
| 66 |
-
if not uri.startswith("gs://"):
|
| 67 |
-
raise SystemExit("--gcs-data-uri must start with gs://")
|
| 68 |
-
rest = uri[5:]
|
| 69 |
-
bucket = rest.split("/", 1)[0]
|
| 70 |
-
staging = f"gs://{bucket}"
|
| 71 |
-
|
| 72 |
archs = [a.strip() for a in args.architectures.split(",") if a.strip()]
|
| 73 |
seeds = [int(s.strip()) for s in args.seeds.split(",") if s.strip()]
|
| 74 |
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
if
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
str(fold_idx),
|
| 103 |
-
]
|
| 104 |
-
env = [{"name": "EXCEL_CHUNKER_DATA_ROOT", "value": fuse_root}]
|
| 105 |
-
worker_pool = {
|
| 106 |
-
"machine_spec": {
|
| 107 |
-
"machine_type": args.machine_type,
|
| 108 |
-
"accelerator_type": args.accelerator_type,
|
| 109 |
-
"accelerator_count": args.accelerator_count,
|
| 110 |
-
},
|
| 111 |
-
"replica_count": 1,
|
| 112 |
-
"container_spec": {
|
| 113 |
-
"image_uri": args.image,
|
| 114 |
-
"command": ["python"],
|
| 115 |
-
"args": train_args,
|
| 116 |
-
"env": env,
|
| 117 |
-
},
|
| 118 |
-
}
|
| 119 |
-
job = aiplatform.CustomJob(
|
| 120 |
-
display_name=display[:128],
|
| 121 |
-
worker_pool_specs=[worker_pool],
|
| 122 |
-
)
|
| 123 |
-
job.run(sync=False)
|
| 124 |
-
print(f"Submitted: {display}")
|
| 125 |
-
jobs += 1
|
| 126 |
-
|
| 127 |
-
print(f"Total jobs submitted: {jobs}")
|
| 128 |
|
| 129 |
|
| 130 |
if __name__ == "__main__":
|
|
|
|
| 16 |
--architectures gat,gcn \\
|
| 17 |
--seeds 42 --folds 5
|
| 18 |
|
| 19 |
+
By default: **NVIDIA_L4** on **g2-standard-8** (quota e.g. ``CustomModelTrainingL4GPUsPerProjectPerRegion``).
|
| 20 |
+
L4 requires a **G2** machine type; the launcher fixes N1+T4-style combos automatically.
|
| 21 |
+
Use ``--accelerator-type NVIDIA_TESLA_T4 --machine-type n1-standard-8`` if you only have T4 quota.
|
| 22 |
+
Pass ``--spot`` + ``--accelerator-type NVIDIA_TESLA_P100`` if you only have preemptible P100 quota.
|
| 23 |
"""
|
| 24 |
|
| 25 |
from __future__ import annotations
|
| 26 |
|
| 27 |
import argparse
|
| 28 |
import sys
|
| 29 |
+
from pathlib import Path
|
| 30 |
+
|
| 31 |
+
_ROOT = Path(__file__).resolve().parents[1]
|
| 32 |
+
if str(_ROOT) not in sys.path:
|
| 33 |
+
sys.path.insert(0, str(_ROOT))
|
| 34 |
+
|
| 35 |
+
from vertex_launch import submit_vertex_training_jobs # noqa: E402
|
| 36 |
|
| 37 |
|
| 38 |
def main() -> None:
|
| 39 |
try:
|
| 40 |
+
import google.cloud.aiplatform # noqa: F401
|
| 41 |
except ImportError:
|
| 42 |
print("Install: pip install -r requirements-gcp.txt", file=sys.stderr)
|
| 43 |
raise SystemExit(1)
|
|
|
|
| 55 |
parser.add_argument("--seeds", default="42", help="Comma-separated ints")
|
| 56 |
parser.add_argument("--folds", type=int, default=5)
|
| 57 |
parser.add_argument("--epochs", type=int, default=100)
|
| 58 |
+
parser.add_argument(
|
| 59 |
+
"--subset-percent",
|
| 60 |
+
type=float,
|
| 61 |
+
default=100.0,
|
| 62 |
+
help="Random subset of sheets on worker (1-100; 100 = all)",
|
| 63 |
+
)
|
| 64 |
parser.add_argument(
|
| 65 |
"--machine-type",
|
| 66 |
+
default="g2-standard-8",
|
| 67 |
+
help="G2 for L4 (default); N1 for T4/P100",
|
| 68 |
+
)
|
| 69 |
+
parser.add_argument(
|
| 70 |
+
"--accelerator-type",
|
| 71 |
+
default="NVIDIA_L4",
|
| 72 |
+
help="Vertex enum: NVIDIA_L4 (default), NVIDIA_TESLA_T4, …",
|
| 73 |
)
|
|
|
|
| 74 |
parser.add_argument("--accelerator-count", type=int, default=1)
|
| 75 |
+
parser.add_argument(
|
| 76 |
+
"--heavy-archs",
|
| 77 |
+
default="",
|
| 78 |
+
help="Comma-separated architectures that use --heavy-accelerator "
|
| 79 |
+
"(e.g. adj_transformer,spatial_edge_transformer,dual_modality_gnn). "
|
| 80 |
+
"Others use --accelerator-type.",
|
| 81 |
+
)
|
| 82 |
+
parser.add_argument(
|
| 83 |
+
"--heavy-accelerator",
|
| 84 |
+
default="NVIDIA_L4",
|
| 85 |
+
help="GPU for --heavy-archs (default: L4; pair resolved with --machine-type)",
|
| 86 |
+
)
|
| 87 |
+
parser.add_argument(
|
| 88 |
+
"--spot",
|
| 89 |
+
action="store_true",
|
| 90 |
+
help="SPOT/preemptible scheduling (use with P100 if you lack on-demand T4 quota)",
|
| 91 |
+
)
|
| 92 |
parser.add_argument(
|
| 93 |
"--staging-bucket",
|
| 94 |
default=None,
|
|
|
|
| 96 |
)
|
| 97 |
args = parser.parse_args()
|
| 98 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 99 |
archs = [a.strip() for a in args.architectures.split(",") if a.strip()]
|
| 100 |
seeds = [int(s.strip()) for s in args.seeds.split(",") if s.strip()]
|
| 101 |
|
| 102 |
+
arch_accelerator = None
|
| 103 |
+
heavy_raw = (args.heavy_archs or "").strip()
|
| 104 |
+
if heavy_raw:
|
| 105 |
+
heavy_set = {a.strip() for a in heavy_raw.split(",") if a.strip()}
|
| 106 |
+
hacc = (args.heavy_accelerator or "").strip() or "NVIDIA_L4"
|
| 107 |
+
arch_accelerator = {a: hacc for a in archs if a in heavy_set}
|
| 108 |
+
|
| 109 |
+
out = submit_vertex_training_jobs(
|
| 110 |
+
project=args.project,
|
| 111 |
+
region=args.region,
|
| 112 |
+
image_uri=args.image,
|
| 113 |
+
gcs_data_uri=args.gcs_data_uri,
|
| 114 |
+
architectures=archs,
|
| 115 |
+
seeds=seeds,
|
| 116 |
+
folds=args.folds,
|
| 117 |
+
epochs=args.epochs,
|
| 118 |
+
subset_percent=args.subset_percent,
|
| 119 |
+
machine_type=args.machine_type,
|
| 120 |
+
accelerator_type=args.accelerator_type,
|
| 121 |
+
accelerator_count=args.accelerator_count,
|
| 122 |
+
staging_bucket=args.staging_bucket,
|
| 123 |
+
use_spot=args.spot,
|
| 124 |
+
arch_accelerator=arch_accelerator,
|
| 125 |
+
)
|
| 126 |
+
for d in out["jobs"]:
|
| 127 |
+
print(f"Submitted: {d}")
|
| 128 |
+
print(f"Total jobs submitted: {out['submitted_count']}")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 129 |
|
| 130 |
|
| 131 |
if __name__ == "__main__":
|
static/rag_dashboard.html
CHANGED
|
@@ -233,6 +233,7 @@ table.lb .best-cell{font-weight:700}
|
|
| 233 |
<h3>Structure vs gold + RAG (same retrieval strategy)</h3>
|
| 234 |
<p style="font-size:12px;color:var(--gray);margin:-6px 0 14px;line-height:1.45">
|
| 235 |
Rows: GNN chunking runs (<code>graph_*</code> / <code>graph_row_*</code>) with human labels on held-out sheets.
|
|
|
|
| 236 |
<strong>Oracle RAG scores</strong> (human chunking) measure a ceiling that still depends on retrieval, generation, and Q&A difficulty — not only perfect structure.
|
| 237 |
</p>
|
| 238 |
<div style="overflow-x:auto">
|
|
@@ -245,6 +246,10 @@ table.lb .best-cell{font-weight:700}
|
|
| 245 |
<th>Judge</th>
|
| 246 |
<th>R@1</th>
|
| 247 |
<th>sR@1</th>
|
|
|
|
|
|
|
|
|
|
|
|
|
| 248 |
</tr></thead>
|
| 249 |
<tbody id="st-body"></tbody>
|
| 250 |
</table>
|
|
@@ -634,12 +639,12 @@ async function loadStructureAnalysis() {
|
|
| 634 |
const spDiv = document.getElementById('st-spearman');
|
| 635 |
const pDiv = document.getElementById('st-pearson');
|
| 636 |
if (!tbody) return;
|
| 637 |
-
tbody.innerHTML = '<tr><td colspan="
|
| 638 |
try {
|
| 639 |
const r = await fetch('/api/rag_eval/structure_analysis?strategy=' + encodeURIComponent(strat)).then(x => x.json());
|
| 640 |
structureAnalysisCache = r;
|
| 641 |
if (!r.exists || !r.joined_rows) {
|
| 642 |
-
tbody.innerHTML = '<tr><td colspan="
|
| 643 |
if (disc) { disc.style.display = 'none'; }
|
| 644 |
if (spDiv) spDiv.innerHTML = '';
|
| 645 |
if (pDiv) pDiv.innerHTML = '';
|
|
@@ -660,6 +665,10 @@ async function loadStructureAnalysis() {
|
|
| 660 |
<td>${fmtF(row.rag_judge_score, 2)}</td>
|
| 661 |
<td>${fmtF(row['rag_recall@1'], 3)}</td>
|
| 662 |
<td>${fmtF(row['rag_source_recall@1'], 3)}</td>
|
|
|
|
|
|
|
|
|
|
|
|
|
| 663 |
</tr>`;
|
| 664 |
}).join('');
|
| 665 |
const sv = an.spearman_structure_vs_rag_judge || {};
|
|
@@ -689,7 +698,7 @@ async function loadStructureAnalysis() {
|
|
| 689 |
} else if (pDiv) pDiv.innerHTML = '<span class="empty">Need ≥2 rows with all structure metrics.</span>';
|
| 690 |
renderScatter(r.joined_rows);
|
| 691 |
} catch (e) {
|
| 692 |
-
tbody.innerHTML = '<tr><td colspan="
|
| 693 |
}
|
| 694 |
}
|
| 695 |
|
|
|
|
| 233 |
<h3>Structure vs gold + RAG (same retrieval strategy)</h3>
|
| 234 |
<p style="font-size:12px;color:var(--gray);margin:-6px 0 14px;line-height:1.45">
|
| 235 |
Rows: GNN chunking runs (<code>graph_*</code> / <code>graph_row_*</code>) with human labels on held-out sheets.
|
| 236 |
+
<strong>Inference timing</strong> is wall-clock <code>predict_sheet</code> per sheet (graph build + forward + post); ms/node uses non-empty <code>value</code> cells as graph nodes (same as the GNN).
|
| 237 |
<strong>Oracle RAG scores</strong> (human chunking) measure a ceiling that still depends on retrieval, generation, and Q&A difficulty — not only perfect structure.
|
| 238 |
</p>
|
| 239 |
<div style="overflow-x:auto">
|
|
|
|
| 246 |
<th>Judge</th>
|
| 247 |
<th>R@1</th>
|
| 248 |
<th>sR@1</th>
|
| 249 |
+
<th title="Mean wall ms per graph node (non-empty value cells); predict_sheet">mean ms/node</th>
|
| 250 |
+
<th title="95th percentile ms per graph node across sheets">p95 ms/node</th>
|
| 251 |
+
<th title="Sum of successful predict_sheet wall times (s)">predict total s</th>
|
| 252 |
+
<th title="total_graph_nodes / predict_total_s">nodes/s</th>
|
| 253 |
</tr></thead>
|
| 254 |
<tbody id="st-body"></tbody>
|
| 255 |
</table>
|
|
|
|
| 639 |
const spDiv = document.getElementById('st-spearman');
|
| 640 |
const pDiv = document.getElementById('st-pearson');
|
| 641 |
if (!tbody) return;
|
| 642 |
+
tbody.innerHTML = '<tr><td colspan="11">Loading…</td></tr>';
|
| 643 |
try {
|
| 644 |
const r = await fetch('/api/rag_eval/structure_analysis?strategy=' + encodeURIComponent(strat)).then(x => x.json());
|
| 645 |
structureAnalysisCache = r;
|
| 646 |
if (!r.exists || !r.joined_rows) {
|
| 647 |
+
tbody.innerHTML = '<tr><td colspan="11" class="empty">No joined data. Run RAG eval with graph methods and ensure rag_eval_structure_metrics.json exists.</td></tr>';
|
| 648 |
if (disc) { disc.style.display = 'none'; }
|
| 649 |
if (spDiv) spDiv.innerHTML = '';
|
| 650 |
if (pDiv) pDiv.innerHTML = '';
|
|
|
|
| 665 |
<td>${fmtF(row.rag_judge_score, 2)}</td>
|
| 666 |
<td>${fmtF(row['rag_recall@1'], 3)}</td>
|
| 667 |
<td>${fmtF(row['rag_source_recall@1'], 3)}</td>
|
| 668 |
+
<td>${row.inf_mean_ms_per_graph_node != null ? fmtF(row.inf_mean_ms_per_graph_node, 2) : '—'}</td>
|
| 669 |
+
<td>${row.inf_p95_ms_per_graph_node != null ? fmtF(row.inf_p95_ms_per_graph_node, 2) : '—'}</td>
|
| 670 |
+
<td>${row.inf_predict_total_s != null ? fmtF(row.inf_predict_total_s, 3) : '—'}</td>
|
| 671 |
+
<td>${row.inf_throughput_nodes_per_s != null ? fmtF(row.inf_throughput_nodes_per_s, 1) : '—'}</td>
|
| 672 |
</tr>`;
|
| 673 |
}).join('');
|
| 674 |
const sv = an.spearman_structure_vs_rag_judge || {};
|
|
|
|
| 698 |
} else if (pDiv) pDiv.innerHTML = '<span class="empty">Need ≥2 rows with all structure metrics.</span>';
|
| 699 |
renderScatter(r.joined_rows);
|
| 700 |
} catch (e) {
|
| 701 |
+
tbody.innerHTML = '<tr><td colspan="11">Error: ' + esc(String(e)) + '</td></tr>';
|
| 702 |
}
|
| 703 |
}
|
| 704 |
|
static/train.html
CHANGED
|
@@ -146,6 +146,15 @@ body{font-family:-apple-system,'Segoe UI',system-ui,sans-serif;background:var(--
|
|
| 146 |
.filter-bar .filter-chip:hover{border-color:var(--text-3);color:var(--text-1)}
|
| 147 |
.filter-bar .filter-chip.active{border-color:var(--blue);color:var(--blue);background:rgba(59,130,246,.08)}
|
| 148 |
.compare-btn{margin-left:auto}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 149 |
|
| 150 |
/* ── Buttons ── */
|
| 151 |
.btn{padding:8px 20px;font-size:12px;font-weight:700;border:none;border-radius:var(--radius-sm);cursor:pointer;
|
|
@@ -231,6 +240,7 @@ body{font-family:-apple-system,'Segoe UI',system-ui,sans-serif;background:var(--
|
|
| 231 |
.form-group label{font-size:10px;color:var(--text-3);text-transform:uppercase;letter-spacing:.5px;font-weight:700}
|
| 232 |
.form-group input,.form-group select{padding:8px 12px;background:var(--bg-0);border:1px solid var(--border);
|
| 233 |
border-radius:var(--radius-sm);color:var(--text-1);font-size:12px;font-family:inherit}
|
|
|
|
| 234 |
.form-group input:focus,.form-group select:focus{outline:none;border-color:var(--blue)}
|
| 235 |
.log-terminal{background:var(--bg-0);border:1px solid var(--border);border-radius:var(--radius);padding:16px;
|
| 236 |
font-family:'SF Mono','Fira Code','Cascadia Code',monospace;font-size:11px;line-height:1.7;
|
|
@@ -477,6 +487,8 @@ body{font-family:-apple-system,'Segoe UI',system-ui,sans-serif;background:var(--
|
|
| 477 |
<div class="form-group" id="fg-seeds"><label>Seeds</label><input id="train-seeds" value="42"></div>
|
| 478 |
<div class="form-group" id="fg-folds"><label>Folds</label><input id="train-folds" type="number" value="5" min="2" max="10"></div>
|
| 479 |
<div class="form-group"><label>Epochs</label><input id="train-epochs" type="number" value="100" min="10" max="1000"></div>
|
|
|
|
|
|
|
| 480 |
</div>
|
| 481 |
<div style="display:flex;gap:12px;align-items:center">
|
| 482 |
<button class="btn btn-primary" id="start-btn" onclick="startTraining()">
|
|
@@ -486,6 +498,80 @@ body{font-family:-apple-system,'Segoe UI',system-ui,sans-serif;background:var(--
|
|
| 486 |
</div>
|
| 487 |
</div>
|
| 488 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 489 |
<div class="card" id="train-progress-card-local" style="display:none">
|
| 490 |
<div class="card-header">
|
| 491 |
<div class="card-title">Run progress (this machine)</div>
|
|
@@ -508,9 +594,10 @@ body{font-family:-apple-system,'Segoe UI',system-ui,sans-serif;background:var(--
|
|
| 508 |
<span id="suite-gcp-meta" style="font-size:11px;color:var(--text-4)"></span>
|
| 509 |
</div>
|
| 510 |
<p style="font-size:12px;color:var(--text-3);margin-bottom:12px">
|
| 511 |
-
Live view of <code style="font-size:11px">experiments/*/progress.json</code> on GCS. Workers must set
|
| 512 |
-
<code style="font-size:11px">EXCEL_CHUNKER_DATA_ROOT</code> so each job writes under your bucket;
|
| 513 |
-
<code style="font-size:11px">EXCEL_CHUNKER_GCS_EXPERIMENTS_URI</code>.
|
|
|
|
| 514 |
</p>
|
| 515 |
<div id="suite-gcp-error" style="display:none;font-size:12px;color:var(--red);margin-bottom:10px"></div>
|
| 516 |
<div class="suite-overall">
|
|
@@ -531,7 +618,9 @@ body{font-family:-apple-system,'Segoe UI',system-ui,sans-serif;background:var(--
|
|
| 531 |
<div class="card-title">Embeddings</div>
|
| 532 |
</div>
|
| 533 |
<p style="font-size:12px;color:var(--text-3);margin-bottom:12px">
|
| 534 |
-
|
|
|
|
|
|
|
| 535 |
</p>
|
| 536 |
<div style="display:flex;gap:12px;align-items:center">
|
| 537 |
<button class="btn btn-secondary btn-sm" id="embed-btn" onclick="doEmbed()">Generate Embeddings</button>
|
|
@@ -553,6 +642,7 @@ let selectedRunIds = new Set();
|
|
| 553 |
let currentSort = {col: 'macro_f1', dir: 'desc'};
|
| 554 |
let archFilter = 'all';
|
| 555 |
let pollTimer = null;
|
|
|
|
| 556 |
|
| 557 |
const LABEL_COLORS = {
|
| 558 |
value:'#3b82f6', attribute:'#22c55e', aggregation:'#f59e0b', metadata:'#a855f7',
|
|
@@ -568,6 +658,79 @@ const ARCH_COLORS = {gat:'#3b82f6', gcn:'#a855f7', mlp:'#f59e0b', transformer:'#
|
|
| 568 |
naive_t0:'#64748b', naive_t1:'#94a3b8', naive_t2:'#cbd5e1'};
|
| 569 |
const METRIC_COLORS = {accuracy:'#3b82f6', macro_f1:'#34d399', chunk_f1:'#f59e0b', ari:'#a855f7'};
|
| 570 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 571 |
/* ──────────────────────────── API ──────────────────────────── */
|
| 572 |
function headers() {
|
| 573 |
const h = {'Content-Type':'application/json'};
|
|
@@ -596,6 +759,11 @@ window.navigate = function(page) {
|
|
| 596 |
if (page === 'experiments') loadExperiments();
|
| 597 |
if (page === 'models') loadModels();
|
| 598 |
if (page === 'compare') renderCompare();
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 599 |
};
|
| 600 |
document.querySelectorAll('.nav-item').forEach(n => n.addEventListener('click', () => navigate(n.dataset.page)));
|
| 601 |
|
|
@@ -1185,12 +1353,15 @@ document.getElementById('train-mode').addEventListener('change', e => {
|
|
| 1185 |
});
|
| 1186 |
|
| 1187 |
window.startTraining = async function() {
|
|
|
|
| 1188 |
const body = {
|
| 1189 |
mode: document.getElementById('train-mode').value,
|
| 1190 |
architectures: document.getElementById('train-arch').value,
|
| 1191 |
seeds: document.getElementById('train-seeds').value,
|
| 1192 |
folds: parseInt(document.getElementById('train-folds').value),
|
| 1193 |
epochs: parseInt(document.getElementById('train-epochs').value),
|
|
|
|
|
|
|
| 1194 |
};
|
| 1195 |
const res = await api('/api/train/start', {method:'POST', body:JSON.stringify(body)});
|
| 1196 |
if (res && res.started) {
|
|
@@ -1201,6 +1372,162 @@ window.startTraining = async function() {
|
|
| 1201 |
}
|
| 1202 |
};
|
| 1203 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1204 |
function startPolling() {
|
| 1205 |
if (pollTimer) clearInterval(pollTimer);
|
| 1206 |
pollTimer = setInterval(pollStatus, 2000);
|
|
@@ -1253,9 +1580,12 @@ function renderSuiteProgress(suite, mode) {
|
|
| 1253 |
? `${src ? src + ' · ' : ''}Overall: ${done} / ${total} units (${pct}%)`
|
| 1254 |
: `${src ? src + ' · ' : ''}Last: ${done} / ${total} units (${pct}%)`;
|
| 1255 |
document.getElementById(ids.bar).style.width = pct + '%';
|
|
|
|
|
|
|
|
|
|
| 1256 |
meta.textContent = suite.started_at
|
| 1257 |
? (suite.finished_at ? `${suite.started_at.slice(11,19)} → ${suite.finished_at.slice(11,19)} UTC` : 'Running…')
|
| 1258 |
-
: (suite.gcs_uri ?
|
| 1259 |
|
| 1260 |
const units = suite.units || [];
|
| 1261 |
const host = document.getElementById(ids.units);
|
|
@@ -1316,7 +1646,7 @@ function renderGcpProgressPanel(data) {
|
|
| 1316 |
}
|
| 1317 |
|
| 1318 |
async function pollStatus() {
|
| 1319 |
-
const data = await api('/api/train/status');
|
| 1320 |
if (!data) return;
|
| 1321 |
const badge = document.getElementById('train-badge');
|
| 1322 |
const globalBadge = document.getElementById('global-status');
|
|
@@ -1334,12 +1664,16 @@ async function pollStatus() {
|
|
| 1334 |
if (data.result.status === 'completed') {
|
| 1335 |
badge.className = 'status-pill done'; badge.textContent = 'Completed';
|
| 1336 |
globalBadge.className = 'status-pill done'; globalBadge.textContent = 'Done';
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1337 |
} else {
|
| 1338 |
badge.className = 'status-pill idle'; badge.textContent = 'Error';
|
| 1339 |
globalBadge.className = 'status-pill idle'; globalBadge.textContent = 'Error';
|
| 1340 |
}
|
| 1341 |
startBtn.disabled = false;
|
| 1342 |
-
|
| 1343 |
} else {
|
| 1344 |
badge.className = 'status-pill idle'; badge.textContent = 'Idle';
|
| 1345 |
globalBadge.className = 'status-pill idle'; globalBadge.textContent = 'Idle';
|
|
@@ -1348,7 +1682,7 @@ async function pollStatus() {
|
|
| 1348 |
|
| 1349 |
const lines = data.log || [];
|
| 1350 |
logBox.innerHTML = lines.map(l => {
|
| 1351 |
-
if (l.startsWith('ERROR')) return `<span class="l-err">${esc(l)}</span>`;
|
| 1352 |
if (l.includes('completed') || l.includes('saved') || l.includes('Best')) return `<span class="l-ok">${esc(l)}</span>`;
|
| 1353 |
if (l.startsWith('Training') || l.startsWith('===')) return `<span class="l-info">${esc(l)}</span>`;
|
| 1354 |
return `<span class="l-dim">${esc(l)}</span>`;
|
|
@@ -1392,14 +1726,15 @@ window.doEmbed = async function() {
|
|
| 1392 |
/* ──────────────────────────── Init ──────────────────────────── */
|
| 1393 |
async function init() {
|
| 1394 |
loadOverview();
|
|
|
|
|
|
|
|
|
|
| 1395 |
const status = await api('/api/train/status');
|
| 1396 |
if (status && status.running) {
|
| 1397 |
navigate('train');
|
| 1398 |
-
startPolling();
|
| 1399 |
-
} else {
|
| 1400 |
-
// Last run may already be finished; still paint log + badges (otherwise refresh shows empty).
|
| 1401 |
-
pollStatus();
|
| 1402 |
}
|
|
|
|
|
|
|
| 1403 |
}
|
| 1404 |
|
| 1405 |
init();
|
|
|
|
| 146 |
.filter-bar .filter-chip:hover{border-color:var(--text-3);color:var(--text-1)}
|
| 147 |
.filter-bar .filter-chip.active{border-color:var(--blue);color:var(--blue);background:rgba(59,130,246,.08)}
|
| 148 |
.compare-btn{margin-left:auto}
|
| 149 |
+
.gcp-check-grid{display:flex;flex-wrap:wrap;gap:8px 14px;align-items:center;margin-top:6px;padding:10px 12px;background:var(--bg-1);border:1px solid var(--border);border-radius:var(--radius-sm)}
|
| 150 |
+
.gcp-check-grid label{display:inline-flex;align-items:center;gap:6px;font-size:12px;color:var(--text-2);cursor:pointer;user-select:none;font-weight:500}
|
| 151 |
+
.gcp-check-grid input[type=checkbox]{width:14px;height:14px;accent-color:var(--blue);cursor:pointer}
|
| 152 |
+
.gcp-gpu-matrix{width:100%;max-width:520px;border-collapse:collapse;font-size:12px;margin-top:8px;background:var(--bg-1);border:1px solid var(--border);border-radius:var(--radius-sm);overflow:hidden}
|
| 153 |
+
.gcp-gpu-matrix th,.gcp-gpu-matrix td{padding:8px 10px;border-bottom:1px solid var(--border);text-align:center}
|
| 154 |
+
.gcp-gpu-matrix th{background:var(--bg-2);color:var(--text-3);font-size:10px;font-weight:700;text-transform:uppercase;letter-spacing:.4px}
|
| 155 |
+
.gcp-gpu-matrix td:first-child{text-align:left;font-weight:600;color:var(--text-2);white-space:nowrap}
|
| 156 |
+
.gcp-gpu-matrix tr:last-child td{border-bottom:none}
|
| 157 |
+
.gcp-gpu-matrix input[type=radio]{width:15px;height:15px;accent-color:var(--blue);cursor:pointer;vertical-align:middle}
|
| 158 |
|
| 159 |
/* ── Buttons ── */
|
| 160 |
.btn{padding:8px 20px;font-size:12px;font-weight:700;border:none;border-radius:var(--radius-sm);cursor:pointer;
|
|
|
|
| 240 |
.form-group label{font-size:10px;color:var(--text-3);text-transform:uppercase;letter-spacing:.5px;font-weight:700}
|
| 241 |
.form-group input,.form-group select{padding:8px 12px;background:var(--bg-0);border:1px solid var(--border);
|
| 242 |
border-radius:var(--radius-sm);color:var(--text-1);font-size:12px;font-family:inherit}
|
| 243 |
+
.form-group input.gcp-wide{max-width:100%}
|
| 244 |
.form-group input:focus,.form-group select:focus{outline:none;border-color:var(--blue)}
|
| 245 |
.log-terminal{background:var(--bg-0);border:1px solid var(--border);border-radius:var(--radius);padding:16px;
|
| 246 |
font-family:'SF Mono','Fira Code','Cascadia Code',monospace;font-size:11px;line-height:1.7;
|
|
|
|
| 487 |
<div class="form-group" id="fg-seeds"><label>Seeds</label><input id="train-seeds" value="42"></div>
|
| 488 |
<div class="form-group" id="fg-folds"><label>Folds</label><input id="train-folds" type="number" value="5" min="2" max="10"></div>
|
| 489 |
<div class="form-group"><label>Epochs</label><input id="train-epochs" type="number" value="100" min="10" max="1000"></div>
|
| 490 |
+
<div class="form-group"><label>Data subset (%)</label><input id="train-subset-percent" type="number" value="100" min="1" max="100" step="1" title="Random share of labeled sheets after filters (100 = all). Smaller = faster."></div>
|
| 491 |
+
<div class="form-group"><label>Subset seed</label><input id="train-subset-seed" type="number" value="0" min="0" step="1" title="Reproducible random pick when subset < 100%."></div>
|
| 492 |
</div>
|
| 493 |
<div style="display:flex;gap:12px;align-items:center">
|
| 494 |
<button class="btn btn-primary" id="start-btn" onclick="startTraining()">
|
|
|
|
| 498 |
</div>
|
| 499 |
</div>
|
| 500 |
|
| 501 |
+
<div class="card" id="gcp-vertex-card">
|
| 502 |
+
<div class="card-header">
|
| 503 |
+
<div class="card-title">Vertex AI (GCP)</div>
|
| 504 |
+
</div>
|
| 505 |
+
<p style="font-size:12px;color:var(--text-3);margin-bottom:14px;line-height:1.5">
|
| 506 |
+
Obraz Dockera <strong>nie</strong> siedzi „na Macu” jako usługa: budujesz go lokalnie
|
| 507 |
+
(<code style="font-size:11px">docker build -f Dockerfile.training …</code>),
|
| 508 |
+
<strong>tagujesz</strong> i <strong>wypychasz</strong> do Artifact Registry — stamtąd Vertex pobiera obraz przy starcie joba.
|
| 509 |
+
<strong>Typ GPU (T4 / L4 / …)</strong> nie wymaga przebudowy obrazu — ten sam tag w Registry; zmieniasz tylko maszynę Vertex.
|
| 510 |
+
Tu wysyłasz tylko żądania uruchomienia jobów (wymaga <code style="font-size:11px">gcloud auth application-default login</code>
|
| 511 |
+
lub konta serwisowego z uprawnieniami Vertex).
|
| 512 |
+
</p>
|
| 513 |
+
<div class="form-grid">
|
| 514 |
+
<div class="form-group"><label>Project ID</label><input id="gcp-project" value="acoustic-atom-386613" autocomplete="off"></div>
|
| 515 |
+
<div class="form-group"><label>Region</label><input id="gcp-region" value="europe-west1" autocomplete="off"></div>
|
| 516 |
+
<div class="form-group" style="grid-column:1/-1"><label>GCS data URI</label><input id="gcp-gcs-data" class="gcp-wide" value="gs://excel-chunker/data" autocomplete="off"></div>
|
| 517 |
+
<div class="form-group" style="grid-column:1/-1"><label>Training image (Artifact Registry)</label><input id="gcp-image" class="gcp-wide" value="europe-west1-docker.pkg.dev/acoustic-atom-386613/excel-chunker/excel-chunker-train:latest" autocomplete="off"></div>
|
| 518 |
+
<div class="form-group" style="grid-column:1/-1">
|
| 519 |
+
<label>Architectures</label>
|
| 520 |
+
<p style="font-size:11px;color:var(--text-4);margin:0 0 8px">Zaznacz, które modele wysłać na Vertex (osobno od pola „Launch Training” powyżej — tamtego nie nadpisujemy przy każdym wejściu na tę stronę).</p>
|
| 521 |
+
<div class="gcp-check-grid" id="gcp-arch-grid" aria-label="Architectures"></div>
|
| 522 |
+
<button type="button" class="btn btn-secondary btn-sm" style="margin-top:8px;font-size:11px" onclick="syncGcpArchFromTrainField()">Skopiuj architektury z pola Launch Training</button>
|
| 523 |
+
</div>
|
| 524 |
+
<div class="form-group" style="grid-column:1/-1">
|
| 525 |
+
<label>Seeds</label>
|
| 526 |
+
<div class="gcp-check-grid" id="gcp-seed-grid" aria-label="Seeds"></div>
|
| 527 |
+
</div>
|
| 528 |
+
<div class="form-group"><label>Folds</label><input id="gcp-folds" type="number" min="2" max="10" value="5"></div>
|
| 529 |
+
<div class="form-group"><label>Epochs</label><input id="gcp-epochs" type="number" min="2" max="1000" value="100"></div>
|
| 530 |
+
<div class="form-group"><label>Data subset (%)</label><input id="gcp-subset-percent" type="number" value="100" min="1" max="100" step="1" title="Same as local Train card — random share of labeled sheets on the worker."></div>
|
| 531 |
+
<div class="form-group" style="grid-column:1/-1">
|
| 532 |
+
<label>GPU na architekturę (Vertex)</label>
|
| 533 |
+
<p style="font-size:11px;color:var(--text-4);margin:6px 0 0">Wybierz <strong>jedną</strong> kolumnę (T4 / L4 / A100) dla każdej architektury. Dotyczy tylko zaznaczonych wyżej „Architectures” przy starcie jobów.</p>
|
| 534 |
+
<table class="gcp-gpu-matrix" aria-label="GPU per architecture">
|
| 535 |
+
<thead><tr><th>Architektura</th><th>T4</th><th>L4</th><th>A100</th></tr></thead>
|
| 536 |
+
<tbody>
|
| 537 |
+
<tr><td>GAT</td><td><input type="radio" name="gcp-gpu-row-gat" value="t4" checked></td><td><input type="radio" name="gcp-gpu-row-gat" value="l4"></td><td><input type="radio" name="gcp-gpu-row-gat" value="a100"></td></tr>
|
| 538 |
+
<tr><td>GCN</td><td><input type="radio" name="gcp-gpu-row-gcn" value="t4" checked></td><td><input type="radio" name="gcp-gpu-row-gcn" value="l4"></td><td><input type="radio" name="gcp-gpu-row-gcn" value="a100"></td></tr>
|
| 539 |
+
<tr><td>MLP</td><td><input type="radio" name="gcp-gpu-row-mlp" value="t4" checked></td><td><input type="radio" name="gcp-gpu-row-mlp" value="l4"></td><td><input type="radio" name="gcp-gpu-row-mlp" value="a100"></td></tr>
|
| 540 |
+
<tr><td>Adj. transformer</td><td><input type="radio" name="gcp-gpu-row-adj_transformer" value="t4"></td><td><input type="radio" name="gcp-gpu-row-adj_transformer" value="l4" checked></td><td><input type="radio" name="gcp-gpu-row-adj_transformer" value="a100"></td></tr>
|
| 541 |
+
<tr><td>Spatial edge</td><td><input type="radio" name="gcp-gpu-row-spatial_edge_transformer" value="t4"></td><td><input type="radio" name="gcp-gpu-row-spatial_edge_transformer" value="l4" checked></td><td><input type="radio" name="gcp-gpu-row-spatial_edge_transformer" value="a100"></td></tr>
|
| 542 |
+
<tr><td>Dual-modality</td><td><input type="radio" name="gcp-gpu-row-dual_modality_gnn" value="t4"></td><td><input type="radio" name="gcp-gpu-row-dual_modality_gnn" value="l4" checked></td><td><input type="radio" name="gcp-gpu-row-dual_modality_gnn" value="a100"></td></tr>
|
| 543 |
+
</tbody>
|
| 544 |
+
</table>
|
| 545 |
+
<p style="font-size:11px;color:var(--text-4);margin-top:8px">A100 wymaga osobnej quota (np. Custom model training A100). Maszyna dopasowywana automatycznie (np. <code style="font-size:10px">a2-highgpu-1g</code>).</p>
|
| 546 |
+
</div>
|
| 547 |
+
<div class="form-group"><label>GPU count (na job)</label><input id="gcp-accel-count" type="number" min="1" max="8" value="1" title="Zwykle 1 na job"></div>
|
| 548 |
+
<div class="form-group" style="grid-column:1/-1">
|
| 549 |
+
<div class="gcp-check-grid" style="margin-top:0">
|
| 550 |
+
<label title="SPOT = preemptible VM — tańsze, mogą przerywać job. Użyj gdy masz quota preemptible (np. P100 / L4).">
|
| 551 |
+
<input type="checkbox" id="gcp-vertex-spot"> Spot (preemptible)
|
| 552 |
+
</label>
|
| 553 |
+
</div>
|
| 554 |
+
</div>
|
| 555 |
+
</div>
|
| 556 |
+
<div style="display:flex;gap:12px;align-items:center;flex-wrap:wrap;margin-top:8px">
|
| 557 |
+
<button type="button" class="btn btn-primary" id="gcp-launch-btn" onclick="launchGcpVertex()">Uruchom joby na Vertex</button>
|
| 558 |
+
<span id="gcp-launch-msg" style="font-size:12px;color:var(--text-3)"></span>
|
| 559 |
+
</div>
|
| 560 |
+
<p style="font-size:11px;color:var(--text-4);margin-top:10px">Anulowanie: Vertex → Custom jobs → job → Cancel, albo <code style="font-size:10px">gcloud ai custom-jobs cancel JOB_ID --region=REGION --project=...</code>.</p>
|
| 561 |
+
<div id="gcp-docker-build-block" style="display:none;margin-top:16px;padding-top:16px;border-top:1px solid var(--border)">
|
| 562 |
+
<p style="font-size:12px;color:var(--text-3);margin-bottom:10px">
|
| 563 |
+
<strong>Build & push obrazu</strong> — tylko gdy aplikacja działa <strong>lokalnie</strong> na maszynie z Dockerem i ustawisz
|
| 564 |
+
<code style="font-size:11px">ENABLE_DOCKER_BUILD_FROM_DASHBOARD=1</code>. Wymaga wcześniej
|
| 565 |
+
<code style="font-size:11px">gcloud auth configure-docker REGION-docker.pkg.dev</code>.
|
| 566 |
+
</p>
|
| 567 |
+
<div style="display:flex;gap:12px;align-items:center;flex-wrap:wrap">
|
| 568 |
+
<button type="button" class="btn btn-secondary" id="gcp-docker-build-btn" onclick="startGcpDockerBuild()">Zbuduj i wypchnij obraz</button>
|
| 569 |
+
<span id="gcp-docker-build-msg" style="font-size:12px;color:var(--text-3)"></span>
|
| 570 |
+
</div>
|
| 571 |
+
<pre class="log-terminal compact" id="gcp-docker-log" style="margin-top:10px;max-height:200px;display:none"></pre>
|
| 572 |
+
</div>
|
| 573 |
+
</div>
|
| 574 |
+
|
| 575 |
<div class="card" id="train-progress-card-local" style="display:none">
|
| 576 |
<div class="card-header">
|
| 577 |
<div class="card-title">Run progress (this machine)</div>
|
|
|
|
| 594 |
<span id="suite-gcp-meta" style="font-size:11px;color:var(--text-4)"></span>
|
| 595 |
</div>
|
| 596 |
<p style="font-size:12px;color:var(--text-3);margin-bottom:12px">
|
| 597 |
+
Live view of <code style="font-size:11px">experiments/*/progress.json</code> on GCS (page auto-refreshes ~every 2s). Workers must set
|
| 598 |
+
<code style="font-size:11px">EXCEL_CHUNKER_DATA_ROOT</code> so each job writes under your bucket; set
|
| 599 |
+
<code style="font-size:11px">EXCEL_CHUNKER_GCS_EXPERIMENTS_URI</code> (or <code style="font-size:11px">GCP_GCS_DATA_URI</code>) so this card is enabled.
|
| 600 |
+
<strong style="color:var(--amber)">Re-run same arch/seed/fold?</strong> If <code style="font-size:11px">results.json</code> already exists on GCS, workers skip training — progress may flip to <code style="font-size:11px">skipped</code> with old metrics. Remove those run folders on GCS or use different seeds/folds to force a full retrain.
|
| 601 |
</p>
|
| 602 |
<div id="suite-gcp-error" style="display:none;font-size:12px;color:var(--red);margin-bottom:10px"></div>
|
| 603 |
<div class="suite-overall">
|
|
|
|
| 618 |
<div class="card-title">Embeddings</div>
|
| 619 |
</div>
|
| 620 |
<p style="font-size:12px;color:var(--text-3);margin-bottom:12px">
|
| 621 |
+
Text embeddings (E5) are generated <strong>automatically before each training run</strong> (local and Vertex)
|
| 622 |
+
for any labeled JSON that is missing or newer than its <code style="font-size:11px">*_embeddings.npz</code>.
|
| 623 |
+
Use the button below only if you want to run a full pass manually (e.g. after bulk label edits).
|
| 624 |
</p>
|
| 625 |
<div style="display:flex;gap:12px;align-items:center">
|
| 626 |
<button class="btn btn-secondary btn-sm" id="embed-btn" onclick="doEmbed()">Generate Embeddings</button>
|
|
|
|
| 642 |
let currentSort = {col: 'macro_f1', dir: 'desc'};
|
| 643 |
let archFilter = 'all';
|
| 644 |
let pollTimer = null;
|
| 645 |
+
let dockerBuildPollTimer = null;
|
| 646 |
|
| 647 |
const LABEL_COLORS = {
|
| 648 |
value:'#3b82f6', attribute:'#22c55e', aggregation:'#f59e0b', metadata:'#a855f7',
|
|
|
|
| 658 |
naive_t0:'#64748b', naive_t1:'#94a3b8', naive_t2:'#cbd5e1'};
|
| 659 |
const METRIC_COLORS = {accuracy:'#3b82f6', macro_f1:'#34d399', chunk_f1:'#f59e0b', ari:'#a855f7'};
|
| 660 |
|
| 661 |
+
/** Vertex launch: same arch set as local default train-arch; keep in sync with train_compare. */
|
| 662 |
+
const GCP_ARCH_CHOICES = [
|
| 663 |
+
['gat', 'GAT'],
|
| 664 |
+
['gcn', 'GCN'],
|
| 665 |
+
['mlp', 'MLP'],
|
| 666 |
+
['adj_transformer', 'Adj. transformer'],
|
| 667 |
+
['spatial_edge_transformer', 'Spatial edge'],
|
| 668 |
+
['dual_modality_gnn', 'Dual-modality GNN'],
|
| 669 |
+
];
|
| 670 |
+
const GCP_SEED_CHOICES = [42, 2137, 10042010];
|
| 671 |
+
|
| 672 |
+
function _ensureGcpCheckboxGrids() {
|
| 673 |
+
const ag = document.getElementById('gcp-arch-grid');
|
| 674 |
+
const sg = document.getElementById('gcp-seed-grid');
|
| 675 |
+
if (ag && !ag.dataset.built) {
|
| 676 |
+
ag.innerHTML = GCP_ARCH_CHOICES.map(([id, label]) =>
|
| 677 |
+
`<label><input type="checkbox" id="gcp-arch-${id}" data-arch="${id}"> ${esc(label)}</label>`
|
| 678 |
+
).join('');
|
| 679 |
+
ag.dataset.built = '1';
|
| 680 |
+
}
|
| 681 |
+
if (sg && !sg.dataset.built) {
|
| 682 |
+
sg.innerHTML = GCP_SEED_CHOICES.map(s =>
|
| 683 |
+
`<label><input type="checkbox" id="gcp-seed-${s}" data-seed="${s}"> ${s}</label>`
|
| 684 |
+
).join('');
|
| 685 |
+
sg.dataset.built = '1';
|
| 686 |
+
}
|
| 687 |
+
}
|
| 688 |
+
|
| 689 |
+
function gcpSelectedArchitectures() {
|
| 690 |
+
return GCP_ARCH_CHOICES.map(([id]) => id).filter(id => {
|
| 691 |
+
const el = document.getElementById(`gcp-arch-${id}`);
|
| 692 |
+
return el && el.checked;
|
| 693 |
+
});
|
| 694 |
+
}
|
| 695 |
+
|
| 696 |
+
function gcpSelectedSeeds() {
|
| 697 |
+
return GCP_SEED_CHOICES.filter(s => {
|
| 698 |
+
const el = document.getElementById(`gcp-seed-${s}`);
|
| 699 |
+
return el && el.checked;
|
| 700 |
+
});
|
| 701 |
+
}
|
| 702 |
+
|
| 703 |
+
function _applyGcpArchFromString(commaList) {
|
| 704 |
+
const want = new Set((commaList || '').split(',').map(s => s.trim()).filter(Boolean));
|
| 705 |
+
GCP_ARCH_CHOICES.forEach(([id]) => {
|
| 706 |
+
const el = document.getElementById(`gcp-arch-${id}`);
|
| 707 |
+
if (el) el.checked = want.has(id);
|
| 708 |
+
});
|
| 709 |
+
if (want.size === 0) {
|
| 710 |
+
const gat = document.getElementById('gcp-arch-gat');
|
| 711 |
+
if (gat) gat.checked = true;
|
| 712 |
+
}
|
| 713 |
+
}
|
| 714 |
+
|
| 715 |
+
/** Jawna synchronizacja checkboxów Vertex z polem #train-arch (nie wywołuj automatycznie przy każdym wejściu na Train). */
|
| 716 |
+
window.syncGcpArchFromTrainField = function() {
|
| 717 |
+
_ensureGcpCheckboxGrids();
|
| 718 |
+
const v = document.getElementById('train-arch') && document.getElementById('train-arch').value;
|
| 719 |
+
_applyGcpArchFromString(v || '');
|
| 720 |
+
};
|
| 721 |
+
|
| 722 |
+
function _applyGcpSeedsFromString(commaList) {
|
| 723 |
+
const nums = (commaList || '').split(',').map(s => parseInt(s.trim(), 10)).filter(n => !isNaN(n));
|
| 724 |
+
GCP_SEED_CHOICES.forEach(s => {
|
| 725 |
+
const el = document.getElementById(`gcp-seed-${s}`);
|
| 726 |
+
if (el) el.checked = nums.includes(s);
|
| 727 |
+
});
|
| 728 |
+
if (!GCP_SEED_CHOICES.some(s => document.getElementById(`gcp-seed-${s}`)?.checked)) {
|
| 729 |
+
const el = document.getElementById('gcp-seed-42');
|
| 730 |
+
if (el) el.checked = true;
|
| 731 |
+
}
|
| 732 |
+
}
|
| 733 |
+
|
| 734 |
/* ──────────────────────────── API ──────────────────────────── */
|
| 735 |
function headers() {
|
| 736 |
const h = {'Content-Type':'application/json'};
|
|
|
|
| 759 |
if (page === 'experiments') loadExperiments();
|
| 760 |
if (page === 'models') loadModels();
|
| 761 |
if (page === 'compare') renderCompare();
|
| 762 |
+
if (page === 'train') {
|
| 763 |
+
loadGcpTrainDefaults();
|
| 764 |
+
if (!pollTimer) startPolling();
|
| 765 |
+
/* Nie wywołuj tutaj _applyGcpArchFromString — nadpisywałoby checkboxy Vertex przy każdym powrocie na Train. */
|
| 766 |
+
}
|
| 767 |
};
|
| 768 |
document.querySelectorAll('.nav-item').forEach(n => n.addEventListener('click', () => navigate(n.dataset.page)));
|
| 769 |
|
|
|
|
| 1353 |
});
|
| 1354 |
|
| 1355 |
window.startTraining = async function() {
|
| 1356 |
+
const pct = parseFloat(document.getElementById('train-subset-percent').value);
|
| 1357 |
const body = {
|
| 1358 |
mode: document.getElementById('train-mode').value,
|
| 1359 |
architectures: document.getElementById('train-arch').value,
|
| 1360 |
seeds: document.getElementById('train-seeds').value,
|
| 1361 |
folds: parseInt(document.getElementById('train-folds').value),
|
| 1362 |
epochs: parseInt(document.getElementById('train-epochs').value),
|
| 1363 |
+
subset_percent: Number.isFinite(pct) ? Math.min(100, Math.max(1, pct)) : 100,
|
| 1364 |
+
subset_seed: parseInt(document.getElementById('train-subset-seed').value, 10) || 0,
|
| 1365 |
};
|
| 1366 |
const res = await api('/api/train/start', {method:'POST', body:JSON.stringify(body)});
|
| 1367 |
if (res && res.started) {
|
|
|
|
| 1372 |
}
|
| 1373 |
};
|
| 1374 |
|
| 1375 |
+
/** Same defaults as ``config.py`` (Vertex dashboard; local use). */
|
| 1376 |
+
const GCP_FORM_DEFAULTS = {
|
| 1377 |
+
project_id: 'acoustic-atom-386613',
|
| 1378 |
+
region: 'europe-west1',
|
| 1379 |
+
gcs_data_uri: 'gs://excel-chunker/data',
|
| 1380 |
+
image_uri: 'europe-west1-docker.pkg.dev/acoustic-atom-386613/excel-chunker/excel-chunker-train:latest',
|
| 1381 |
+
accelerator_count: 1,
|
| 1382 |
+
vertex_spot: false,
|
| 1383 |
+
};
|
| 1384 |
+
|
| 1385 |
+
/** Vertex accelerator enum z wiersza tabeli GPU (radio: t4 | l4 | a100). */
|
| 1386 |
+
function gcpGpuVertexEnumForArch(archId) {
|
| 1387 |
+
const sel = document.querySelector(`input[name="gcp-gpu-row-${archId}"]:checked`);
|
| 1388 |
+
const v = (sel && sel.value) || 't4';
|
| 1389 |
+
if (v === 'l4') return 'NVIDIA_L4';
|
| 1390 |
+
if (v === 'a100') return 'NVIDIA_TESLA_A100';
|
| 1391 |
+
return 'NVIDIA_TESLA_T4';
|
| 1392 |
+
}
|
| 1393 |
+
|
| 1394 |
+
async function loadGcpTrainDefaults() {
|
| 1395 |
+
_ensureGcpCheckboxGrids();
|
| 1396 |
+
const raw = await api('/api/train/gcp-defaults');
|
| 1397 |
+
const d = raw && !raw.detail ? raw : {};
|
| 1398 |
+
const set = (id, v) => {
|
| 1399 |
+
const el = document.getElementById(id);
|
| 1400 |
+
if (el && v != null && v !== '') el.value = v;
|
| 1401 |
+
};
|
| 1402 |
+
set('gcp-project', (d.project_id || GCP_FORM_DEFAULTS.project_id).trim());
|
| 1403 |
+
set('gcp-region', (d.region || GCP_FORM_DEFAULTS.region).trim());
|
| 1404 |
+
set('gcp-gcs-data', (d.gcs_data_uri || GCP_FORM_DEFAULTS.gcs_data_uri).trim());
|
| 1405 |
+
set('gcp-image', (d.image_uri || GCP_FORM_DEFAULTS.image_uri).trim());
|
| 1406 |
+
set('gcp-accel-count', String(d.accelerator_count != null ? d.accelerator_count : GCP_FORM_DEFAULTS.accelerator_count));
|
| 1407 |
+
const spotEl = document.getElementById('gcp-vertex-spot');
|
| 1408 |
+
if (spotEl) spotEl.checked = d.vertex_spot === true || d.vertex_spot === 1;
|
| 1409 |
+
const dock = document.getElementById('gcp-docker-build-block');
|
| 1410 |
+
if (dock) dock.style.display = d.docker_build_from_dashboard ? 'block' : 'none';
|
| 1411 |
+
_applyGcpSeedsFromString(document.getElementById('train-seeds').value);
|
| 1412 |
+
document.getElementById('gcp-folds').value = document.getElementById('train-folds').value;
|
| 1413 |
+
document.getElementById('gcp-epochs').value = document.getElementById('train-epochs').value;
|
| 1414 |
+
const gsp = document.getElementById('gcp-subset-percent');
|
| 1415 |
+
const tsp = document.getElementById('train-subset-percent');
|
| 1416 |
+
if (gsp && tsp) gsp.value = tsp.value;
|
| 1417 |
+
}
|
| 1418 |
+
|
| 1419 |
+
async function pollDockerBuildStatus() {
|
| 1420 |
+
const s = await api('/api/train/gcp-image-build/status');
|
| 1421 |
+
const logEl = document.getElementById('gcp-docker-log');
|
| 1422 |
+
const msg = document.getElementById('gcp-docker-build-msg');
|
| 1423 |
+
const btn = document.getElementById('gcp-docker-build-btn');
|
| 1424 |
+
if (!s || s.detail) return;
|
| 1425 |
+
if (s.log && s.log.length) {
|
| 1426 |
+
logEl.style.display = 'block';
|
| 1427 |
+
logEl.textContent = s.log.join('\n');
|
| 1428 |
+
logEl.scrollTop = logEl.scrollHeight;
|
| 1429 |
+
}
|
| 1430 |
+
if (!s.running && s.ok !== null && s.ok !== undefined) {
|
| 1431 |
+
if (dockerBuildPollTimer) { clearInterval(dockerBuildPollTimer); dockerBuildPollTimer = null; }
|
| 1432 |
+
btn.disabled = false;
|
| 1433 |
+
if (s.ok) {
|
| 1434 |
+
msg.textContent = 'Gotowe — obraz w Artifact Registry.';
|
| 1435 |
+
msg.style.color = 'var(--green)';
|
| 1436 |
+
} else {
|
| 1437 |
+
msg.textContent = s.error || 'Błąd';
|
| 1438 |
+
msg.style.color = 'var(--red)';
|
| 1439 |
+
}
|
| 1440 |
+
}
|
| 1441 |
+
}
|
| 1442 |
+
|
| 1443 |
+
window.startGcpDockerBuild = async function() {
|
| 1444 |
+
const msg = document.getElementById('gcp-docker-build-msg');
|
| 1445 |
+
const btn = document.getElementById('gcp-docker-build-btn');
|
| 1446 |
+
const logEl = document.getElementById('gcp-docker-log');
|
| 1447 |
+
msg.textContent = '';
|
| 1448 |
+
logEl.style.display = 'none';
|
| 1449 |
+
logEl.textContent = '';
|
| 1450 |
+
btn.disabled = true;
|
| 1451 |
+
const res = await api('/api/train/gcp-image-build', { method: 'POST', body: '{}' });
|
| 1452 |
+
if (!res) {
|
| 1453 |
+
btn.disabled = false;
|
| 1454 |
+
msg.textContent = 'Nie uruchomiono (403: dodaj ENABLE_DOCKER_BUILD_FROM_DASHBOARD=1 do .env i zrestartuj uvicorn; albo brak uprawnień).';
|
| 1455 |
+
msg.style.color = 'var(--red)';
|
| 1456 |
+
return;
|
| 1457 |
+
}
|
| 1458 |
+
if (res.detail) {
|
| 1459 |
+
btn.disabled = false;
|
| 1460 |
+
msg.textContent = _formatVertexErr(res);
|
| 1461 |
+
msg.style.color = 'var(--red)';
|
| 1462 |
+
return;
|
| 1463 |
+
}
|
| 1464 |
+
msg.textContent = 'Trwa build… (log poniżej)';
|
| 1465 |
+
msg.style.color = 'var(--blue)';
|
| 1466 |
+
if (dockerBuildPollTimer) clearInterval(dockerBuildPollTimer);
|
| 1467 |
+
dockerBuildPollTimer = setInterval(pollDockerBuildStatus, 1500);
|
| 1468 |
+
pollDockerBuildStatus();
|
| 1469 |
+
};
|
| 1470 |
+
|
| 1471 |
+
function _formatVertexErr(res) {
|
| 1472 |
+
if (!res) return 'Brak odpowiedzi';
|
| 1473 |
+
const d = res.detail;
|
| 1474 |
+
if (typeof d === 'string') return d;
|
| 1475 |
+
if (Array.isArray(d)) return d.map(x => (x.msg != null ? x.msg : JSON.stringify(x))).join('; ');
|
| 1476 |
+
return 'Błąd';
|
| 1477 |
+
}
|
| 1478 |
+
|
| 1479 |
+
window.launchGcpVertex = async function() {
|
| 1480 |
+
const msg = document.getElementById('gcp-launch-msg');
|
| 1481 |
+
const btn = document.getElementById('gcp-launch-btn');
|
| 1482 |
+
_ensureGcpCheckboxGrids();
|
| 1483 |
+
const archs = gcpSelectedArchitectures();
|
| 1484 |
+
const seeds = gcpSelectedSeeds();
|
| 1485 |
+
if (!archs.length) {
|
| 1486 |
+
msg.textContent = 'Zaznacz co najmniej jedną architekturę.';
|
| 1487 |
+
msg.style.color = 'var(--amber)';
|
| 1488 |
+
return;
|
| 1489 |
+
}
|
| 1490 |
+
if (!seeds.length) {
|
| 1491 |
+
msg.textContent = 'Zaznacz co najmniej jeden seed (42, 2137 lub 10042010).';
|
| 1492 |
+
msg.style.color = 'var(--amber)';
|
| 1493 |
+
return;
|
| 1494 |
+
}
|
| 1495 |
+
const gPct = parseFloat(document.getElementById('gcp-subset-percent').value);
|
| 1496 |
+
const archAccelerators = {};
|
| 1497 |
+
archs.forEach(arch => {
|
| 1498 |
+
archAccelerators[arch] = gcpGpuVertexEnumForArch(arch);
|
| 1499 |
+
});
|
| 1500 |
+
const firstAccel = archAccelerators[archs[0]] || 'NVIDIA_TESLA_T4';
|
| 1501 |
+
const body = {
|
| 1502 |
+
architectures: archs.join(','),
|
| 1503 |
+
seeds: seeds.join(','),
|
| 1504 |
+
folds: parseInt(document.getElementById('gcp-folds').value, 10),
|
| 1505 |
+
epochs: parseInt(document.getElementById('gcp-epochs').value, 10),
|
| 1506 |
+
subset_percent: Number.isFinite(gPct) ? Math.min(100, Math.max(1, gPct)) : 100,
|
| 1507 |
+
project_id: document.getElementById('gcp-project').value.trim() || null,
|
| 1508 |
+
region: document.getElementById('gcp-region').value.trim() || null,
|
| 1509 |
+
image_uri: document.getElementById('gcp-image').value.trim() || null,
|
| 1510 |
+
gcs_data_uri: document.getElementById('gcp-gcs-data').value.trim() || null,
|
| 1511 |
+
machine_type: 'n1-standard-8',
|
| 1512 |
+
accelerator_type: firstAccel,
|
| 1513 |
+
accelerator_count: parseInt(document.getElementById('gcp-accel-count').value, 10) || 1,
|
| 1514 |
+
vertex_spot: !!document.getElementById('gcp-vertex-spot')?.checked,
|
| 1515 |
+
arch_accelerators: archAccelerators,
|
| 1516 |
+
};
|
| 1517 |
+
msg.textContent = 'Wysyłanie…';
|
| 1518 |
+
msg.style.color = 'var(--blue)';
|
| 1519 |
+
btn.disabled = true;
|
| 1520 |
+
const res = await api('/api/train/gcp-launch', { method: 'POST', body: JSON.stringify(body) });
|
| 1521 |
+
btn.disabled = false;
|
| 1522 |
+
if (res && res.status === 'ok') {
|
| 1523 |
+
msg.textContent = `Wysłano ${res.submitted_count} job(ów). Sprawdź Vertex → Custom jobs.`;
|
| 1524 |
+
msg.style.color = 'var(--green)';
|
| 1525 |
+
} else {
|
| 1526 |
+
msg.textContent = _formatVertexErr(res);
|
| 1527 |
+
msg.style.color = 'var(--red)';
|
| 1528 |
+
}
|
| 1529 |
+
};
|
| 1530 |
+
|
| 1531 |
function startPolling() {
|
| 1532 |
if (pollTimer) clearInterval(pollTimer);
|
| 1533 |
pollTimer = setInterval(pollStatus, 2000);
|
|
|
|
| 1580 |
? `${src ? src + ' · ' : ''}Overall: ${done} / ${total} units (${pct}%)`
|
| 1581 |
: `${src ? src + ' · ' : ''}Last: ${done} / ${total} units (${pct}%)`;
|
| 1582 |
document.getElementById(ids.bar).style.width = pct + '%';
|
| 1583 |
+
const gcpRead = suite.gcs_uri && suite.refreshed_at
|
| 1584 |
+
? ` · last GCS read ${suite.refreshed_at.slice(0, 19).replace('T', ' ')} UTC`
|
| 1585 |
+
: '';
|
| 1586 |
meta.textContent = suite.started_at
|
| 1587 |
? (suite.finished_at ? `${suite.started_at.slice(11,19)} → ${suite.finished_at.slice(11,19)} UTC` : 'Running…')
|
| 1588 |
+
: (suite.gcs_uri ? `Poll ~2s · server cache ~2s${gcpRead}` : '');
|
| 1589 |
|
| 1590 |
const units = suite.units || [];
|
| 1591 |
const host = document.getElementById(ids.units);
|
|
|
|
| 1646 |
}
|
| 1647 |
|
| 1648 |
async function pollStatus() {
|
| 1649 |
+
const data = await api('/api/train/status', { cache: 'no-store' });
|
| 1650 |
if (!data) return;
|
| 1651 |
const badge = document.getElementById('train-badge');
|
| 1652 |
const globalBadge = document.getElementById('global-status');
|
|
|
|
| 1664 |
if (data.result.status === 'completed') {
|
| 1665 |
badge.className = 'status-pill done'; badge.textContent = 'Completed';
|
| 1666 |
globalBadge.className = 'status-pill done'; globalBadge.textContent = 'Done';
|
| 1667 |
+
} else if (data.result.status === 'failed') {
|
| 1668 |
+
const n = data.result.failed_units != null ? ` (${data.result.failed_units})` : '';
|
| 1669 |
+
badge.className = 'status-pill idle'; badge.textContent = 'Failed units' + n;
|
| 1670 |
+
globalBadge.className = 'status-pill idle'; globalBadge.textContent = 'Some folds failed';
|
| 1671 |
} else {
|
| 1672 |
badge.className = 'status-pill idle'; badge.textContent = 'Error';
|
| 1673 |
globalBadge.className = 'status-pill idle'; globalBadge.textContent = 'Error';
|
| 1674 |
}
|
| 1675 |
startBtn.disabled = false;
|
| 1676 |
+
/* Keep polling: Vertex/GCS progress and log tail should update without a local run. */
|
| 1677 |
} else {
|
| 1678 |
badge.className = 'status-pill idle'; badge.textContent = 'Idle';
|
| 1679 |
globalBadge.className = 'status-pill idle'; globalBadge.textContent = 'Idle';
|
|
|
|
| 1682 |
|
| 1683 |
const lines = data.log || [];
|
| 1684 |
logBox.innerHTML = lines.map(l => {
|
| 1685 |
+
if (l.startsWith('ERROR') || l.includes('failed unit')) return `<span class="l-err">${esc(l)}</span>`;
|
| 1686 |
if (l.includes('completed') || l.includes('saved') || l.includes('Best')) return `<span class="l-ok">${esc(l)}</span>`;
|
| 1687 |
if (l.startsWith('Training') || l.startsWith('===')) return `<span class="l-info">${esc(l)}</span>`;
|
| 1688 |
return `<span class="l-dim">${esc(l)}</span>`;
|
|
|
|
| 1726 |
/* ──────────────────────────── Init ──────────────────────────── */
|
| 1727 |
async function init() {
|
| 1728 |
loadOverview();
|
| 1729 |
+
_ensureGcpCheckboxGrids();
|
| 1730 |
+
_applyGcpArchFromString(document.getElementById('train-arch').value);
|
| 1731 |
+
await loadGcpTrainDefaults();
|
| 1732 |
const status = await api('/api/train/status');
|
| 1733 |
if (status && status.running) {
|
| 1734 |
navigate('train');
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1735 |
}
|
| 1736 |
+
/* Always poll: local training log, badges, and GCS progress (when EXCEL_CHUNKER_GCS_EXPERIMENTS_URI / data URI is set). */
|
| 1737 |
+
startPolling();
|
| 1738 |
}
|
| 1739 |
|
| 1740 |
init();
|
train_compare.py
CHANGED
|
@@ -19,6 +19,7 @@ from __future__ import annotations
|
|
| 19 |
import argparse
|
| 20 |
import json
|
| 21 |
import random
|
|
|
|
| 22 |
import time
|
| 23 |
from collections import defaultdict
|
| 24 |
from datetime import datetime, timezone
|
|
@@ -700,13 +701,18 @@ def run_experiments(
|
|
| 700 |
fold_index: int | None = None,
|
| 701 |
sheet_list_file: Path | str | None = None,
|
| 702 |
max_graphs: int | None = None,
|
|
|
|
| 703 |
subset_seed: int = 0,
|
| 704 |
progress: Any = None,
|
| 705 |
-
) -> dict:
|
| 706 |
"""Run the full experiment suite.
|
| 707 |
|
| 708 |
If *progress* is a ``training_progress.SuiteProgress`` instance, it is updated
|
| 709 |
for each (architecture, seed, fold) unit and each training epoch (dashboard).
|
|
|
|
|
|
|
|
|
|
|
|
|
| 710 |
"""
|
| 711 |
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 712 |
print(f"Device: {device}")
|
|
@@ -738,6 +744,16 @@ def run_experiments(
|
|
| 738 |
}
|
| 739 |
print(f" Sheet allowlist: {len(sheet_allowlist)} stems from {sheet_list_file}")
|
| 740 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 741 |
graphs = load_all_graphs(
|
| 742 |
max_nodes=max_nodes,
|
| 743 |
use_row_col_edges=use_row_col_edges,
|
|
@@ -745,6 +761,7 @@ def run_experiments(
|
|
| 745 |
graph_cfg=graph_cfg,
|
| 746 |
sheet_allowlist=sheet_allowlist,
|
| 747 |
max_graphs=max_graphs,
|
|
|
|
| 748 |
subset_seed=subset_seed,
|
| 749 |
)
|
| 750 |
loaded_splits: dict | None = None
|
|
@@ -764,7 +781,7 @@ def run_experiments(
|
|
| 764 |
if progress is not None:
|
| 765 |
progress.reset_for_new_run([])
|
| 766 |
progress.mark_suite_end(datetime.now(timezone.utc).isoformat())
|
| 767 |
-
return {}
|
| 768 |
|
| 769 |
in_dim = graphs[0].x.size(1)
|
| 770 |
class_weights = compute_class_weights(graphs)
|
|
@@ -825,7 +842,7 @@ def _run_experiment_suite_inner(
|
|
| 825 |
class_weights: torch.Tensor,
|
| 826 |
loaded_splits: dict | None,
|
| 827 |
fold_index: int | None,
|
| 828 |
-
) -> dict:
|
| 829 |
"""Body of ``run_experiments`` after data load (factored for ``try/finally`` on dashboard progress)."""
|
| 830 |
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 831 |
|
|
@@ -834,6 +851,7 @@ def _run_experiment_suite_inner(
|
|
| 834 |
SpatialEdgeTransformerModel.precompute_spatial(g)
|
| 835 |
|
| 836 |
all_results: dict[str, list] = defaultdict(list)
|
|
|
|
| 837 |
|
| 838 |
for seed in seeds:
|
| 839 |
random.seed(seed)
|
|
@@ -967,6 +985,7 @@ def _run_experiment_suite_inner(
|
|
| 967 |
progress_epoch_callback=_epoch_hook,
|
| 968 |
)
|
| 969 |
except Exception as e:
|
|
|
|
| 970 |
print(f" ERROR in {run_tag}: {e}")
|
| 971 |
print(f" Skipping {run_tag} and continuing...")
|
| 972 |
if progress is not None:
|
|
@@ -1250,7 +1269,12 @@ def _run_experiment_suite_inner(
|
|
| 1250 |
/ "best_model.pt")
|
| 1251 |
if src.exists():
|
| 1252 |
import shutil
|
| 1253 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1254 |
shutil.copy2(src, dst)
|
| 1255 |
print(f"\nBest model ({best['arch']} seed={best['seed']} "
|
| 1256 |
f"fold={best['fold']}) promoted to {dst}")
|
|
@@ -1265,7 +1289,15 @@ def _run_experiment_suite_inner(
|
|
| 1265 |
}, f, indent=2, default=str)
|
| 1266 |
print(f"Ranking saved to {summary_path}")
|
| 1267 |
|
| 1268 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1269 |
|
| 1270 |
|
| 1271 |
# ═══════════════════════════════════════════════════════════════════════════
|
|
@@ -1307,8 +1339,14 @@ def main():
|
|
| 1307 |
help="File with one labeled JSON stem per line (# comments allowed)")
|
| 1308 |
parser.add_argument("--max-graphs", type=int, default=None,
|
| 1309 |
help="After filters, randomly keep at most this many graphs")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1310 |
parser.add_argument("--subset-seed", type=int, default=0,
|
| 1311 |
-
help="RNG seed for --max-graphs subsampling")
|
| 1312 |
|
| 1313 |
# GraphConfig options
|
| 1314 |
parser.add_argument("--graph-config", type=str, default=None,
|
|
@@ -1349,7 +1387,7 @@ def main():
|
|
| 1349 |
if cli_overrides:
|
| 1350 |
graph_cfg = GraphConfig(**cli_overrides)
|
| 1351 |
|
| 1352 |
-
run_experiments(
|
| 1353 |
architectures=archs,
|
| 1354 |
seeds=seeds,
|
| 1355 |
n_folds=args.folds,
|
|
@@ -1369,8 +1407,11 @@ def main():
|
|
| 1369 |
fold_index=args.fold_index,
|
| 1370 |
sheet_list_file=args.sheet_list,
|
| 1371 |
max_graphs=args.max_graphs,
|
|
|
|
| 1372 |
subset_seed=args.subset_seed,
|
| 1373 |
)
|
|
|
|
|
|
|
| 1374 |
|
| 1375 |
|
| 1376 |
if __name__ == "__main__":
|
|
|
|
| 19 |
import argparse
|
| 20 |
import json
|
| 21 |
import random
|
| 22 |
+
import sys
|
| 23 |
import time
|
| 24 |
from collections import defaultdict
|
| 25 |
from datetime import datetime, timezone
|
|
|
|
| 701 |
fold_index: int | None = None,
|
| 702 |
sheet_list_file: Path | str | None = None,
|
| 703 |
max_graphs: int | None = None,
|
| 704 |
+
subset_percent: float | None = None,
|
| 705 |
subset_seed: int = 0,
|
| 706 |
progress: Any = None,
|
| 707 |
+
) -> tuple[dict, int]:
|
| 708 |
"""Run the full experiment suite.
|
| 709 |
|
| 710 |
If *progress* is a ``training_progress.SuiteProgress`` instance, it is updated
|
| 711 |
for each (architecture, seed, fold) unit and each training epoch (dashboard).
|
| 712 |
+
|
| 713 |
+
Returns ``(results_by_arch, failed_unit_count)``. Failed units are those where
|
| 714 |
+
``_train_one`` raised (e.g. CUDA OOM); the process used to exit 0 anyway — callers
|
| 715 |
+
should check ``failed_unit_count`` or rely on CLI ``sys.exit(1)``.
|
| 716 |
"""
|
| 717 |
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 718 |
print(f"Device: {device}")
|
|
|
|
| 744 |
}
|
| 745 |
print(f" Sheet allowlist: {len(sheet_allowlist)} stems from {sheet_list_file}")
|
| 746 |
|
| 747 |
+
try:
|
| 748 |
+
from embed_text import ensure_embeddings_for_labeled
|
| 749 |
+
|
| 750 |
+
ensure_embeddings_for_labeled()
|
| 751 |
+
except ImportError:
|
| 752 |
+
print(
|
| 753 |
+
" Warning: embed_text not available; training uses zero vectors where "
|
| 754 |
+
"*_embeddings.npz is missing."
|
| 755 |
+
)
|
| 756 |
+
|
| 757 |
graphs = load_all_graphs(
|
| 758 |
max_nodes=max_nodes,
|
| 759 |
use_row_col_edges=use_row_col_edges,
|
|
|
|
| 761 |
graph_cfg=graph_cfg,
|
| 762 |
sheet_allowlist=sheet_allowlist,
|
| 763 |
max_graphs=max_graphs,
|
| 764 |
+
subset_percent=subset_percent,
|
| 765 |
subset_seed=subset_seed,
|
| 766 |
)
|
| 767 |
loaded_splits: dict | None = None
|
|
|
|
| 781 |
if progress is not None:
|
| 782 |
progress.reset_for_new_run([])
|
| 783 |
progress.mark_suite_end(datetime.now(timezone.utc).isoformat())
|
| 784 |
+
return {}, 0
|
| 785 |
|
| 786 |
in_dim = graphs[0].x.size(1)
|
| 787 |
class_weights = compute_class_weights(graphs)
|
|
|
|
| 842 |
class_weights: torch.Tensor,
|
| 843 |
loaded_splits: dict | None,
|
| 844 |
fold_index: int | None,
|
| 845 |
+
) -> tuple[dict, int]:
|
| 846 |
"""Body of ``run_experiments`` after data load (factored for ``try/finally`` on dashboard progress)."""
|
| 847 |
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 848 |
|
|
|
|
| 851 |
SpatialEdgeTransformerModel.precompute_spatial(g)
|
| 852 |
|
| 853 |
all_results: dict[str, list] = defaultdict(list)
|
| 854 |
+
failed_units = 0
|
| 855 |
|
| 856 |
for seed in seeds:
|
| 857 |
random.seed(seed)
|
|
|
|
| 985 |
progress_epoch_callback=_epoch_hook,
|
| 986 |
)
|
| 987 |
except Exception as e:
|
| 988 |
+
failed_units += 1
|
| 989 |
print(f" ERROR in {run_tag}: {e}")
|
| 990 |
print(f" Skipping {run_tag} and continuing...")
|
| 991 |
if progress is not None:
|
|
|
|
| 1269 |
/ "best_model.pt")
|
| 1270 |
if src.exists():
|
| 1271 |
import shutil
|
| 1272 |
+
# Unique name: parallel Vertex jobs share GCS (gcsfuse); copying to the same
|
| 1273 |
+
# ``models/best_model.pt`` from two workers causes Errno 116 (stale file handle).
|
| 1274 |
+
dst = (
|
| 1275 |
+
config.MODELS_DIR
|
| 1276 |
+
/ f"best_model_{best['arch']}_s{best['seed']}_f{best['fold']}.pt"
|
| 1277 |
+
)
|
| 1278 |
shutil.copy2(src, dst)
|
| 1279 |
print(f"\nBest model ({best['arch']} seed={best['seed']} "
|
| 1280 |
f"fold={best['fold']}) promoted to {dst}")
|
|
|
|
| 1289 |
}, f, indent=2, default=str)
|
| 1290 |
print(f"Ranking saved to {summary_path}")
|
| 1291 |
|
| 1292 |
+
if failed_units:
|
| 1293 |
+
print(f"\n{'!'*80}")
|
| 1294 |
+
print(
|
| 1295 |
+
f"FAILED: {failed_units} training unit(s) raised an exception "
|
| 1296 |
+
f"(e.g. CUDA OOM). This run is incomplete."
|
| 1297 |
+
)
|
| 1298 |
+
print(f"{'!'*80}\n")
|
| 1299 |
+
|
| 1300 |
+
return dict(all_results), failed_units
|
| 1301 |
|
| 1302 |
|
| 1303 |
# ═══════════════════════════════════════════════════════════════════════════
|
|
|
|
| 1339 |
help="File with one labeled JSON stem per line (# comments allowed)")
|
| 1340 |
parser.add_argument("--max-graphs", type=int, default=None,
|
| 1341 |
help="After filters, randomly keep at most this many graphs")
|
| 1342 |
+
parser.add_argument(
|
| 1343 |
+
"--subset-percent",
|
| 1344 |
+
type=float,
|
| 1345 |
+
default=100.0,
|
| 1346 |
+
help="Random fraction of candidate sheets before CV (1-100; 100 = all, same as dashboard)",
|
| 1347 |
+
)
|
| 1348 |
parser.add_argument("--subset-seed", type=int, default=0,
|
| 1349 |
+
help="RNG seed for --subset-percent / --max-graphs subsampling")
|
| 1350 |
|
| 1351 |
# GraphConfig options
|
| 1352 |
parser.add_argument("--graph-config", type=str, default=None,
|
|
|
|
| 1387 |
if cli_overrides:
|
| 1388 |
graph_cfg = GraphConfig(**cli_overrides)
|
| 1389 |
|
| 1390 |
+
_, failed = run_experiments(
|
| 1391 |
architectures=archs,
|
| 1392 |
seeds=seeds,
|
| 1393 |
n_folds=args.folds,
|
|
|
|
| 1407 |
fold_index=args.fold_index,
|
| 1408 |
sheet_list_file=args.sheet_list,
|
| 1409 |
max_graphs=args.max_graphs,
|
| 1410 |
+
subset_percent=None if args.subset_percent >= 100.0 else args.subset_percent,
|
| 1411 |
subset_seed=args.subset_seed,
|
| 1412 |
)
|
| 1413 |
+
if failed:
|
| 1414 |
+
sys.exit(1)
|
| 1415 |
|
| 1416 |
|
| 1417 |
if __name__ == "__main__":
|
train_gnn.py
CHANGED
|
@@ -1811,6 +1811,7 @@ def load_all_graphs(max_nodes: int = 5000,
|
|
| 1811 |
graph_cfg: GraphConfig | None = None,
|
| 1812 |
sheet_allowlist: set[str] | None = None,
|
| 1813 |
max_graphs: int | None = None,
|
|
|
|
| 1814 |
subset_seed: int = 0) -> list[Data]:
|
| 1815 |
"""Load all labeled JSONs as PyG Data objects.
|
| 1816 |
|
|
@@ -1825,6 +1826,9 @@ def load_all_graphs(max_nodes: int = 5000,
|
|
| 1825 |
|
| 1826 |
*max_graphs*: after other filters, randomly keep at most this many graphs
|
| 1827 |
(reproducible with *subset_seed*).
|
|
|
|
|
|
|
|
|
|
| 1828 |
"""
|
| 1829 |
held_out = _load_rag_held_out_ids() if exclude_rag_holdout else set()
|
| 1830 |
if held_out:
|
|
@@ -1849,6 +1853,20 @@ def load_all_graphs(max_nodes: int = 5000,
|
|
| 1849 |
continue
|
| 1850 |
candidates.append(fp)
|
| 1851 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1852 |
if max_graphs is not None and len(candidates) > max_graphs:
|
| 1853 |
rng = np.random.RandomState(subset_seed)
|
| 1854 |
perm = rng.permutation(len(candidates))
|
|
@@ -2162,13 +2180,32 @@ def _next_model_version() -> str:
|
|
| 2162 |
return f"{last_num + 1:03d}"
|
| 2163 |
|
| 2164 |
|
| 2165 |
-
def train_model(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2166 |
from datetime import datetime, timezone
|
| 2167 |
|
| 2168 |
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 2169 |
print(f"Using device: {device}")
|
| 2170 |
|
| 2171 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2172 |
if len(graphs) < 2:
|
| 2173 |
print(f"Need at least 2 labeled sheets, found {len(graphs)}. Label more data first.")
|
| 2174 |
return
|
|
|
|
| 1811 |
graph_cfg: GraphConfig | None = None,
|
| 1812 |
sheet_allowlist: set[str] | None = None,
|
| 1813 |
max_graphs: int | None = None,
|
| 1814 |
+
subset_percent: float | None = None,
|
| 1815 |
subset_seed: int = 0) -> list[Data]:
|
| 1816 |
"""Load all labeled JSONs as PyG Data objects.
|
| 1817 |
|
|
|
|
| 1826 |
|
| 1827 |
*max_graphs*: after other filters, randomly keep at most this many graphs
|
| 1828 |
(reproducible with *subset_seed*).
|
| 1829 |
+
|
| 1830 |
+
*subset_percent*: if set and ``< 100``, keep a random ``pct%`` of candidates
|
| 1831 |
+
(after allowlist / RAG hold-out / max_nodes), then apply *max_graphs* if set.
|
| 1832 |
"""
|
| 1833 |
held_out = _load_rag_held_out_ids() if exclude_rag_holdout else set()
|
| 1834 |
if held_out:
|
|
|
|
| 1853 |
continue
|
| 1854 |
candidates.append(fp)
|
| 1855 |
|
| 1856 |
+
if subset_percent is not None and subset_percent < 100:
|
| 1857 |
+
n_cand = len(candidates)
|
| 1858 |
+
n_keep = max(1, int(round(n_cand * (subset_percent / 100.0))))
|
| 1859 |
+
if n_keep < n_cand:
|
| 1860 |
+
rng = np.random.RandomState(subset_seed)
|
| 1861 |
+
perm = rng.permutation(n_cand)
|
| 1862 |
+
sel_idx = perm[:n_keep]
|
| 1863 |
+
candidates = [candidates[i] for i in sel_idx]
|
| 1864 |
+
candidates.sort(key=lambda p: p.stem)
|
| 1865 |
+
print(
|
| 1866 |
+
f" Subset: {n_keep}/{n_cand} graphs ({subset_percent:g}% of candidates, "
|
| 1867 |
+
f"subset_seed={subset_seed})"
|
| 1868 |
+
)
|
| 1869 |
+
|
| 1870 |
if max_graphs is not None and len(candidates) > max_graphs:
|
| 1871 |
rng = np.random.RandomState(subset_seed)
|
| 1872 |
perm = rng.permutation(len(candidates))
|
|
|
|
| 2180 |
return f"{last_num + 1:03d}"
|
| 2181 |
|
| 2182 |
|
| 2183 |
+
def train_model(
|
| 2184 |
+
num_epochs: int = 100,
|
| 2185 |
+
lr: float = 1e-3,
|
| 2186 |
+
skip_predict: bool = False,
|
| 2187 |
+
*,
|
| 2188 |
+
subset_percent: float | None = None,
|
| 2189 |
+
subset_seed: int = 0,
|
| 2190 |
+
):
|
| 2191 |
from datetime import datetime, timezone
|
| 2192 |
|
| 2193 |
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 2194 |
print(f"Using device: {device}")
|
| 2195 |
|
| 2196 |
+
try:
|
| 2197 |
+
from embed_text import ensure_embeddings_for_labeled
|
| 2198 |
+
|
| 2199 |
+
ensure_embeddings_for_labeled()
|
| 2200 |
+
except ImportError:
|
| 2201 |
+
print(
|
| 2202 |
+
"Warning: embed_text not available; using zero vectors where *_embeddings.npz is missing."
|
| 2203 |
+
)
|
| 2204 |
+
|
| 2205 |
+
graphs = load_all_graphs(
|
| 2206 |
+
subset_percent=subset_percent,
|
| 2207 |
+
subset_seed=subset_seed,
|
| 2208 |
+
)
|
| 2209 |
if len(graphs) < 2:
|
| 2210 |
print(f"Need at least 2 labeled sheets, found {len(graphs)}. Label more data first.")
|
| 2211 |
return
|
vertex_launch.py
ADDED
|
@@ -0,0 +1,153 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Submit Vertex AI Custom Training jobs (shared by CLI and dashboard)."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import logging
|
| 6 |
+
from typing import Any
|
| 7 |
+
|
| 8 |
+
logger = logging.getLogger(__name__)
|
| 9 |
+
|
| 10 |
+
# Default accelerator comes from config / env (often L4 + g2 once quota exists).
|
| 11 |
+
_DEFAULT_ACCEL = "NVIDIA_L4"
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def resolve_machine_type_for_accelerator(machine_type: str, accelerator_type: str) -> str:
|
| 15 |
+
"""Return a Vertex-compatible machine type for the given GPU.
|
| 16 |
+
|
| 17 |
+
L4 must use **G2** (e.g. ``g2-standard-8``). T4/P100 typically use **N1**.
|
| 18 |
+
If the pair is inconsistent, we pick a safe default for the GPU family.
|
| 19 |
+
"""
|
| 20 |
+
mt = (machine_type or "").strip()
|
| 21 |
+
at = (accelerator_type or "").strip().upper()
|
| 22 |
+
if "L4" in at:
|
| 23 |
+
if not mt.startswith("g2-"):
|
| 24 |
+
return "g2-standard-8"
|
| 25 |
+
return mt
|
| 26 |
+
if "A100" in at:
|
| 27 |
+
if mt.startswith("a2-") or mt.startswith("a3-"):
|
| 28 |
+
return mt
|
| 29 |
+
# Single A100 — typical Vertex pairing (region/quota dependent).
|
| 30 |
+
return "a2-highgpu-1g"
|
| 31 |
+
if "T4" in at or "P100" in at:
|
| 32 |
+
if mt.startswith("g2-"):
|
| 33 |
+
return "n1-standard-8"
|
| 34 |
+
return mt or "n1-standard-8"
|
| 35 |
+
return mt or "n1-standard-8"
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def submit_vertex_training_jobs(
|
| 39 |
+
*,
|
| 40 |
+
project: str,
|
| 41 |
+
region: str,
|
| 42 |
+
image_uri: str,
|
| 43 |
+
gcs_data_uri: str,
|
| 44 |
+
architectures: list[str],
|
| 45 |
+
seeds: list[int],
|
| 46 |
+
folds: int,
|
| 47 |
+
epochs: int,
|
| 48 |
+
machine_type: str = "g2-standard-8",
|
| 49 |
+
accelerator_type: str = _DEFAULT_ACCEL,
|
| 50 |
+
accelerator_count: int = 1,
|
| 51 |
+
staging_bucket: str | None = None,
|
| 52 |
+
subset_percent: float = 100.0,
|
| 53 |
+
use_spot: bool = False,
|
| 54 |
+
arch_accelerator: dict[str, str] | None = None,
|
| 55 |
+
) -> dict[str, Any]:
|
| 56 |
+
"""One Custom Job per (seed, fold, arch). Returns counts and display names.
|
| 57 |
+
|
| 58 |
+
*arch_accelerator*: optional map ``architecture -> accelerator_type`` (e.g.
|
| 59 |
+
``{"adj_transformer": "NVIDIA_L4", "gat": "NVIDIA_TESLA_T4"}``). Omitted archs
|
| 60 |
+
use *accelerator_type*. Machine type is resolved per job via
|
| 61 |
+
``resolve_machine_type_for_accelerator``.
|
| 62 |
+
|
| 63 |
+
use_spot: if True, passes scheduling SPOT to ``job.run()`` (preemptible VM; often
|
| 64 |
+
P100 preemptible quota). If False, on-demand GPUs (e.g. ``CustomModelTrainingT4GPUsPerProjectPerRegion``).
|
| 65 |
+
"""
|
| 66 |
+
from google.cloud import aiplatform
|
| 67 |
+
from google.cloud.aiplatform.compat.types import custom_job as gca_custom_job
|
| 68 |
+
|
| 69 |
+
staging = staging_bucket
|
| 70 |
+
if not staging:
|
| 71 |
+
uri = gcs_data_uri.rstrip("/")
|
| 72 |
+
if not uri.startswith("gs://"):
|
| 73 |
+
raise ValueError("gcs_data_uri must start with gs://")
|
| 74 |
+
rest = uri[5:]
|
| 75 |
+
bucket = rest.split("/", 1)[0]
|
| 76 |
+
staging = f"gs://{bucket}"
|
| 77 |
+
|
| 78 |
+
aiplatform.init(project=project, location=region, staging_bucket=staging)
|
| 79 |
+
|
| 80 |
+
if gcs_data_uri.startswith("gs://"):
|
| 81 |
+
without = gcs_data_uri[5:]
|
| 82 |
+
parts = without.split("/", 1)
|
| 83 |
+
bucket = parts[0]
|
| 84 |
+
prefix = parts[1] if len(parts) > 1 else ""
|
| 85 |
+
fuse_root = f"/gcs/{bucket}" + (f"/{prefix}" if prefix else "")
|
| 86 |
+
else:
|
| 87 |
+
fuse_root = gcs_data_uri
|
| 88 |
+
|
| 89 |
+
# Use CustomJob.submit(), not job.run(sync=False). The latter schedules _run() in a
|
| 90 |
+
# background thread and returns immediately, so the HTTP handler could finish (and report
|
| 91 |
+
# N submitted jobs) before all create_custom_job RPCs complete — or surface errors only
|
| 92 |
+
# on the Future. submit() runs create_custom_job synchronously (does not wait for training).
|
| 93 |
+
submitted: list[str] = []
|
| 94 |
+
job_resource_names: list[str] = []
|
| 95 |
+
for seed in seeds:
|
| 96 |
+
for fold_idx in range(folds):
|
| 97 |
+
for arch in architectures:
|
| 98 |
+
accel = accelerator_type
|
| 99 |
+
if arch_accelerator and arch in arch_accelerator:
|
| 100 |
+
accel = (arch_accelerator[arch] or "").strip() or accelerator_type
|
| 101 |
+
mach = resolve_machine_type_for_accelerator(machine_type, accel)
|
| 102 |
+
display = f"train-{arch}-s{seed}-f{fold_idx}"
|
| 103 |
+
train_args = [
|
| 104 |
+
"train_compare.py",
|
| 105 |
+
"--architectures",
|
| 106 |
+
arch,
|
| 107 |
+
"--seeds",
|
| 108 |
+
str(seed),
|
| 109 |
+
"--folds",
|
| 110 |
+
str(folds),
|
| 111 |
+
"--epochs",
|
| 112 |
+
str(epochs),
|
| 113 |
+
"--fold-index",
|
| 114 |
+
str(fold_idx),
|
| 115 |
+
]
|
| 116 |
+
if subset_percent < 100.0:
|
| 117 |
+
train_args.extend(["--subset-percent", str(subset_percent)])
|
| 118 |
+
env = [{"name": "EXCEL_CHUNKER_DATA_ROOT", "value": fuse_root}]
|
| 119 |
+
worker_pool = {
|
| 120 |
+
"machine_spec": {
|
| 121 |
+
"machine_type": mach,
|
| 122 |
+
"accelerator_type": accel,
|
| 123 |
+
"accelerator_count": accelerator_count,
|
| 124 |
+
},
|
| 125 |
+
"replica_count": 1,
|
| 126 |
+
"container_spec": {
|
| 127 |
+
"image_uri": image_uri,
|
| 128 |
+
"command": ["python"],
|
| 129 |
+
"args": train_args,
|
| 130 |
+
"env": env,
|
| 131 |
+
},
|
| 132 |
+
}
|
| 133 |
+
job = aiplatform.CustomJob(
|
| 134 |
+
display_name=display[:128],
|
| 135 |
+
worker_pool_specs=[worker_pool],
|
| 136 |
+
)
|
| 137 |
+
submit_kw: dict[str, Any] = {}
|
| 138 |
+
if use_spot:
|
| 139 |
+
submit_kw["scheduling_strategy"] = gca_custom_job.Scheduling.Strategy.SPOT
|
| 140 |
+
job.submit(**submit_kw)
|
| 141 |
+
resource = getattr(job, "resource_name", None) or getattr(
|
| 142 |
+
job, "name", None
|
| 143 |
+
)
|
| 144 |
+
if resource:
|
| 145 |
+
job_resource_names.append(resource)
|
| 146 |
+
logger.info("Vertex CustomJob created: %s (%s)", display, resource)
|
| 147 |
+
submitted.append(display)
|
| 148 |
+
|
| 149 |
+
return {
|
| 150 |
+
"submitted_count": len(submitted),
|
| 151 |
+
"jobs": submitted,
|
| 152 |
+
"job_resource_names": job_resource_names,
|
| 153 |
+
}
|