fix pi052 runtime and training safety

This commit is contained in:
Pepijn
2026-07-15 18:17:23 +02:00
parent 3cec067795
commit 0fe31bfae1
16 changed files with 861 additions and 643 deletions
+65 -1
View File
@@ -19,7 +19,71 @@ import torch
pytest.importorskip("transformers")
from lerobot.policies.pi052.modeling_pi052 import _lin_ce_flat
from lerobot.policies.pi052.modeling_pi052 import _lin_ce_flat, _shifted_lin_ce
def test_shifted_ce_none_retains_distinct_per_sample_losses():
hidden = torch.tensor(
[
[[8.0, 0.0], [0.0, 8.0], [0.0, 0.0]],
[[0.0, 8.0], [8.0, 0.0], [0.0, 0.0]],
]
)
labels = torch.tensor([[0, 0, 1], [0, 0, 1]])
losses = _shifted_lin_ce(hidden, torch.eye(2), labels, reduction="none")
assert losses.shape == (2,)
assert losses[0] < losses[1]
def test_checkpoint_resolution_forwards_explicit_hub_options(monkeypatch, tmp_path):
import lerobot.policies.pi052.modeling_pi052 as modeling_pi052
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_pi052, "cached_file", fake_cached_file)
files = modeling_pi052._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.pi052.modeling_pi052 as modeling_pi052
with pytest.raises(FileNotFoundError, match="model.safetensors"):
modeling_pi052._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])