refactor pi052 to reuse pi05

This commit is contained in:
Pepijn
2026-07-15 19:26:55 +02:00
parent 55d9ff740e
commit 5b8e6ffe8e
5 changed files with 223 additions and 563 deletions
@@ -37,7 +37,7 @@ def test_shifted_ce_none_retains_distinct_per_sample_losses():
def test_checkpoint_resolution_forwards_explicit_hub_options(monkeypatch, tmp_path):
import lerobot.policies.pi052.modeling_pi052 as modeling_pi052
import lerobot.policies.pi05.modeling_pi05 as modeling_pi05
checkpoint = tmp_path / "model.safetensors"
checkpoint.touch()
@@ -47,8 +47,8 @@ def test_checkpoint_resolution_forwards_explicit_hub_options(monkeypatch, tmp_pa
calls.append((model_id, filename, kwargs))
return None if filename.endswith("index.json") else str(checkpoint)
monkeypatch.setattr(modeling_pi052, "cached_file", fake_cached_file)
files = modeling_pi052._resolve_weight_files(
monkeypatch.setattr(modeling_pi05, "cached_file", fake_cached_file)
files = modeling_pi05._resolve_weight_files(
"org/model",
force_download=True,
resume_download=True,
@@ -71,10 +71,10 @@ def test_checkpoint_resolution_forwards_explicit_hub_options(monkeypatch, tmp_pa
def test_checkpoint_resolution_rejects_local_directory_without_weights(tmp_path):
import lerobot.policies.pi052.modeling_pi052 as modeling_pi052
import lerobot.policies.pi05.modeling_pi05 as modeling_pi05
with pytest.raises(FileNotFoundError, match="model.safetensors"):
modeling_pi052._resolve_weight_files(
modeling_pi05._resolve_weight_files(
tmp_path,
force_download=False,
resume_download=None,
@@ -41,14 +41,12 @@ def _checkpoint_model():
tower = _MockVisionTower()
language_model = SimpleNamespace(gradient_checkpointing=False)
expert_model = SimpleNamespace(gradient_checkpointing=False)
model = SimpleNamespace(
gradient_checkpointing_enabled=False,
paligemma_with_expert=SimpleNamespace(
paligemma=SimpleNamespace(
model=SimpleNamespace(language_model=language_model, vision_tower=tower)
),
gemma_expert=SimpleNamespace(model=expert_model),
),
model = PI05Pytorch.__new__(PI05Pytorch)
nn.Module.__init__(model)
model.gradient_checkpointing_enabled = False
model.paligemma_with_expert = SimpleNamespace(
paligemma=SimpleNamespace(model=SimpleNamespace(language_model=language_model, vision_tower=tower)),
gemma_expert=SimpleNamespace(model=expert_model),
)
return model, tower, language_model, expert_model