#!/usr/bin/env python # Copyright 2024 The HuggingFace Inc. team. All rights reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. import logging from collections import deque from typing import Any import numpy as np import torch from torch import nn from src.config.policies import PreTrainedConfig from src.config.types import FeatureType, PolicyFeature from src.utils.constants import ACTION, OBS_STR def populate_queues( queues: dict[str, deque], batch: dict[str, torch.Tensor], exclude_keys: list[str] | None = None ): if exclude_keys is None: exclude_keys = [] for key in batch: # Ignore keys not in the queues already (leaving the responsibility to the caller to make sure the # queues have the keys they want). if key not in queues or key in exclude_keys: continue if len(queues[key]) != queues[key].maxlen: # initialize by copying the first observation several times until the queue is full while len(queues[key]) != queues[key].maxlen: queues[key].append(batch[key]) else: # add latest observation to the queue queues[key].append(batch[key]) return queues def get_device_from_parameters(module: nn.Module) -> torch.device: """Get a module's device by checking one of its parameters. Note: assumes that all parameters have the same device """ return next(iter(module.parameters())).device def get_dtype_from_parameters(module: nn.Module) -> torch.dtype: """Get a module's parameter dtype by checking one of its parameters. Note: assumes that all parameters have the same dtype. """ return next(iter(module.parameters())).dtype def get_output_shape(module: nn.Module, input_shape: tuple) -> tuple: """ Calculates the output shape of a PyTorch module given an input shape. Args: module (nn.Module): a PyTorch module input_shape (tuple): A tuple representing the input shape, e.g., (batch_size, channels, height, width) Returns: tuple: The output shape of the module. """ dummy_input = torch.zeros(size=input_shape) with torch.inference_mode(): output = module(dummy_input) return tuple(output.shape) def log_model_loading_keys(missing_keys: list[str], unexpected_keys: list[str]) -> None: """Log missing and unexpected keys when loading a model. Args: missing_keys (list[str]): Keys that were expected but not found. unexpected_keys (list[str]): Keys that were found but not expected. """ if missing_keys: logging.warning(f"Missing key(s) when loading model: {missing_keys}") if unexpected_keys: logging.warning(f"Unexpected key(s) when loading model: {unexpected_keys}") def raise_feature_mismatch_error( provided_features: set[str], expected_features: set[str], ) -> None: """ Raises a standardized ValueError for feature mismatches between dataset/environment and policy config. """ missing = expected_features - provided_features extra = provided_features - expected_features # TODO (jadechoghari): provide a dynamic rename map suggestion to the user. raise ValueError( f"Feature mismatch between dataset/environment and policy config.\n" f"- Missing features: {sorted(missing) if missing else 'None'}\n" f"- Extra features: {sorted(extra) if extra else 'None'}\n\n" f"Please ensure your dataset and policy use consistent feature names.\n" f"If your dataset uses different observation keys (e.g., cameras named differently), " f"use the `--rename_map` argument, for example:\n" f' --rename_map=\'{{"observation.images.left": "observation.images.camera1", ' f'"observation.images.top": "observation.images.camera2"}}\'' ) def validate_visual_features_consistency( cfg: PreTrainedConfig, features: dict[str, PolicyFeature], ) -> None: """ Validates visual feature consistency between a policy config and provided dataset/environment features. Args: cfg (PreTrainedConfig): The model or policy configuration containing input_features and type. features (Dict[str, PolicyFeature]): A mapping of feature names to PolicyFeature objects. """ expected_visuals = {k for k, v in cfg.input_features.items() if v.type == FeatureType.VISUAL} provided_visuals = {k for k, v in features.items() if v.type == FeatureType.VISUAL} if not provided_visuals.issubset(expected_visuals): raise_feature_mismatch_error(provided_visuals, expected_visuals)