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())
|