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

"""FineART-VLA, built on Physical Intelligence's pi0.5 and openpi implementation."""

import math
from dataclasses import asdict, dataclass, field
from pathlib import Path

from lerobot.configs import PreTrainedConfig

from ..pi05.configuration_pi05 import PI05Config


def _fineart_vla_default_recipe() -> dict:
    """Embed the subtask contract without importing optional dataset dependencies."""
    return {
        "blend": {
            "high_level_subtask": {
                "weight": 0.30,
                "messages": [
                    {"role": "user", "content": "${task}", "stream": "high_level"},
                    {
                        "role": "assistant",
                        "content": "${subtask}",
                        "stream": "high_level",
                        "target": True,
                        "if_present": "subtask",
                    },
                ],
            },
            "low_level_execution": {
                "weight": 0.70,
                "messages": [
                    {
                        "role": "user",
                        "content": "${subtask}",
                        "stream": "low_level",
                        "if_present": "subtask",
                    }
                ],
            },
        }
    }


@PreTrainedConfig.register_subclass("fineart_vla")
@dataclass
class FineARTVLAConfig(PI05Config):
    """FineART-VLA: pi0.5 with recipe-driven text and action supervision."""

    # Recipe / language stack ---------------------------------------------
    recipe_path: str | None = None
    """Optional file override for the embedded default recipe; resolved into ``recipe`` once, then cleared."""

    recipe: dict | None = field(default_factory=_fineart_vla_default_recipe)
    """Serialized training/runtime contract; ``None`` selects the plain PI0.5 prompt."""

    memory_scratchpad: bool = False
    """Legacy checkpoint guard: combined responses require a scratchpad-aware controller."""

    # Balance frequent recipe text supervision against the paper's α=10 flow weight.
    text_loss_weight: float = 1.0
    """Text cross-entropy weight; ``0`` disables it."""

    flow_loss_weight: float = 10.0
    """Flow-matching loss weight."""

    # Backbone training ---------------------------------------------------
    unfreeze_lm_head: bool = True
    """Deprecated and ignored: text and FAST supervision always train the VLM, language head included.
    Kept so existing checkpoint configs still load."""

    # FAST adds discrete-action CE to the text and flow objectives from paper §III.B-C.
    enable_fast_action_loss: bool = True
    """Add FAST action-token cross-entropy."""

    action_tokenizer_name: str = "physical-intelligence/fast"
    """FAST tokenizer identifier."""

    max_action_tokens: int = 256
    """Maximum FAST tokens per action chunk."""

    fast_skip_tokens: int = 1152
    """Reserved vocabulary IDs skipped by FAST token mapping."""

    fast_action_loss_weight: float = 1.0
    """FAST action-token loss weight."""

    auto_fit_fast_tokenizer: bool = True
    """Fit (and cache) a FAST tokenizer on the training dataset, starting from ``action_tokenizer_name``."""

    fast_tokenizer_cache_dir: str = "~/.cache/lerobot/fast_tokenizers"
    """Cache directory for fitted FAST tokenizers."""

    fast_tokenizer_fit_samples: int = 1024
    """Action chunks sampled for tokenizer fitting."""

    fast_tokenizer_validation_samples: int = 256
    """Held-out chunks used for tokenizer validation."""

    fast_tokenizer_max_reconstruction_rmse: float = 0.10
    """Maximum validation reconstruction RMSE."""

    fast_tokenizer_max_dim_rmse: float = 0.20
    """Maximum per-dimension validation RMSE."""

    # Knowledge insulation detaches VLM K/V from action-loss gradients (paper §III.B).
    knowledge_insulation: bool = True
    """Detach VLM keys and values from action-loss gradients."""

    # Optional training backends. Defaults preserve the eager/SDPA path.
    use_liger_kernels: bool = False
    """Fuse PaliGemma RoPE/GeGLU with Liger. Opt-in: it patches transformers for the whole process."""

    use_flashrt_adarms: bool = False
    """Use FlashRT adaptive RMSNorm kernels."""

    use_compiled_text_ce: bool = False
    """Compile text and FAST cross-entropy."""

    use_compiled_vision: bool = False
    """Compile the SigLIP vision tower."""

    use_flex_attention: bool = False
    """Use FlexAttention for knowledge insulation."""

    use_manual_attention: bool = False
    """Use manual attention for profiled KI shapes."""

    manual_attention_scope: str = "all"
    """Manual-attention scope: ``all`` or ``action``."""

    # Scale language-head updates relative to the base optimizer schedule.
    lm_head_lr_scale: float = 1.0

    # Scale backbone and action-expert optimizer groups independently.
    backbone_lr_scale: float = 1.0
    action_expert_lr_scale: float = 1.0

    # Opt-in stabilization: leave saved checkpoints and default training unchanged.
    conditioning_lr_scale: float = 1.0
    """Additional LR multiplier for time MLPs and adaptive scale/shift/gate projections."""

    # Reuse each VLM prefix across independent denoising draws; 1 restores single-draw flow.
    flow_num_repeats: int = 5

    # PaLM-style z-loss stabilizes large-vocabulary CE; 0 disables it.
    text_ce_z_loss_weight: float = 1e-4

    # Applied to FineART-VLA parameter groups, leaving the shared AdamW config unchanged.
    optimizer_foreach: bool | None = False
    optimizer_fused: bool | None = True

    def __post_init__(self) -> None:
        super().__post_init__()
        if not math.isfinite(self.conditioning_lr_scale) or self.conditioning_lr_scale < 0:
            raise ValueError("conditioning_lr_scale must be finite and nonnegative")
        if self.recipe_path is not None:
            from lerobot.datasets.recipe import resolve_recipe_override

            path = self.recipe_path
            if path is not None and not Path(path).exists():
                packaged = Path(__file__).parents[2] / "configs" / path
                if packaged.exists():
                    path = str(packaged)
            resolved_recipe = resolve_recipe_override(self.recipe, path)
            if resolved_recipe is None:
                raise ValueError("recipe_path must resolve to a training recipe")
            self.recipe = asdict(resolved_recipe)
            # The resolved recipe is now part of the config: later loads must not re-read a file
            # that may have changed since, so the saved checkpoint keeps the recipe it trained with.
            self.recipe_path = None
        if self.enable_fast_action_loss and self.recipe is None:
            raise ValueError("FineART-VLA FAST action loss requires recipe_path to build action supervision.")
        # Text and FAST cross-entropy supervise the VLM itself (knowledge insulation only keeps the
        # flow loss out of it), so they need a trainable backbone.
        if self.train_expert_only and (self.text_loss_weight > 0 or self.enable_fast_action_loss):
            raise ValueError(
                "train_expert_only=true freezes the VLM that the text and FAST losses train. "
                "Set train_expert_only=false, or disable both (text_loss_weight=0, "
                "enable_fast_action_loss=false) to train only the action expert."
            )
        if self.flow_num_repeats < 1:
            raise ValueError(f"flow_num_repeats must be >= 1, got {self.flow_num_repeats}")
        if self.fast_tokenizer_validation_samples < 1:
            raise ValueError("fast_tokenizer_validation_samples must be >= 1")
        if self.fast_tokenizer_max_reconstruction_rmse <= 0 or self.fast_tokenizer_max_dim_rmse <= 0:
            raise ValueError("FAST tokenizer reconstruction thresholds must be positive")
        if self.manual_attention_scope not in {"all", "action"}:
            raise ValueError(
                f"manual_attention_scope must be 'all' or 'action', got {self.manual_attention_scope!r}"
            )
        if self.use_flex_attention and self.use_manual_attention:
            raise ValueError("use_flex_attention and use_manual_attention are mutually exclusive")
        if self.use_flex_attention and self.flow_num_repeats == 1:
            raise ValueError("use_flex_attention requires flow_num_repeats > 1")
        if not self.knowledge_insulation and (
            self.use_flex_attention or self.use_manual_attention or self.use_flashrt_adarms
        ):
            raise ValueError("KI attention and AdaRMS optimizations require knowledge_insulation=True")
