File size: 4,753 Bytes
adc02fa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python
from __future__ import annotations

import argparse
import json
import os
import sys
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
    sys.path.insert(0, str(ROOT))

from dovla_cil.eval.external_vla_baseline import (  # noqa: E402
    ExternalVLABaselineSpec,
    assess_external_vla_baseline,
    build_external_vla_plan,
    run_external_vla_entrypoint,
    write_external_vla_plan,
)


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description=(
            "Prepare or run an isolated public VLA baseline adapter. This command never downloads "
            "weights or imports heavy VLA packages unless a user-provided adapter entrypoint is "
            "run."
        )
    )
    parser.add_argument("--model-family", default="smolvla", choices=["smolvla", "openvla"])
    parser.add_argument("--checkpoint", default=None, help="Local checkpoint/model directory.")
    parser.add_argument("--dataset", default=None, help="DoVLA-CIL dataset directory to evaluate.")
    parser.add_argument("--out", required=True, help="Output directory for plan and metrics.")
    parser.add_argument("--revision", default=None, help="Pinned public checkpoint revision.")
    parser.add_argument("--repo-id", default=None, help="Public Hugging Face repo id.")
    parser.add_argument("--package-name", default=None, help="External package to check/import.")
    parser.add_argument(
        "--python", default="python", help="Python executable to use in the generated env plan."
    )
    parser.add_argument(
        "--adapter-entrypoint",
        default=None,
        help=(
            "External adapter formatted as module:function. The function receives "
            "(spec_dict, plan_dict) and returns JSON-serializable metrics."
        ),
    )
    parser.add_argument(
        "--adapter-config",
        type=Path,
        default=None,
        help="Secret-free JSON object passed to the adapter as spec metadata.",
    )
    parser.add_argument(
        "--dry-run",
        action="store_true",
        help="Write a plan and exit without requiring the external baseline to be ready.",
    )
    parser.add_argument(
        "--require-ready",
        action="store_true",
        help="Exit nonzero unless package, checkpoint, dataset, and adapter entrypoint are ready.",
    )
    return parser.parse_args()


def main() -> int:
    args = parse_args()
    adapter_metadata = _load_adapter_config(args.adapter_config)
    spec = ExternalVLABaselineSpec(
        model_family=args.model_family,
        checkpoint_path=args.checkpoint,
        dataset_dir=args.dataset,
        out_dir=args.out,
        revision=args.revision,
        repo_id=args.repo_id,
        package_name=args.package_name,
        python=args.python,
        adapter_entrypoint=args.adapter_entrypoint,
        metadata=adapter_metadata,
    )
    out_dir = Path(args.out)
    out_dir.mkdir(parents=True, exist_ok=True)
    plan_path = write_external_vla_plan(spec, out_dir)
    plan = build_external_vla_plan(spec)
    status = assess_external_vla_baseline(spec)

    print(f"external VLA plan: {plan_path}")
    print(json.dumps(status.to_dict(), indent=2, sort_keys=True))
    if args.dry_run:
        return 0
    if args.require_ready and not status.ready:
        print(
            "External VLA baseline is not ready. Use the generated plan to create an isolated "
            "environment, download the public checkpoint, and provide an adapter entrypoint.",
            file=sys.stderr,
        )
        return 2
    if not args.adapter_entrypoint:
        print(
            "No adapter entrypoint was provided; wrote a reproducible plan but did not run "
            "metrics.",
            file=sys.stderr,
        )
        return 2

    metrics = run_external_vla_entrypoint(args.adapter_entrypoint, spec, plan)
    metrics_path = out_dir / "external_vla_metrics.json"
    metrics_path.write_text(json.dumps(metrics, indent=2, sort_keys=True), encoding="utf-8")
    print(f"external VLA metrics: {metrics_path}")
    return 0


def _load_adapter_config(path: Path | None) -> dict[str, object]:
    if path is None:
        return {}
    payload = json.loads(os.path.expandvars(path.read_text(encoding="utf-8")))
    if not isinstance(payload, dict):
        raise ValueError("adapter config must be a JSON object")
    forbidden = {"api_key", "apikey", "token", "secret", "password"}
    unsafe = [key for key in payload if key.lower() in forbidden]
    if unsafe:
        raise ValueError(f"adapter config must not contain secrets: {', '.join(sorted(unsafe))}")
    return payload


if __name__ == "__main__":
    raise SystemExit(main())