# Copyright 2026 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.

from __future__ import annotations

import logging
from typing import Any

import torch
import torch.nn.functional as F  # noqa: N812

from lerobot.configs import PipelineFeatureType, PolicyFeature
from lerobot.configs.types import NormalizationMode
from lerobot.policies.vla_jepa.configuration_vla_jepa import VLAJEPAConfig
from lerobot.processor import (
    AbsoluteActionsProcessorStep,
    EnvTransition,
    ObservationProcessorStep,
    PolicyAction,
    PolicyProcessorPipeline,
    ProcessorStep,
    ProcessorStepRegistry,
    RelativeActionsProcessorStep,
    TransitionKey,
    UnnormalizerProcessorStep,
    make_default_policy_processor_steps,
    make_policy_processor_pipelines,
)
from lerobot.utils.constants import ACTION


@ProcessorStepRegistry.register(name="vla_jepa_image_prep")
class ImagePrepProcessorStep(ObservationProcessorStep):
    """Prepares image observations for the VLA-JEPA model: float cast, 1->3 channel expand, resize.

    Surfaces in the serialized pipeline the prep the model used to do internally. The model keeps the
    same operations as idempotent guards, so older checkpoints saved without this step still get
    prepped and newer ones no-op there.

    Mirrors `Qwen3VLInterface.to_pixel_values` plus the `F.interpolate(mode="area")` resize in
    `VLAJEPAPolicy._prepare_model_inputs`/`predict_action`, and does not clamp (neither does the model
    path), so values stay bit-identical. Handles [C,H,W], [B,C,H,W]/[T,C,H,W] and [B,T,C,H,W].
    """

    def __init__(self, resize_to: tuple[int, int] | None = None, expand_channels: bool = True):
        self.resize_to = tuple(resize_to) if resize_to is not None else None
        self.expand_channels = expand_channels

    def observation(self, observation: dict) -> dict:
        new_observation = dict(observation)
        for key in observation:
            if "image" not in key:
                continue
            image = observation[key].float()
            if self.expand_channels and image.shape[-3] == 1:
                repeats = [1] * image.ndim
                repeats[-3] = 3
                image = image.repeat(*repeats)
            if self.resize_to is not None and tuple(image.shape[-2:]) != self.resize_to:
                device = image.device
                # NOTE: no "area" kernel on mps; resize on cpu then move back.
                if device.type == "mps":
                    image = image.cpu()
                lead = image.shape[:-3]
                c, h, w = image.shape[-3:]
                flat = image.reshape(-1, c, h, w)
                flat = F.interpolate(flat, size=self.resize_to, mode="area")
                image = flat.reshape(*lead, c, *self.resize_to).to(device)
            new_observation[key] = image
        return new_observation

    def get_config(self) -> dict[str, Any]:
        return {
            "resize_to": list(self.resize_to) if self.resize_to is not None else None,
            "expand_channels": self.expand_channels,
        }

    def transform_features(self, features):
        for key in features[PipelineFeatureType.OBSERVATION]:
            if "image" not in key:
                continue
            feat = features[PipelineFeatureType.OBSERVATION][key]
            # Match `to_pixel_values`: only a single channel is expanded to 3.
            nb_channel = 3 if (self.expand_channels and feat.shape[0] == 1) else feat.shape[0]
            spatial = self.resize_to if self.resize_to is not None else tuple(feat.shape[1:])
            features[PipelineFeatureType.OBSERVATION][key] = PolicyFeature(
                type=feat.type, shape=(nb_channel, *spatial)
            )
        return features


def _tensor_action(transition: EnvTransition) -> torch.Tensor | None:
    """The transition's action, checked to be a tensor (None when absent)."""
    action = transition.get(TransitionKey.ACTION)
    if action is not None and not isinstance(action, torch.Tensor):
        raise ValueError(f"Expected a tensor action, got {type(action).__name__}.")
    return action


@ProcessorStepRegistry.register(name="vla_jepa_clip_actions")
class ClipActionsProcessorStep(ProcessorStep):
    """Clips action tensor to [-1, 1] before unnormalization."""

    def __call__(self, transition: EnvTransition) -> EnvTransition:
        action = _tensor_action(transition)
        if action is not None:
            transition = transition.copy()
            transition[TransitionKey.ACTION] = action.clamp(-1.0, 1.0)
        return transition

    def transform_features(self, features):
        return features


@ProcessorStepRegistry.register(name="vla_jepa_pre_snap_gripper")
class PreSnapGripperProcessorStep(ProcessorStep):
    """Snaps a gripper dimension to {0, 1} BEFORE unnormalization.

    Mirrors the original starVLA LIBERO eval:
      normalized[:, gripper_dim] = np.where(normalized[:, gripper_dim] < threshold, 0, 1)
    This ensures the unnormalizer receives an exact binary value, which is
    required when the model was trained with gripper in identity (mask=False)
    space where 0=open and 1=close.
    """

    def __init__(self, gripper_dim: int = 6, threshold: float = 0.5):
        self.gripper_dim = gripper_dim
        self.threshold = threshold

    def __call__(self, transition: EnvTransition) -> EnvTransition:
        action = _tensor_action(transition)
        if action is not None and action.shape[-1] > self.gripper_dim:
            transition = transition.copy()
            a = action.clone()
            a[..., self.gripper_dim] = (a[..., self.gripper_dim] >= self.threshold).float()
            transition[TransitionKey.ACTION] = a
        return transition

    def get_config(self) -> dict[str, Any]:
        # Without this the base class serializes `{}` and a reloaded pipeline silently reverts
        # to the class defaults, discarding whatever the training config set.
        return {"gripper_dim": self.gripper_dim, "threshold": self.threshold}

    def transform_features(self, features):
        return features


@ProcessorStepRegistry.register(name="vla_jepa_binarize_gripper")
class BinarizeGripperProcessorStep(ProcessorStep):
    """Binarizes a gripper dimension after unnormalization.

    Maps continuous value to {-1, 1}: > threshold → -1, <= threshold → 1 (matches starVLA convention).
    Only applied when action has more dimensions than gripper_dim.

    WARNING: runs *below* the unnormalizer, so `threshold` is compared against the gripper's
    **physical** value while its 0.5 default comes from the model's [0, 1]/±1 convention. For a
    gripper in degrees, mm or [0, 100] every value exceeds 0.5 and the output collapses to -1;
    `make_vla_jepa_pre_post_processors` warns when the dataset stats say that is the case.
    """

    def __init__(self, gripper_dim: int = 6, threshold: float = 0.5):
        self.gripper_dim = gripper_dim
        self.threshold = threshold

    def __call__(self, transition: EnvTransition) -> EnvTransition:
        action = _tensor_action(transition)
        if action is not None and action.shape[-1] > self.gripper_dim:
            transition = transition.copy()
            a = action.clone()
            a[..., self.gripper_dim] = 1.0 - 2.0 * (a[..., self.gripper_dim] > self.threshold).float()
            transition[TransitionKey.ACTION] = a
        return transition

    def get_config(self) -> dict[str, Any]:
        # See PreSnapGripperProcessorStep.get_config: `{}` would reload as the class defaults.
        return {"gripper_dim": self.gripper_dim, "threshold": self.threshold}

    def transform_features(self, features):
        return features


def _warn_if_gripper_steps_are_misconfigured(
    config: VLAJEPAConfig,
    gripper_dim: int,
    dataset_stats: dict[str, dict[str, torch.Tensor]] | None,
) -> None:
    """Warn when the gripper post-steps would pin the gripper to a constant.

    `BinarizeGripperProcessorStep` thresholds the *unnormalized* gripper at `gripper_threshold`. When
    the dataset's physical range sits well above it, every value lands on the same side and the
    gripper never moves — detectable from the stats already here, so say so.
    """
    if not (config.pre_snap_gripper_action or config.binarize_gripper_action):
        return
    action_stats = (dataset_stats or {}).get(ACTION)
    if not action_stats or "min" not in action_stats or "max" not in action_stats:
        return
    try:
        low = float(action_stats["min"][gripper_dim])
        high = float(action_stats["max"][gripper_dim])
    except (IndexError, TypeError, ValueError):
        return
    threshold = config.gripper_threshold
    # `pre_snap` writes {0, 1} in normalized space, which unnormalizes to the midpoint and the
    # max. Both landing on the same side of the threshold means a constant output.
    midpoint = (low + high) / 2.0
    if (midpoint > threshold) == (high > threshold):
        name = (
            config.action_feature_names[gripper_dim]
            if config.action_feature_names and gripper_dim < len(config.action_feature_names)
            else f"dim {gripper_dim}"
        )
        logging.warning(
            f"vla_jepa gripper post-processing looks misconfigured: action {name} has a physical "
            f"range of [{low:.3g}, {high:.3g}], and `gripper_threshold={threshold}` is compared "
            f"against that unnormalized value. Both {midpoint:.3g} and {high:.3g} fall on the same "
            f"side of it, so the commanded gripper will be constant. Set `gripper_threshold` in "
            f"the gripper's own units, or set `pre_snap_gripper_action=false` and "
            f"`binarize_gripper_action=false` (the defaults) unless you are running LIBERO."
        )


def make_vla_jepa_pre_post_processors(
    config: VLAJEPAConfig,
    dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
) -> tuple[
    PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
    PolicyProcessorPipeline[PolicyAction, PolicyAction],
]:
    if config.input_features is None or config.output_features is None:
        raise ValueError(
            "`input_features` and `output_features` must be resolved before building the processors."
        )
    features = {**config.input_features, **config.output_features}
    steps = make_default_policy_processor_steps(config, dataset_stats)

    # Shared relative-action step (OpenPI order: raw -> relative -> normalize -> model ->
    # unnormalize -> absolute). The SAME instance is passed to AbsoluteActionsProcessorStep
    # below so its cached raw state (set during preprocessing) flows to postprocessing.
    relative_step = RelativeActionsProcessorStep(
        enabled=config.use_relative_actions,
        exclude_joints=getattr(config, "relative_exclude_joints", []),
        action_names=getattr(config, "action_feature_names", None),
    )

    input_steps = [
        steps.rename_observations,
        steps.add_batch_dim,
        steps.to_device,
        ImagePrepProcessorStep(resize_to=config.resize_images_to or None),
        relative_step,
        steps.normalize,
    ]
    gripper_dim = config.resolved_gripper_dim
    _warn_if_gripper_steps_are_misconfigured(config, gripper_dim, dataset_stats)

    output_steps: list[ProcessorStep] = []
    if config.clip_normalized_actions:
        # Clipping to [-1, 1] is a range assertion under MIN_MAX, but under MEAN_STD the same
        # clamp truncates every action beyond 1 sigma. That shows up as a hesitant, low-amplitude
        # policy with no error anywhere, so refuse to add the step instead of honoring the flag.
        action_norm_mode = config.normalization_mapping.get("ACTION")
        if action_norm_mode == NormalizationMode.MIN_MAX:
            output_steps.append(ClipActionsProcessorStep())
        else:
            logging.warning(
                f"`clip_normalized_actions=True` is ignored: it clips normalized actions to "
                f"[-1, 1], which is only a no-op bound under MIN_MAX, but ACTION uses "
                f"{getattr(action_norm_mode, 'value', action_norm_mode)}. Under MEAN_STD this "
                f"would clamp every action to 1 sigma."
            )
    if config.pre_snap_gripper_action:
        output_steps.append(
            PreSnapGripperProcessorStep(gripper_dim=gripper_dim, threshold=config.gripper_threshold)
        )
    # NOTE: unlike the default policy unnormalizer (output features only), VLA-JEPA
    # unnormalizes over BOTH input and output features.
    output_steps.append(
        UnnormalizerProcessorStep(
            features=features,
            norm_map=config.normalization_mapping,
            stats=dataset_stats,
        )
    )
    # Reverse the relative conversion on the unnormalized action, before gripper binarization.
    # gripper is kept absolute by relative_exclude_joints, so the two steps touch disjoint dims.
    output_steps.append(
        AbsoluteActionsProcessorStep(enabled=config.use_relative_actions, relative_step=relative_step)
    )
    if config.binarize_gripper_action:
        output_steps.append(
            BinarizeGripperProcessorStep(gripper_dim=gripper_dim, threshold=config.gripper_threshold)
        )
    output_steps.append(steps.to_cpu)
    return make_policy_processor_pipelines(input_steps=input_steps, output_steps=output_steps)
