Instructions to use zuoyerumeng/xvla-m2w-multitask-1gpu with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- LeRobot
How to use zuoyerumeng/xvla-m2w-multitask-1gpu with LeRobot:
- Notebooks
- Google Colab
- Kaggle
File size: 6,805 Bytes
74a3a91 | 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 | #!/usr/bin/env python3
"""Validate pinned X-VLA weights and LeRobot training YAML files without loading the model."""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any
import draccus
import numpy as np
import yaml
from safetensors import safe_open
from lerobot.configs import parser
from lerobot.configs.train import TrainPipelineConfig
import lerobot.configs.policies as policy_configs
from lerobot.policies.xvla.configuration_xvla import XVLAConfig # noqa: F401
MODEL_SHA256 = "f05bc0fab1c9523d7f5d6b41a651313641ca2227a88822249845da8e20e036c9"
def parse_training_config(path: Path) -> TrainPipelineConfig:
# Config parsing normally replaces unavailable CUDA with CPU. This check is
# intentionally hardware-independent, so preserve the requested CUDA device.
original_device_check = policy_configs.is_torch_device_available
policy_configs.is_torch_device_available = lambda _device: True
try:
clean_path = parser.extract_path_fields_from_config(
str(path), TrainPipelineConfig.__get_path_fields__()
)
cfg = draccus.parse(TrainPipelineConfig, config_path=clean_path, args=[])
cfg.validate()
return cfg
finally:
policy_configs.is_torch_device_available = original_device_check
def validate_config(
path: Path,
expected_steps: int,
expected_batch_size: int,
expected_save_checkpoint: bool = True,
) -> dict[str, Any]:
raw_config = yaml.safe_load(path.read_text(encoding="utf-8"))
raw_policy = raw_config.get("policy") or {}
for feature_field in ("input_features", "output_features"):
if feature_field not in raw_policy or raw_policy[feature_field] is not None:
raise ValueError(
f"{path}: policy.{feature_field} must be an unquoted YAML null so "
"the production CLI infers dataset features"
)
cfg = parse_training_config(path)
policy = cfg.policy
if policy is None or policy.type != "xvla":
raise ValueError(f"{path}: did not resolve to an XVLA policy")
expected = {
"device": "cuda",
"dtype": "bfloat16",
"action_mode": "auto",
"max_action_dim": 20,
"max_state_dim": 20,
"num_image_views": 3,
"freeze_vision_encoder": False,
"freeze_language_encoder": False,
"train_policy_transformer": True,
"train_soft_prompts": True,
}
for key, value in expected.items():
if getattr(policy, key) != value:
raise ValueError(f"{path}: policy.{key} did not resolve to {value!r}")
if policy.input_features is not None or policy.output_features is not None:
raise ValueError(f"{path}: policy feature dictionaries must be inferred from the dataset")
if cfg.steps != expected_steps or cfg.batch_size != expected_batch_size:
raise ValueError(f"{path}: unexpected steps or batch size")
if cfg.save_checkpoint is not expected_save_checkpoint:
raise ValueError(f"{path}: unexpected checkpoint-saving setting")
if cfg.tolerance_s != 0.001 or cfg.dataset.eval_split != 0.1:
raise ValueError(f"{path}: unexpected timestamp tolerance or eval split")
if cfg.dataset.use_imagenet_stats:
raise ValueError(
f"{path}: dataset.use_imagenet_stats must be false because X-VLA normalizes "
"images in its policy processor and merged camera stats are intentionally absent"
)
return {
"path": str(path),
"steps": cfg.steps,
"batch_size_per_process": cfg.batch_size,
"save_checkpoint": cfg.save_checkpoint,
"dataset_root": str(cfg.dataset.root),
"policy_path": str(policy.pretrained_path),
"dtype": policy.dtype,
"action_mode": policy.action_mode,
"num_image_views": policy.num_image_views,
"use_imagenet_stats": cfg.dataset.use_imagenet_stats,
"eval_split": cfg.dataset.eval_split,
"tolerance_s": cfg.tolerance_s,
}
def main() -> None:
project_root = Path(__file__).resolve().parents[1]
model_root = project_root / "models" / "xvla-base"
model_manifest = json.loads(
(model_root / "download_manifest.json").read_text(encoding="utf-8")
)
if model_manifest["model_sha256"] != MODEL_SHA256:
raise ValueError("Pinned X-VLA model manifest SHA256 mismatch")
tensor_count = 0
parameter_count = 0
dtype_counts: dict[str, int] = {}
with safe_open(model_root / "model.safetensors", framework="pt", device="cpu") as stream:
for key in stream.keys():
tensor = stream.get_slice(key)
count = int(np.prod(tensor.get_shape()))
dtype = str(tensor.get_dtype())
tensor_count += 1
parameter_count += count
dtype_counts[dtype] = dtype_counts.get(dtype, 0) + count
merged_validation = json.loads(
(project_root / "data" / "merged" / "validation_report.json").read_text(encoding="utf-8")
)
if merged_validation["status"] != "passed" or merged_validation["frames"] != 120_469:
raise ValueError("Merged training dataset validation report is not usable")
configs = [
validate_config(project_root / "configs" / "xvla_smoke.yaml", 20, 1),
validate_config(
project_root / "configs" / "xvla_pilot.yaml",
100,
16,
expected_save_checkpoint=False,
),
validate_config(project_root / "configs" / "xvla_full.yaml", 20_000, 4),
validate_config(project_root / "configs" / "xvla_full_1gpu.yaml", 20_000, 16),
]
report = {
"status": "passed",
"lerobot_version": "0.6.0",
"model": {
"root": str(model_root),
"revision": model_manifest["revision"],
"sha256": model_manifest["model_sha256"],
"tensor_count": tensor_count,
"parameter_count": parameter_count,
"dtype_parameter_counts": dtype_counts,
},
"dataset": {
"root": merged_validation["root"],
"episodes": merged_validation["episodes"],
"frames": merged_validation["frames"],
"tasks": len(merged_validation["tasks"]),
},
"configs": configs,
}
report_path = project_root / "configs" / "train_setup_validation.json"
report_path.write_text(
json.dumps(report, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
print(
f"[passed] X-VLA setup: {parameter_count:,} parameters, "
f"{merged_validation['frames']:,} data frames, {len(configs)} configs",
flush=True,
)
print(f"Validation report: {report_path}", flush=True)
if __name__ == "__main__":
main()
|