# Copyright 2026 The HuggingFace Inc. team. All rights reserved. from __future__ import annotations import json import os from pathlib import Path import pytest import torch from torch import nn from lerobot.configs.policies import PreTrainedConfig from lerobot.configs.types import FeatureType, PolicyFeature from lerobot.policies.factory import get_policy_class, make_policy_config, make_pre_post_processors from lerobot.policies.g05.configuration_g05 import G05_EMBODIMENT_MAPPINGS, G05Config from lerobot.policies.g05.convert_g05_checkpoint import ( convert_checkpoint, convert_state_dict, save_converted_state_dict, ) from lerobot.policies.g05.modeling_g05 import G05Policy from lerobot.processor import PolicyProcessorPipeline from lerobot.utils.constants import ACTION, OBS_STATE, POLICY_PREPROCESSOR_DEFAULT_NAME class TinyG05Backend(nn.Module): def __init__(self): super().__init__() self.proj = nn.Linear(20, 20) self.last_samples = None def predict_action(self, batch): self.last_samples = batch["samples"] state = batch[OBS_STATE] if state.ndim == 2: state = state.unsqueeze(1) step = self.proj(state[:, -1]) return { ACTION: step.unsqueeze(1).expand(-1, 4, -1), "ar_action": (step + 1).unsqueeze(1).expand(-1, 4, -1), "cot_text": ["Subtask: move carefully"] * step.shape[0], } def forward(self, batch): prediction = self.proj(batch[OBS_STATE][:, -1]) target = batch[ACTION][:, 0] loss = torch.nn.functional.mse_loss(prediction, target) return loss, {"fm_loss": loss.detach()} def _features(): return { OBS_STATE: PolicyFeature(type=FeatureType.STATE, shape=(7,)), "observation.images.image": PolicyFeature(type=FeatureType.VISUAL, shape=(3, 8, 8)), "observation.images.wrist_image": PolicyFeature(type=FeatureType.VISUAL, shape=(3, 8, 8)), } def _config(**kwargs): normalization_mode = kwargs.pop("normalization_mode", "identity") return G05Config( checkpoint_profile="custom", normalization_mode=normalization_mode, input_features=_features(), output_features={ACTION: PolicyFeature(type=FeatureType.ACTION, shape=(7,))}, chunk_size=4, device="cpu", **kwargs, ) def _policy_batch(task: str = " Pick café cup\nverbatim "): return { OBS_STATE: torch.zeros(1, 1, 20), ACTION: torch.zeros(1, 4, 20), "observation.images.image": torch.zeros(1, 3, 8, 8), "observation.images.wrist_image": torch.zeros(1, 3, 8, 8), "task": [task], "proprio_dim_is_pad": torch.zeros(20, dtype=torch.bool), } def test_factory_wiring_is_lazy(): assert make_policy_config("g05", checkpoint_profile="custom").type == "g05" assert get_policy_class("g05") is G05Policy def test_system2_fm_only_builder_uses_exact_cot_template_without_action_tokens(): config = _config( action_head="flow", runtime_system="system2", predict_cot=True, discrete_action=False, continuous_action=True, return_continuous_action=True, processor_metadata={ "samples_builder": { "_target_": ("g05.data_processor.processor.samples_builder.SubtaskCoTBuilderFMOnly") } }, ) assert "\n|Action: " in config.prompt_template assert " 0 and torch.isfinite(grad_norm) optimizer.step() assert metrics is not None and metrics["fm_loss"] >= 0 policy.save_pretrained(tmp_path) reloaded = G05Policy.from_pretrained( tmp_path, backend=TinyG05Backend(), local_files_only=True, strict=True ) expected = policy.predict_action_chunk(_policy_batch("save")) actual = reloaded.predict_action_chunk(_policy_batch("save")) torch.testing.assert_close(actual, expected) def test_save_pretrained_copies_required_gated_sidecars_portably(tmp_path: Path): source = tmp_path / "converted" processor = source / "hf_processor" processor.mkdir(parents=True) (processor / "tokenizer.json").write_text("{}") tokenizer = source / "action_tokenizer.pt" torch.save({"codec": "ActionCodec"}, tokenizer) for name in ("LICENSE-G0.5", "NOTICE", "conversion_report.json"): (source / name).write_text("{}") config = _config( author_model_config={ "hf_processor_path": str(processor), "AT_CONFIG": {"ckpt_dir": str(tokenizer)}, } ) output = tmp_path / "saved" G05Policy(config, backend=TinyG05Backend()).save_pretrained(output) assert (output / "hf_processor" / "tokenizer.json").is_file() assert (output / "action_tokenizer.pt").is_file() assert (output / "LICENSE-G0.5").is_file() loaded_config = PreTrainedConfig.from_pretrained(output) assert isinstance(loaded_config, G05Config) assert loaded_config.author_model_config["hf_processor_path"] == "hf_processor" assert loaded_config.author_model_config["AT_CONFIG"]["ckpt_dir"] == "action_tokenizer.pt" def test_tiny_fixed_batch_overfit_reduces_loss(): policy = G05Policy(_config(), backend=TinyG05Backend()) optimizer = torch.optim.AdamW(policy.get_optim_params()["params"], lr=5e-2) batch = _policy_batch("overfit") initial = policy(batch)[0].item() for _ in range(20): optimizer.zero_grad() loss, _ = policy(batch) loss.backward() optimizer.step() final = policy(batch)[0].item() assert final < initial * 0.25 def test_conversion_reports_mapping_duplicates_shapes_and_required_prefixes(): source = { "model.embed_tokens.weight": torch.zeros(2, 3), "model.vision_tower.block.weight": torch.ones(2, 2), "model.action_expert.block.weight": torch.ones(2, 2), } converted, report = convert_state_dict(source) assert "backend.model.vlm.input_proj.weight" in converted assert report.missing == [] expected = { "backend.model.vlm.input_proj.weight": torch.zeros(3, 3), "backend.model.vision_tower.block.weight": torch.zeros(2, 2), "backend.model.action_expert.block.weight": torch.ones(2, 2), } _, strict_report = convert_state_dict(source, expected) assert strict_report.shape_mismatched["backend.model.vlm.input_proj.weight"]["source"] == [2, 3] def test_conversion_records_and_deduplicates_tied_weight_aliases(tmp_path: Path): tied = torch.zeros(2, 3) aliases = save_converted_state_dict( { "backend.model.vlm.input_proj.weight": tied, "backend.model.vlm.output_proj.weight": tied, }, tmp_path / "model.safetensors", ) assert aliases == {"backend.model.vlm.output_proj.weight": "backend.model.vlm.input_proj.weight"} def test_libero_conversion_packages_model_processors_and_provenance(tmp_path: Path): source = tmp_path / "author" output = tmp_path / "lerobot" (source / ".hydra").mkdir(parents=True) (source / "hf_processor").mkdir() (source / "hf_processor" / "tokenizer.json").write_text("{}") (source / ".hydra" / "config.yaml").write_text( """ model: model_arch: num_input_images: 2 horizon_steps: 32 predict_cot: false discrete_action: false continuous_action: true processor: use_stepwise_action_norm: true norm_default_mode: q01/q99 camera_size_config: exterior: [256, 256] wrist_right: [256, 256] data: action_size: 32 processors: libero: shape_meta: state: - {key: right_ee_pose, shape: 6} - {key: right_gripper, shape: 1} action: - {key: right_ee_pose, shape: 6} - {key: right_gripper, shape: 1} images: - {key: image, camera_type: exterior, lerobot_key: observation.images.image, shape: [3, 224, 224]} - {key: wrist_image, camera_type: wrist_right, lerobot_key: observation.images.wrist_image, shape: [3, 224, 224]} tokenizer: vq_config: {block_wise_autoregressive: false} """ ) action_stats = { "right_ee_pose": { "stepwise_q01": torch.zeros(32, 6).tolist(), "stepwise_q99": torch.ones(32, 6).tolist(), }, "right_gripper": { "stepwise_q01": torch.zeros(32, 1).tolist(), "stepwise_q99": torch.ones(32, 1).tolist(), }, } state_stats = { "right_ee_pose": {"global_q01": [0.0] * 6, "global_q99": [1.0] * 6}, "right_gripper": {"global_q01": [0.0], "global_q99": [1.0]}, } (source / "dataset_stats.json").write_text( json.dumps({"libero": {"state": state_stats, "action": action_stats}}) ) torch.save( { "model.vlm.block.weight": torch.zeros(2, 2), "model.vision_tower.block.weight": torch.zeros(2, 2), "model.action_expert.block.weight": torch.zeros(2, 2), }, source / "model.pt", ) torch.save({"tokenizer_meta": {"codec": "ActionCodec"}}, source / "action_tokenizer.pt") license_file = source / "LICENSE-G0.5" license_file.write_text("test fixture license") report = convert_checkpoint(source, output, "g05-libero", license_file=license_file) assert len(report.mapped) == 3 assert (output / "model.safetensors").is_file() assert (output / "policy_preprocessor.json").is_file() assert (output / "policy_postprocessor.json").is_file() assert (output / "conversion_report.json").is_file() config = PreTrainedConfig.from_pretrained(output) assert isinstance(config, G05Config) assert config.source_checkpoint_revision assert config.prompt_template.startswith("") assert config.camera_sizes == { "observation.images.image": (256, 256), "observation.images.wrist_image": (256, 256), } @pytest.mark.skipif( not os.environ.get("LEROBOT_G05_CHECKPOINT"), reason="requires an accepted gated OpenGalaxea/G05 checkpoint and author CUDA environment", ) def test_gated_checkpoint_loads_strictly(): checkpoint = Path(os.environ["LEROBOT_G05_CHECKPOINT"]) policy = G05Policy.from_pretrained(checkpoint, local_files_only=True, strict=True) assert policy.config.source_checkpoint_revision