mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
fix(libero): resolve LIBERO-plus perturbation init_states path ourselves
Delegating to `task_suite.get_task_init_states(i)` works for path resolution
but LIBERO-plus's method calls `torch.load(path)` without `weights_only=False`,
which fails on PyTorch 2.6+ because the pickled init_states contains numpy
objects not in the default allowlist:
_pickle.UnpicklingError: Weights only load failed.
WeightsUnpickler error: Unsupported global:
GLOBAL numpy.core.multiarray._reconstruct was not an allowed global.
Mirror LIBERO-plus's suffix-stripping logic (`_table_N`, `_tb_N`, `_view_`,
`_language_`, `_light_`, `_add_`, `_level`) in our own helper so we can pass
`weights_only=False` ourselves. Vanilla LIBERO task names don't contain any
of these patterns except for `_table_` when followed by the word `center`
(e.g. `pick_up_the_black_bowl_from_table_center_...`), and the regex
requires `_table_\\d+` so semantic uses are preserved.
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
+40
-11
@@ -16,6 +16,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
import re
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||||
from functools import partial
|
from functools import partial
|
||||||
@@ -69,19 +70,47 @@ def _select_task_ids(total_tasks: int, task_ids: Iterable[int] | None) -> list[i
|
|||||||
return ids
|
return ids
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_libero_init_states_path(task: Any) -> tuple[Path, bool]:
|
||||||
|
"""Resolve the on-disk init_states path for a LIBERO task.
|
||||||
|
|
||||||
|
LIBERO-plus perturbation variants encode the perturbation in the filename
|
||||||
|
(_table_N, _tb_N, _view_, _language_, _light_, _add_, _level) but only the
|
||||||
|
base `.pruned_init` exists on disk. Mirror LIBERO-plus's own suffix-stripping
|
||||||
|
(libero/libero/benchmark/__init__.py) so we can load with `weights_only=False`
|
||||||
|
ourselves instead of delegating to the suite (whose `torch.load` call omits
|
||||||
|
that flag and fails on PyTorch 2.6+ numpy pickles).
|
||||||
|
|
||||||
|
Returns (path, needs_reshape) — `_add_`/`_level` variants store a flat
|
||||||
|
init_state that must be reshaped to (1, -1).
|
||||||
|
"""
|
||||||
|
filename = task.init_states_file
|
||||||
|
problem_folder = task.problem_folder
|
||||||
|
ext = filename.rsplit(".", 1)[-1]
|
||||||
|
root = Path(get_libero_path("init_states"))
|
||||||
|
|
||||||
|
needs_reshape = "_add_" in filename or "_level" in filename
|
||||||
|
if needs_reshape:
|
||||||
|
return root / "libero_newobj" / problem_folder / filename, True
|
||||||
|
if "_language_" in filename:
|
||||||
|
filename = filename.split("_language_")[0] + "." + ext
|
||||||
|
elif "_view_" in filename:
|
||||||
|
filename = filename.split("_view_")[0] + "." + ext
|
||||||
|
else:
|
||||||
|
if "_table_" in filename:
|
||||||
|
filename = re.sub(r"_table_\d+", "", filename)
|
||||||
|
if "_tb_" in filename:
|
||||||
|
filename = re.sub(r"_tb_\d+", "", filename)
|
||||||
|
if "_light_" in filename:
|
||||||
|
filename = filename.split("_light_")[0] + "." + ext
|
||||||
|
return root / problem_folder / filename, False
|
||||||
|
|
||||||
|
|
||||||
def get_task_init_states(task_suite: Any, i: int) -> np.ndarray:
|
def get_task_init_states(task_suite: Any, i: int) -> np.ndarray:
|
||||||
# LIBERO-plus's Benchmark exposes a suffix-aware loader that maps perturbation
|
task = task_suite.tasks[i]
|
||||||
# variants (_table_N, _tb_N, _view_, _language_, _light_, _add_, _level) back to
|
init_states_path, needs_reshape = _resolve_libero_init_states_path(task)
|
||||||
# the base init_states file. Vanilla LIBERO has no such method — fall through.
|
|
||||||
suite_loader = getattr(task_suite, "get_task_init_states", None)
|
|
||||||
if callable(suite_loader):
|
|
||||||
return suite_loader(i)
|
|
||||||
init_states_path = (
|
|
||||||
Path(get_libero_path("init_states"))
|
|
||||||
/ task_suite.tasks[i].problem_folder
|
|
||||||
/ task_suite.tasks[i].init_states_file
|
|
||||||
)
|
|
||||||
init_states = torch.load(init_states_path, weights_only=False) # nosec B614
|
init_states = torch.load(init_states_path, weights_only=False) # nosec B614
|
||||||
|
if needs_reshape:
|
||||||
|
init_states = init_states.reshape(1, -1)
|
||||||
return init_states
|
return init_states
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user