ZibinDong's picture
Upload pretrained ActionCodec2 artifact
fee0e43 verified
Raw History Blame Contribute Delete
2.86 kB
"""Shared validation primitives for ActionCodec2 public boundaries."""
from __future__ import annotations
from collections.abc import Iterable, Sequence
from typing import TypeVar
import numpy as np
ErrorT = TypeVar("ErrorT", bound=Exception)
def active_dimensions(
active_dims: Sequence[int] | None,
dimension: int,
*,
subject: str = "active_dims",
) -> tuple[int, ...]:
"""Normalize and validate a non-empty increasing subset of dimensions."""
dims = (
tuple(range(dimension))
if active_dims is None
else tuple(int(item) for item in active_dims)
)
if not dims:
raise ValueError(f"{subject} must not be empty")
if any(left >= right for left, right in zip(dims, dims[1:])):
raise ValueError(f"{subject} must be strictly increasing")
if dims[0] < 0 or dims[-1] >= dimension:
raise ValueError(f"{subject} contains a dimension outside the vocabulary")
return dims
def optional_finite_array(
value: np.ndarray | None,
*,
shape: tuple[int, ...],
name: str = "current_state",
) -> np.ndarray | None:
"""Convert an optional float32 array and enforce its exact shape/finiteness."""
if value is None:
return None
result = np.asarray(value, dtype=np.float32)
if result.shape != shape:
rendered = ", ".join(str(item) for item in shape)
raise ValueError(f"{name} must have shape [{rendered}]")
if not np.all(np.isfinite(result)):
raise ValueError(f"{name} contains a non-finite value")
return result
def broadcast_vector(
value: object,
width: int,
*,
dtype: np.dtype,
name: str,
) -> np.ndarray:
"""Broadcast a scalar or width-sized value to one numerical vector."""
array = np.asarray(value, dtype=dtype)
try:
return np.broadcast_to(array, (width,)).copy()
except ValueError as error:
raise ValueError(f"{name} must be scalar or contain {width} values") from error
def validate_dimension_partition(
groups: Iterable[tuple[str, Sequence[int]]],
dimension: int,
*,
label: str,
error_type: type[ErrorT] = ValueError,
) -> None:
"""Validate that named dimension groups form one exact non-overlapping partition."""
owners: dict[int, str] = {}
for name, dims in groups:
for index in dims:
if index < 0 or index >= dimension:
raise error_type(
f"{label} {index} used by {name!r} is outside input_dim"
)
if index in owners:
raise error_type(
f"{label} {index} belongs to both {owners[index]!r} and {name!r}"
)
owners[index] = name
missing = sorted(set(range(dimension)) - set(owners))
if missing:
raise error_type(f"{label}(s) {missing} are not assigned")