zofiasmolenasana commited on
Commit
93e4108
·
unverified ·
1 Parent(s): 75e64b2

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 CHANGED
@@ -1,8 +1,13 @@
1
  # GPU training image for Vertex AI / local CUDA (PyTorch + PyTorch Geometric).
2
- # Build: docker build -f Dockerfile.training -t excel-chunker-train .
 
 
 
 
 
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-cudnn9-runtime
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
- """Sync latest label files from Drive to local data/labeled/."""
2
- import drive_client, metadata_client, config
3
 
4
- rows = metadata_client.get_all_rows()
5
- labelled = [r for r in rows if r["status"] == "labelled" and r.get("labels_file_id") and r["labels_file_id"] != "local-only"]
6
- print(f"Labelled sheets on Drive: {len(labelled)}")
7
 
8
- downloaded = 0
9
- failed = 0
10
- for r in labelled:
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 = f"{payload.spreadsheet_id}_{payload.sheet_name}".replace("/", "_").replace(" ", "_")
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 = f"{payload.spreadsheet_id}_{payload.sheet_name}".replace("/", "_").replace(" ", "_")
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 = f"{sid}_{sn}".replace("/", "_").replace(" ", "_") + ".json"
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 = f"{spreadsheet_id}_{sheet_name}".replace("/", "_").replace(" ", "_")
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
- _GCP_SUITE_TTL_SEC = 5.0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- uri = (os.environ.get("EXCEL_CHUNKER_GCS_EXPERIMENTS_URI") or "").strip()
 
 
 
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
- return _GCP_SUITE_CACHE["data"], _GCP_SUITE_CACHE["err"]
 
 
 
 
 
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
- return data, err
 
 
 
 
 
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 = (os.environ.get("EXCEL_CHUNKER_GCS_EXPERIMENTS_URI") or "").strip()
1482
- gcp_suite, gcp_err = _get_gcp_suite_cached()
1483
- return {
 
 
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
- result = run_experiments(
 
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
- _train_status["result"] = {"status": "completed", "architectures": list(result.keys())}
1578
- _train_status["log"].append("Training completed successfully!")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1579
 
1580
  elif req.mode == "production":
1581
- _train_status["log"].append(f"Training production model for {req.epochs} epochs...")
 
 
 
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
- train_model(num_epochs=req.epochs, skip_predict=True)
 
 
 
 
 
 
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
- """Generate text embeddings for all labeled files that don't have them yet."""
1629
- count = 0
1630
- skipped = 0
1631
- errors = []
1632
  try:
1633
- from embed_text import embed_file
1634
- for fp in sorted(config.LABELED_DIR.glob("*.json")):
1635
- embed_path = fp.with_name(fp.stem + "_embeddings.npz")
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 {"embedded": count, "skipped": skipped, "errors": errors}
 
 
 
 
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 = sorted(labeled_dir.glob("*.json"))
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
- ) -> dict[str, list[dict] | None]:
141
- """One ``predict_sheet`` per held-out sheet (under GNN_CELL_LIMIT). Used for chunking + structure metrics."""
 
 
 
 
 
 
 
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
- return cache
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- if any(k in val for k in ("recall@1", "mrr", "judge_score", "judge_binary_acc")):
 
 
 
 
 
 
 
 
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(key, resume, existing_agg, log):
 
 
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(sheet_data, model, log)
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- structure_metrics_doc[method_key_base] = entry
 
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(key, resume, existing_agg, log):
 
 
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(key, resume, existing_agg, log):
 
 
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
- Spot VMs: enable in the Cloud Console for the job, or use ``gcloud ai custom-jobs create`` with
20
- ``--scheduling-strategy=SPOT``; the Python SDK flag varies by version.
 
 
21
  """
22
 
23
  from __future__ import annotations
24
 
25
  import argparse
26
  import sys
 
 
 
 
 
 
 
27
 
28
 
29
  def main() -> None:
30
  try:
31
- from google.cloud import aiplatform
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="n1-standard-8",
52
- help="e.g. n1-standard-8",
 
 
 
 
 
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
- aiplatform.init(project=args.project, location=args.region, staging_bucket=staging)
76
-
77
- if args.gcs_data_uri.startswith("gs://"):
78
- without = args.gcs_data_uri[5:]
79
- parts = without.split("/", 1)
80
- bucket = parts[0]
81
- prefix = parts[1] if len(parts) > 1 else ""
82
- fuse_root = f"/gcs/{bucket}" + (f"/{prefix}" if prefix else "")
83
- else:
84
- fuse_root = args.gcs_data_uri
85
-
86
- jobs = 0
87
- for seed in seeds:
88
- for fold_idx in range(args.folds):
89
- for arch in archs:
90
- display = f"train-{arch}-s{seed}-f{fold_idx}"
91
- train_args = [
92
- "train_compare.py",
93
- "--architectures",
94
- arch,
95
- "--seeds",
96
- str(seed),
97
- "--folds",
98
- str(args.folds),
99
- "--epochs",
100
- str(args.epochs),
101
- "--fold-index",
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&amp;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="7">Loading…</td></tr>';
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="7" class="empty">No joined data. Run RAG eval with graph methods and ensure rag_eval_structure_metrics.json exists.</td></tr>';
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="7">Error: ' + esc(String(e)) + '</td></tr>';
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&amp;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; this app reads
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
- Generate E5 text embeddings for all labeled files. Already-embedded files are skipped.
 
 
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 ? 'Refreshed every ~5s (server cache)' : '');
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
- if (pollTimer) { clearInterval(pollTimer); pollTimer = null; }
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 &lt; 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 &amp; 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
- dst = config.MODELS_DIR / "best_model.pt"
 
 
 
 
 
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
- return dict(all_results)
 
 
 
 
 
 
 
 
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(num_epochs: int = 100, lr: float = 1e-3, skip_predict: bool = False):
 
 
 
 
 
 
 
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
- graphs = load_all_graphs()
 
 
 
 
 
 
 
 
 
 
 
 
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
+ }