Files

161 lines
6.0 KiB
Python
Raw Permalink Normal View History

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
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}."
)
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")
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:
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