2026-06-15 11:40:00 +08:00
|
|
|
"""Build action/proprio normalizers and resolve the effective norm key.
|
|
|
|
|
|
|
|
|
|
Public inference artifacts must carry their own normalization data. This module
|
2026-09-23 20:58:39 +08:00
|
|
|
uses checkpoint-local ``norm_stats.json`` for the 7D Euler layout, and uses
|
|
|
|
|
checkpoint-side normalizer state dicts for checkpoints whose train config contains
|
|
|
|
|
6D rotation fields. It finally supports an explicit
|
|
|
|
|
``customized_action_statistic_dof`` path.
|
2026-06-15 11:40:00 +08:00
|
|
|
It does not fall back to internal default action statistics.
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
import json
|
|
|
|
|
import logging
|
|
|
|
|
import os
|
|
|
|
|
|
|
|
|
|
import torch
|
|
|
|
|
|
|
|
|
|
from wall_x.data.backends.lerobot.utils import NormStats
|
|
|
|
|
from wall_x.model.core.action.normalizer import Normalizer, pad_normalizer_to_dim
|
|
|
|
|
from wall_x._vendor.harrix.utils.train_config import (
|
|
|
|
|
resolve_agent_pos_config,
|
|
|
|
|
resolve_dof_config,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _load_norm_stats(norm_stats_path: str, action_key: str) -> NormStats:
|
|
|
|
|
with open(norm_stats_path, "r") as f:
|
|
|
|
|
norm_stats = json.load(f)
|
|
|
|
|
q01 = torch.tensor(norm_stats["norm_stats"][action_key]["q01"])
|
|
|
|
|
q99 = torch.tensor(norm_stats["norm_stats"][action_key]["q99"])
|
|
|
|
|
return NormStats(min=q01, max=q99, delta=q99 - q01)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _load_custom_action_stats(train_config: dict) -> dict | None:
|
|
|
|
|
custom = train_config.get("customized_action_statistic_dof", None)
|
|
|
|
|
if not custom:
|
|
|
|
|
return None
|
|
|
|
|
with open(custom, "r") as f:
|
|
|
|
|
return json.load(f)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _normalizer_from_stats(action_stats: dict, train_config: dict, key: str) -> Normalizer:
|
|
|
|
|
return Normalizer(
|
|
|
|
|
action_stats,
|
|
|
|
|
train_config[key],
|
|
|
|
|
min_key=train_config.get("min_key", "min"),
|
|
|
|
|
delta_key=train_config.get("delta_key", "delta"),
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _missing_normalizer_error(checkpoint_path: str, train_config: dict) -> FileNotFoundError:
|
|
|
|
|
custom = train_config.get("customized_action_statistic_dof", None)
|
|
|
|
|
return FileNotFoundError(
|
|
|
|
|
"Public inference requires normalization data. Expected one of: "
|
|
|
|
|
f"{os.path.join(checkpoint_path, 'norm_stats.json')}; checkpoint-side "
|
|
|
|
|
"normalizer_action.pth and normalizer_propri.pth; or an explicit "
|
|
|
|
|
f"customized_action_statistic_dof path. Current customized_action_statistic_dof={custom!r}."
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
2026-09-23 20:58:39 +08:00
|
|
|
def _uses_rotation_6d(train_config: dict) -> bool:
|
|
|
|
|
"""Return whether the checkpoint was trained with 6D rotation fields."""
|
|
|
|
|
for layout_name in ("dof_config", "agent_pos_config", "ar_dof_config"):
|
|
|
|
|
layout = train_config.get(layout_name)
|
|
|
|
|
if layout is None and isinstance(train_config.get("task"), dict):
|
|
|
|
|
layout = train_config["task"].get(layout_name)
|
|
|
|
|
if isinstance(layout, dict) and any(
|
|
|
|
|
"rotation_6d" in str(key).lower() for key in layout
|
|
|
|
|
):
|
|
|
|
|
return True
|
|
|
|
|
return False
|
|
|
|
|
|
|
|
|
|
|
2026-06-15 11:40:00 +08:00
|
|
|
def build_normalizers(
|
|
|
|
|
checkpoint_path: str,
|
|
|
|
|
train_config: dict,
|
|
|
|
|
norm_key: str,
|
|
|
|
|
) -> tuple[Normalizer, Normalizer, str]:
|
|
|
|
|
"""Return action/proprio normalizers and the resolved norm key."""
|
|
|
|
|
norm_stats_path = os.path.join(checkpoint_path, "norm_stats.json")
|
2026-09-23 20:58:39 +08:00
|
|
|
action_pth = os.path.join(checkpoint_path, "normalizer_action.pth")
|
|
|
|
|
propri_pth = os.path.join(checkpoint_path, "normalizer_propri.pth")
|
|
|
|
|
prefer_checkpoint_pth = (
|
|
|
|
|
os.path.exists(action_pth)
|
|
|
|
|
and os.path.exists(propri_pth)
|
|
|
|
|
and _uses_rotation_6d(train_config)
|
|
|
|
|
)
|
|
|
|
|
if os.path.exists(norm_stats_path) and not prefer_checkpoint_pth:
|
2026-06-15 11:40:00 +08:00
|
|
|
propri_stats = _load_norm_stats(norm_stats_path, "observation.state")
|
|
|
|
|
action_stats = _load_norm_stats(norm_stats_path, "action")
|
|
|
|
|
normalizer_propri = Normalizer.from_lerobot_norm_stats(propri_stats, norm_key)
|
|
|
|
|
normalizer_action = Normalizer.from_lerobot_norm_stats(action_stats, norm_key)
|
|
|
|
|
else:
|
2026-09-23 20:58:39 +08:00
|
|
|
if prefer_checkpoint_pth:
|
|
|
|
|
logger.info(
|
|
|
|
|
"Using checkpoint normalizer .pth files because train config "
|
|
|
|
|
"contains rotation_6D; ignoring 7D norm_stats.json"
|
|
|
|
|
)
|
2026-06-15 11:40:00 +08:00
|
|
|
custom_stats = _load_custom_action_stats(train_config)
|
|
|
|
|
if custom_stats is None and (not os.path.exists(action_pth) or not os.path.exists(propri_pth)):
|
|
|
|
|
raise _missing_normalizer_error(checkpoint_path, train_config)
|
|
|
|
|
|
|
|
|
|
if os.path.exists(action_pth):
|
|
|
|
|
normalizer_action = Normalizer.from_ckpt(action_pth)
|
|
|
|
|
else:
|
|
|
|
|
normalizer_action = _normalizer_from_stats(custom_stats, train_config, "dof_config")
|
|
|
|
|
|
|
|
|
|
if os.path.exists(propri_pth):
|
|
|
|
|
normalizer_propri = Normalizer.from_ckpt(propri_pth)
|
|
|
|
|
else:
|
|
|
|
|
normalizer_propri = _normalizer_from_stats(custom_stats, train_config, "agent_pos_config")
|
|
|
|
|
|
|
|
|
|
action_dim = sum(resolve_dof_config(train_config).values())
|
|
|
|
|
propri_dim = sum(resolve_agent_pos_config(train_config).values())
|
|
|
|
|
pad_normalizer_to_dim(normalizer_action, action_dim, "action")
|
|
|
|
|
pad_normalizer_to_dim(normalizer_propri, propri_dim, "propri")
|
|
|
|
|
|
|
|
|
|
resolved = _resolve_norm_key(norm_key, normalizer_action, normalizer_propri)
|
|
|
|
|
return normalizer_action, normalizer_propri, resolved
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _resolve_norm_key(
|
|
|
|
|
norm_key: str,
|
|
|
|
|
normalizer_action: Normalizer,
|
|
|
|
|
normalizer_propri: Normalizer,
|
|
|
|
|
) -> str:
|
|
|
|
|
"""Resolve a requested norm key against normalizer keys."""
|
|
|
|
|
available = sorted(
|
|
|
|
|
set(normalizer_action.min.keys()) & set(normalizer_propri.min.keys())
|
|
|
|
|
)
|
|
|
|
|
if norm_key in available:
|
|
|
|
|
return norm_key
|
|
|
|
|
if not available:
|
|
|
|
|
return norm_key
|
|
|
|
|
|
|
|
|
|
prefix_matches = [k for k in available if k.startswith(f"{norm_key}_")]
|
|
|
|
|
if len(prefix_matches) == 1:
|
|
|
|
|
logger.warning(
|
|
|
|
|
"norm_key=%r not found; using prefix fallback %r",
|
|
|
|
|
norm_key,
|
|
|
|
|
prefix_matches[0],
|
|
|
|
|
)
|
|
|
|
|
return prefix_matches[0]
|
|
|
|
|
if len(available) == 1:
|
|
|
|
|
logger.warning(
|
|
|
|
|
"norm_key=%r not found; using the only available key %r",
|
|
|
|
|
norm_key,
|
|
|
|
|
available[0],
|
|
|
|
|
)
|
|
|
|
|
return available[0]
|
|
|
|
|
|
|
|
|
|
logger.warning(
|
|
|
|
|
"norm_key=%r not found; available=%s; returning the requested key unchanged",
|
|
|
|
|
norm_key,
|
|
|
|
|
available,
|
|
|
|
|
)
|
|
|
|
|
return norm_key
|