File size: 19,393 Bytes
aade474
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
75f4985
aade474
75f4985
 
aade474
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
75f4985
 
 
 
 
 
aade474
 
 
 
75f4985
 
 
 
aade474
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
75f4985
 
 
 
 
 
 
 
 
 
 
 
 
 
 
aade474
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
75f4985
 
 
 
 
 
 
 
 
 
aade474
75f4985
aade474
 
75f4985
aade474
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
75f4985
 
 
 
 
aade474
 
 
 
 
 
75f4985
aade474
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
"""GCP Vertex AI Custom Job entrypoint — **evaluation** (not training).

Mirrors gcp_entrypoint.py but:
  1. Pulls dataset from HF Hub (same logic).
  2. Pulls a trained checkpoint from <HF_RUNS_REPO>/<RUN_ID>/stage2/{best|last}/
     into /workspace/run_pull/<RUN_ID>/stage2/<CKPT_PICK>/.
  3. Pulls <RUN_ID>/configs/{model,train}_config.yaml (if present) so the model
     is rebuilt with the EXACT architecture used at training time. Falls back
     to the repo defaults when the snapshot is absent.
  4. Patches the configs for the detected GPU (T4 → fp16+SDPA; L4/A100 →
     bf16+FA2 if flash-attn installed) and pins paths.
  5. Pins run_id.txt to RUN_ID so evaluate.py writes under {output_dir}/{RUN_ID}/.
  6. Execs `python -m evaluation.evaluate --checkpoint ... --run_id ...
     --task all --split test --output_dir /workspace/results`.

Required env vars:
  HF_TOKEN         — HuggingFace token (read code+data+runs, write runs)
  DATASET_NAME     — 'IU-Xray' | 'MIMIC-CXR' | 'MIMIC-CXR_resized'
  RUN_ID           — folder name on HF_RUNS_REPO to evaluate, e.g.
                     'MIMIC-CXR_resized_run_1'

Optional env vars (defaults shown):
  HF_USER             = hieu3636
  HF_RUNS_REPO        = hieu3636/cxr-vlm-runs
  CKPT_PICK           = best                # 'best' | 'last' (stage2 sub-folder)
  REPORT_MODE         =                     # blank → use the run's saved snapshot
  IMAGE_MODE          =                     # blank → use the run's saved snapshot
  TASK                = all                 # 'all', one task, or comma list ('findings,vqa')
  SPLIT               = test
  PNU_SOURCE          = oracle              # 'oracle' | 'predicted' (CheXpert classifier)
  CHEXPERT_CKPT_PATH  = chexpert_classifier/chexpert_mimic_resized.pt  # on HF_RUNS_REPO; predicted only
  BATCH_SIZE          =                     # blank → from GPU profile
  MAX_NEW_TOKENS      = 300
  UPLOAD_RESULTS_TO_HF= 1                   # set '0' to skip upload
  LLM_JUDGE           = 0                   # set '1' to enable VQA judge
  LLM_JUDGE_MODEL     = gpt-4o-mini
  LLM_JUDGE_BASE_URL  =                     # for OpenAI-compat endpoints
  LLM_JUDGE_MAX_SAMPLES=                    # cap samples for cost control
  WORK                = /workspace
"""

from __future__ import annotations

import os
import shutil
import subprocess
import sys
import tarfile
import zipfile
from pathlib import Path

# Tame HF/transformers chatter so logs are readable in Cloud Logging.
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
os.environ.setdefault("BITSANDBYTES_NOWELCOME", "1")
os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1")
os.environ.setdefault("TRANSFORMERS_VERBOSITY", "warning")
os.environ.setdefault("PYTHONUNBUFFERED", "1")
os.environ.setdefault("CUDA_VISIBLE_DEVICES", "0")
# tqdm flushes ≥ every 1s — keeps Cloud Logging looking alive during inference.
os.environ.setdefault("TQDM_MININTERVAL", "1.0")


def env(name: str, default: str | None = None, *, required: bool = False) -> str:
    val = os.environ.get(name, default)
    if required and not val:
        sys.exit(f"[gcp_eval_entrypoint] ERROR: required env var {name} not set")
    return val or ""


# ── 1) Resolve config from env ────────────────────────────────────────────────
HF_TOKEN              = env("HF_TOKEN", required=True)
DATASET_NAME          = env("DATASET_NAME", required=True)
RUN_ID                = env("RUN_ID", required=True)
HF_USER               = env("HF_USER", "hieu3636")
HF_RUNS_REPO          = env("HF_RUNS_REPO", "hieu3636/cxr-vlm-runs")
CKPT_PICK             = env("CKPT_PICK", "best")
REPORT_MODE_OVERRIDE  = env("REPORT_MODE", "")
IMAGE_MODE_OVERRIDE   = env("IMAGE_MODE", "")
TASK                  = env("TASK", "all")
SPLIT                 = env("SPLIT", "test")
BATCH_SIZE_OVERRIDE   = env("BATCH_SIZE", "")
MAX_NEW_TOKENS        = int(env("MAX_NEW_TOKENS", "300"))
UPLOAD_RESULTS_TO_HF  = env("UPLOAD_RESULTS_TO_HF", "1") not in ("0", "false", "False", "")
LLM_JUDGE             = env("LLM_JUDGE", "0") in ("1", "true", "True")
LLM_JUDGE_MODEL       = env("LLM_JUDGE_MODEL", "gpt-4o-mini")
LLM_JUDGE_BASE_URL    = env("LLM_JUDGE_BASE_URL", "")
LLM_JUDGE_MAX_SAMPLES = env("LLM_JUDGE_MAX_SAMPLES", "")
# PNU prompt-condition source — mirrors the notebook. 'oracle' = GT chex_*
# labels from the JSON; 'predicted' = run the Stage-0 CheXpert classifier on
# each image (needs CHEXPERT_CKPT_PATH on the runs repo).
PNU_SOURCE            = env("PNU_SOURCE", "oracle")
CHEXPERT_CKPT_PATH    = env("CHEXPERT_CKPT_PATH",
                            "chexpert_classifier/chexpert_mimic_resized.pt")
WORK                  = Path(env("WORK", "/workspace"))

assert DATASET_NAME in ("IU-Xray", "MIMIC-CXR", "MIMIC-CXR_resized"), DATASET_NAME
assert CKPT_PICK in ("best", "last"), CKPT_PICK
# TASK may be 'all', a single task, or a comma-separated subset ('findings,vqa').
for _t in TASK.split(","):
    assert _t.strip() in ("all", "findings", "impression", "report", "vqa"), TASK
assert PNU_SOURCE in ("oracle", "predicted"), PNU_SOURCE

PROJECT       = Path(__file__).resolve().parent.parent         # /workspace/code
DATA_SRC      = WORK / "data"
RUN_PULL_ROOT = WORK / "run_pull"
RESULTS_DIR   = WORK / "results"
CKPT_ROOT     = WORK / "ckpt_eval"
for d in (DATA_SRC, RUN_PULL_ROOT, RESULTS_DIR, CKPT_ROOT):
    d.mkdir(parents=True, exist_ok=True)

print(f"[gcp_eval] PROJECT      = {PROJECT}")
print(f"[gcp_eval] WORK         = {WORK}")
print(f"[gcp_eval] DATASET_NAME = {DATASET_NAME}")
print(f"[gcp_eval] RUN_ID       = {RUN_ID}  (ckpt: stage2/{CKPT_PICK})")
print(f"[gcp_eval] TASK         = {TASK}    SPLIT = {SPLIT}")
print(f"[gcp_eval] UPLOAD_TO_HF = {UPLOAD_RESULTS_TO_HF}")

# ── 2) Download dataset payload from HF Hub ───────────────────────────────────
from huggingface_hub import HfApi, hf_hub_download, snapshot_download  # noqa: E402

if DATASET_NAME == "MIMIC-CXR_resized":
    mr_dir = DATA_SRC / "MIMIC-CXR_resized"
    mr_dir.mkdir(parents=True, exist_ok=True)
    files_dir = mr_dir / "files"
    manifests_present = all(
        (mr_dir / f).is_file()
        for f in ("manifest_train.csv", "manifest_val.csv", "manifest_test.csv")
    )
    if manifests_present and files_dir.is_dir() and any(files_dir.glob("p*")):
        print(f"[gcp_eval] {mr_dir} already populated — skipping download.")
    else:
        api = HfApi(token=HF_TOKEN)
        all_files = api.list_repo_files(
            repo_id=f"{HF_USER}/cxr-vlm-data", repo_type="dataset"
        )
        mr_files = [f for f in all_files if f.startswith("MIMIC-CXR_resized/")]
        tar_files = sorted(f for f in mr_files if f.endswith(".tar"))
        print(f"[gcp_eval] {len(tar_files)} tar shards on HF")

        snapshot_download(
            repo_id=f"{HF_USER}/cxr-vlm-data",
            repo_type="dataset",
            allow_patterns=[
                "MIMIC-CXR_resized/*.csv",
                "MIMIC-CXR_resized/*.json",
                "MIMIC-CXR_resized/*.txt",
                "MIMIC-CXR_resized/vqa/**",
            ],
            token=HF_TOKEN,
            local_dir=str(DATA_SRC),
        )

        for i, tf in enumerate(tar_files, 1):
            print(f"[gcp_eval]  [{i}/{len(tar_files)}] {tf}", flush=True)
            tp = Path(hf_hub_download(
                repo_id=f"{HF_USER}/cxr-vlm-data",
                repo_type="dataset",
                filename=tf,
                token=HF_TOKEN,
                local_dir=str(DATA_SRC),
            ))
            with tarfile.open(tp) as t:
                t.extractall(mr_dir)
            tp.unlink(missing_ok=True)
        print(f"[gcp_eval] {mr_dir} ready.")

    DATA_ROOT_RESIZED = mr_dir

else:
    zip_name = f"{DATASET_NAME}.zip"
    marker = DATA_SRC / DATASET_NAME
    if not marker.exists():
        print(f"[gcp_eval] downloading {zip_name} ...")
        zpath = hf_hub_download(
            repo_id=f"{HF_USER}/cxr-vlm-data",
            filename=zip_name,
            repo_type="dataset",
            token=HF_TOKEN,
            local_dir=str(DATA_SRC),
        )
        with zipfile.ZipFile(zpath) as zf:
            zf.extractall(DATA_SRC)
        try:
            os.remove(zpath)
        except OSError:
            pass
    else:
        print(f"[gcp_eval] {marker} already present — skipping download.")

print(f"[gcp_eval] DATA_SRC contents: {sorted(os.listdir(DATA_SRC))}")

# ── 3) Pull checkpoint + saved config snapshot from runs repo ────────────────
print(f"[gcp_eval] pulling {RUN_ID}/{{configs,stage2/{CKPT_PICK}}} from {HF_RUNS_REPO} …")
snapshot_download(
    repo_id=HF_RUNS_REPO,
    repo_type="model",
    token=HF_TOKEN,
    allow_patterns=[
        f"{RUN_ID}/configs/**",
        f"{RUN_ID}/run_meta.json",
        f"{RUN_ID}/stage2/{CKPT_PICK}/**",
    ],
    local_dir=str(RUN_PULL_ROOT),
)
RUN_DIR_PULLED  = RUN_PULL_ROOT / RUN_ID
CKPT_DIR_PULLED = RUN_DIR_PULLED / "stage2" / CKPT_PICK
PROJ_PT         = CKPT_DIR_PULLED / "checkpoint_projection.pt"
LORA_DIR        = CKPT_DIR_PULLED / "checkpoint_lora"
CHEXPERT_PT     = CKPT_DIR_PULLED / "checkpoint_chexpert_classifier.pt"
SAVED_CFG_DIR   = RUN_DIR_PULLED / "configs"
SAVED_TRAIN_CFG = SAVED_CFG_DIR / "train_config.yaml"
SAVED_MODEL_CFG = SAVED_CFG_DIR / "model_config.yaml"

assert PROJ_PT.is_file(), f"projection weights not found: {PROJ_PT}"
assert (LORA_DIR / "adapter_config.json").is_file(), \
    f"LoRA adapter_config.json missing in {LORA_DIR}"
print(f"[gcp_eval]   projection : {PROJ_PT}  ({PROJ_PT.stat().st_size/1e6:.1f} MB)")
print(f"[gcp_eval]   lora       : {LORA_DIR}/")
print(f"[gcp_eval]   chexpert   : exists={CHEXPERT_PT.is_file()}")

# CheXpert classifier checkpoint (separate location on the runs repo, outside
# {RUN_ID}/). Only needed for PNU_SOURCE=predicted.
CHEXPERT_LOCAL = None
if PNU_SOURCE == "predicted":
    CHEXPERT_LOCAL = Path(hf_hub_download(
        repo_id=HF_RUNS_REPO, repo_type="model",
        filename=CHEXPERT_CKPT_PATH, token=HF_TOKEN,
        local_dir=str(RUN_PULL_ROOT),
    ))
    assert CHEXPERT_LOCAL.is_file(), \
        f"CheXpert checkpoint not pulled: {CHEXPERT_LOCAL} " \
        f"(expected {CHEXPERT_CKPT_PATH!r} on {HF_RUNS_REPO})."
    print(f"[gcp_eval]   chexpert clf (predicted PNU): {CHEXPERT_LOCAL} "
          f"({CHEXPERT_LOCAL.stat().st_size/1e6:.2f} MB)")

# ── 4) Build configs ─────────────────────────────────────────────────────────
import torch  # noqa: E402
from omegaconf import OmegaConf  # noqa: E402

repo_train_cfg_path = PROJECT / "configs" / "train_config.yaml"
repo_model_cfg_path = PROJECT / "configs" / "model_config.yaml"

if SAVED_TRAIN_CFG.is_file():
    train_cfg = OmegaConf.load(SAVED_TRAIN_CFG)
    print(f"[gcp_eval] train_cfg <- {SAVED_TRAIN_CFG}")
else:
    train_cfg = OmegaConf.load(repo_train_cfg_path)
    print(f"[gcp_eval] train_cfg <- repo default (no snapshot)")

if SAVED_MODEL_CFG.is_file():
    model_cfg = OmegaConf.load(SAVED_MODEL_CFG)
    print(f"[gcp_eval] model_cfg <- {SAVED_MODEL_CFG}")
else:
    model_cfg = OmegaConf.load(repo_model_cfg_path)
    print(f"[gcp_eval] model_cfg <- repo default (no snapshot)")

# Optional report/image mode overrides — leave blank to use saved values.
if REPORT_MODE_OVERRIDE:
    train_cfg.data.report_mode = REPORT_MODE_OVERRIDE
if IMAGE_MODE_OVERRIDE:
    train_cfg.data.image_mode  = IMAGE_MODE_OVERRIDE
print(f"[gcp_eval] report_mode = {train_cfg.data.report_mode}  image_mode = {train_cfg.data.image_mode}")

# Dataset paths (mirror gcp_entrypoint.py)
train_cfg.data.dataset_name           = DATASET_NAME
train_cfg.data.max_images_per_sample  = int(getattr(train_cfg.data, "max_images_per_sample", 2))

out_dir = PROJECT / "data" / "data_files"
out_dir.mkdir(parents=True, exist_ok=True)

if DATASET_NAME == "MIMIC-CXR_resized":
    train_cfg.data.mimic_cxr_resized.root          = str(DATA_ROOT_RESIZED)
    train_cfg.data.mimic_cxr_resized.manifest_dir  = None
    train_cfg.data.mimic_cxr_resized.vqa_dir       = None
    train_cfg.data.mimic_cxr_resized.reports_root  = None
    train_cfg.data.mimic_cxr_resized.auto_build    = True
    train_cfg.data.mimic_cxr_resized.instruct_json = str(
        out_dir / "mimic_cxr_resized_instruct.json")
elif DATASET_NAME == "MIMIC-CXR":
    def _find_mimic_root(root: Path) -> Path:
        for cand in [root / "MIMIC-CXR", root]:
            if (cand / "train").exists() and (cand / "valid").exists() and (cand / "test").exists():
                return cand
        for p in root.rglob("train"):
            if p.is_dir() and (p.parent / "valid").exists() and (p.parent / "test").exists():
                return p.parent
        raise FileNotFoundError(f"MIMIC-CXR train/valid/test not found under {root}")
    cxr_root = _find_mimic_root(DATA_SRC)
    train_cfg.data.mimic_cxr_root   = str(cxr_root)
    train_cfg.data.instruct_json    = str(out_dir / "mimic_cxr_instruct_unified.json")
    train_cfg.data.mimic_auto_build = True
    _cx = sorted(DATA_SRC.rglob("*chexpert*.csv")) or sorted(DATA_SRC.rglob("*chexbert*.csv"))
    train_cfg.data.mimic_chexpert_csv = str(_cx[0]) if _cx else None
    _vqa = list(DATA_SRC.rglob("vqa"))
    train_cfg.data.mimic_vqa_root = str(_vqa[0]) if _vqa else None
else:  # IU-Xray
    iu_root = DATA_SRC / "IU-Xray"
    train_cfg.data.iu_xray.images_dir    = str(iu_root / "images")
    train_cfg.data.iu_xray.labels_dir    = str(iu_root / "labels")
    train_cfg.data.iu_xray.instruct_json = str(out_dir / "iu_xray_instruct.json")
    train_cfg.data.iu_xray.auto_build    = True

train_cfg.data.train_split = "train"
train_cfg.data.val_split   = "validate"
train_cfg.data.test_split  = "test"
train_cfg.data.feature_cache_dir = None
train_cfg.training.output_root = str(CKPT_ROOT)

# ── GPU auto-profile (eval batch sizes — smaller than training because the KV
#    cache built during generation scales with batch × max_new_tokens). ──
assert torch.cuda.is_available(), "CUDA not available in container"
_props   = torch.cuda.get_device_properties(0)
_cap     = (_props.major, _props.minor)
_vram_gb = _props.total_memory / 1e9
_bf16_ok = torch.cuda.is_bf16_supported()
_fa2_ok  = _cap >= (8, 0)

_flash_attn_installed = False
if _fa2_ok:
    try:
        import flash_attn  # noqa: F401
        _flash_attn_installed = True
    except Exception:
        _flash_attn_installed = False

if _vram_gb >= 70:
    _label, _eval_bs, _nw = "A100/H100 80GB", 16, 16
elif _vram_gb >= 35:
    _label, _eval_bs, _nw = "A100 40GB", 8, 12
elif _vram_gb >= 22:
    _label, _eval_bs, _nw = "3090 / L4 / A10 (24GB)", 4, 8
elif _vram_gb >= 14:
    _label, _eval_bs, _nw = "T4 / V100 (15-16GB)", 1, 2
else:
    _label, _eval_bs, _nw = f"unknown ({_vram_gb:.0f}GB)", 1, 2

if BATCH_SIZE_OVERRIDE:
    _eval_bs = int(BATCH_SIZE_OVERRIDE)
    print(f"[gcp_eval] BATCH_SIZE override -> {_eval_bs}")

print(f"[gcp_eval] GPU: {_props.name}  {_vram_gb:.1f}GB  sm_{_cap[0]}{_cap[1]}  "
      f"bf16={_bf16_ok}  fa2={_fa2_ok}  fa2_wheel={_flash_attn_installed}")
print(f"[gcp_eval] -> profile {_label}  eval_batch={_eval_bs}")

train_cfg.training.per_device_train_batch_size   = _eval_bs
train_cfg.training.per_device_eval_batch_size    = _eval_bs
train_cfg.training.dataloader_num_workers        = _nw
train_cfg.training.fp16                          = not _bf16_ok
train_cfg.training.bf16                          = bool(_bf16_ok)
train_cfg.training.dataloader_pin_memory         = True
train_cfg.training.dataloader_persistent_workers = True

model_cfg.llm.attn_implementation       = (
    "flash_attention_2" if (_fa2_ok and _flash_attn_installed) else "sdpa"
)
model_cfg.llm.gradient_checkpointing    = False     # eval — no backward
model_cfg.llm.torch_dtype               = "bfloat16" if _bf16_ok else "float16"
model_cfg.llm.bnb_4bit_compute_dtype    = "bfloat16" if _bf16_ok else "float16"
model_cfg.llm.bnb_4bit_quant_type       = "nf4"
model_cfg.llm.bnb_4bit_use_double_quant = True
model_cfg.llm.load_in_8bit              = False
model_cfg.llm.load_in_4bit              = True

# CheXpert classifier.
#   PNU_SOURCE=predicted → enable + load the separately-pulled Stage-0 head.
#   PNU_SOURCE=oracle    → only enable for the legacy case where the head was
#                          saved inside the stage2 checkpoint; otherwise disabled
#                          (oracle PNU comes from the JSON's GT chex_* labels).
if PNU_SOURCE == "predicted":
    model_cfg.chexpert_classifier.enabled    = True
    model_cfg.chexpert_classifier.checkpoint = str(CHEXPERT_LOCAL)
    print(f"[gcp_eval] PNU=predicted -> CheXpert classifier ENABLED, ckpt={CHEXPERT_LOCAL}")
elif CHEXPERT_PT.is_file():
    model_cfg.chexpert_classifier.enabled = True
    print("[gcp_eval] CheXpert classifier checkpoint found in stage2 -> enabled")
else:
    model_cfg.chexpert_classifier.enabled = False
    print("[gcp_eval] PNU=oracle -> CheXpert classifier disabled (GT chex_* labels)")

# HF Hub tracker
train_cfg.wandb.enabled         = False
if UPLOAD_RESULTS_TO_HF:
    train_cfg.hf_hub.enabled    = True
    train_cfg.hf_hub.repo_id    = HF_RUNS_REPO
    train_cfg.hf_hub.token_env  = "HF_TOKEN"
    train_cfg.hf_hub.private    = True
else:
    train_cfg.hf_hub.enabled    = False
train_cfg.hf_hub.run_state_file = str(CKPT_ROOT / "run_id.txt")

# Pin run_id so resolve_run_id picks up exactly the one we're evaluating.
(CKPT_ROOT / "run_id.txt").write_text(RUN_ID)
print(f"[gcp_eval] pinned run_id = {RUN_ID}")

# Save patched configs back into the project so evaluate.py reads them.
OmegaConf.save(train_cfg, repo_train_cfg_path)
OmegaConf.save(model_cfg, repo_model_cfg_path)
print("[gcp_eval] configs patched.")

# ── 5) Launch evaluation ─────────────────────────────────────────────────────
cmd = [
    "python", "-u", "-m", "evaluation.evaluate",
    "--model_config",   str(repo_model_cfg_path),
    "--train_config",   str(repo_train_cfg_path),
    # Pass the DIRECTORY, not the .pt file: load_checkpoint derives
    # checkpoint_projection.pt + checkpoint_lora/ from the dir. Passing the .pt
    # makes it look for checkpoint_projection_projection.pt → silently skips
    # BOTH projection and LoRA (evaluating raw, untrained Vicuna).
    "--checkpoint",     str(CKPT_DIR_PULLED),
    "--run_id",         RUN_ID,
    "--task",           TASK,
    "--split",          SPLIT,
    "--batch_size",     str(_eval_bs),
    "--max_new_tokens", str(MAX_NEW_TOKENS),
    "--output_dir",     str(RESULTS_DIR),
    "--pnu_source",     PNU_SOURCE,
    "--device",         "cuda",
]
if not UPLOAD_RESULTS_TO_HF:
    cmd.append("--no_hf_upload")
if LLM_JUDGE:
    cmd += ["--llm_judge", "--llm_judge_model", LLM_JUDGE_MODEL]
    if LLM_JUDGE_BASE_URL:
        cmd += ["--llm_judge_base_url", LLM_JUDGE_BASE_URL]
    if LLM_JUDGE_MAX_SAMPLES:
        cmd += ["--llm_judge_max_samples", LLM_JUDGE_MAX_SAMPLES]

print(f"[gcp_eval] launching: {' '.join(cmd)}", flush=True)
os.chdir(PROJECT)
sys.exit(subprocess.call(cmd))