refactor(pi052): load native base checkpoint

This commit is contained in:
Pepijn
2026-07-28 12:32:43 +02:00
parent 6ac10f2a13
commit f2c8867df1
6 changed files with 98 additions and 133 deletions
@@ -36,56 +36,6 @@ def test_shifted_ce_none_retains_distinct_per_sample_losses():
assert losses[0] < losses[1]
def test_checkpoint_resolution_forwards_explicit_hub_options(monkeypatch, tmp_path):
import lerobot.policies.pi05.modeling_pi05 as modeling_pi05
checkpoint = tmp_path / "model.safetensors"
checkpoint.touch()
calls = []
def fake_cached_file(model_id, filename, **kwargs):
calls.append((model_id, filename, kwargs))
return None if filename.endswith("index.json") else str(checkpoint)
monkeypatch.setattr(modeling_pi05, "cached_file", fake_cached_file)
files = modeling_pi05._resolve_weight_files(
"org/model",
force_download=True,
resume_download=True,
proxies={"https": "proxy"},
token="secret",
cache_dir=tmp_path / "cache",
local_files_only=True,
revision="commit",
)
assert files == [checkpoint]
for _model_id, _filename, kwargs in calls:
assert kwargs["revision"] == "commit"
assert kwargs["cache_dir"] == tmp_path / "cache"
assert kwargs["force_download"] is True
assert kwargs["resume_download"] is True
assert kwargs["proxies"] == {"https": "proxy"}
assert kwargs["token"] == "secret"
assert kwargs["local_files_only"] is True
def test_checkpoint_resolution_rejects_local_directory_without_weights(tmp_path):
import lerobot.policies.pi05.modeling_pi05 as modeling_pi05
with pytest.raises(FileNotFoundError, match="model.safetensors"):
modeling_pi05._resolve_weight_files(
tmp_path,
force_download=False,
resume_download=None,
proxies=None,
token=None,
cache_dir=None,
local_files_only=False,
revision=None,
)
@pytest.mark.parametrize("z_loss_weight", [0.0, 1e-4])
@pytest.mark.parametrize("rows,valid_rows", [(24, 9), (48, 25)])
def test_bucketed_ce_matches_dense_loss_and_gradients(z_loss_weight, rows, valid_rows):
+67
View File
@@ -46,6 +46,15 @@ class _CheckpointPolicy(PI05Policy):
return [], []
class _NativeCheckpointPolicy(PI05Policy):
use_native_pretrained_loader = True
def __init__(self, config, **kwargs):
nn.Module.__init__(self)
self.config = config
self.weight = nn.Parameter(torch.zeros(1))
def test_from_pretrained_loads_existing_single_file_checkpoint(tmp_path):
save_file({"weight": torch.tensor([1.0])}, tmp_path / "model.safetensors")
@@ -55,6 +64,64 @@ def test_from_pretrained_loads_existing_single_file_checkpoint(tmp_path):
torch.testing.assert_close(policy.loaded_state_dict["model.weight"], torch.tensor([1.0]))
def test_pi05_checkpoint_loader_forwards_hub_options(monkeypatch, tmp_path):
import lerobot.policies.pi05.modeling_pi05 as modeling_pi05
checkpoint = tmp_path / "model.safetensors"
save_file({"weight": torch.tensor([1.0])}, checkpoint)
calls = []
def fake_cached_file(model_id, filename, **kwargs):
calls.append((model_id, filename, kwargs))
return str(checkpoint)
monkeypatch.setattr(modeling_pi05, "cached_file", fake_cached_file)
_CheckpointPolicy.from_pretrained(
"org/model",
config=SimpleNamespace(),
force_download=True,
resume_download=True,
proxies={"https": "proxy"},
token="secret",
cache_dir=tmp_path / "cache",
local_files_only=True,
revision="commit",
)
assert len(calls) == 1
model_id, filename, kwargs = calls[0]
assert model_id == "org/model"
assert filename == "model.safetensors"
assert kwargs["revision"] == "commit"
assert kwargs["cache_dir"] == tmp_path / "cache"
assert kwargs["force_download"] is True
assert kwargs["resume_download"] is True
assert kwargs["proxies"] == {"https": "proxy"}
assert kwargs["token"] == "secret"
assert kwargs["local_files_only"] is True
def test_pi05_checkpoint_loader_rejects_missing_weights(tmp_path):
with pytest.raises(FileNotFoundError, match="model.safetensors"):
_CheckpointPolicy.from_pretrained(tmp_path, config=SimpleNamespace())
def test_native_checkpoint_uses_standard_lerobot_loader(tmp_path):
save_file({"weight": torch.tensor([2.0])}, tmp_path / "model.safetensors")
policy = _NativeCheckpointPolicy.from_pretrained(
tmp_path, config=SimpleNamespace(device="cpu"), strict=True
)
torch.testing.assert_close(policy.weight, torch.tensor([2.0]))
def test_pi052_uses_native_checkpoint_loader():
from lerobot.policies.pi052.modeling_pi052 import PI052Policy
assert PI052Policy.use_native_pretrained_loader
@require_cuda
@require_hf_token
def test_policy_instantiation():