# Copyright 2026 Physical Intelligence and 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: joint flow/text training and hierarchical language inference.

Derived from Physical Intelligence's pi0.5/openpi and LeRobot's PI05 implementation.
Original architecture and implementation: https://github.com/Physical-Intelligence/openpi
"""

from __future__ import annotations

import logging
from contextlib import nullcontext
from typing import TYPE_CHECKING, Any, Unpack

import torch
from safetensors import safe_open
from safetensors.torch import load_file
from torch import Tensor
from torch.nn import functional
from torch.nn.attention.flex_attention import create_block_mask, flex_attention

from lerobot.utils.constants import (
    ACTION,
    ACTION_CODE_TOKEN_MASK,
    ACTION_TOKEN_MASK,
    ACTION_TOKENS,
    OBS_LANGUAGE_ATTENTION_MASK,
    OBS_LANGUAGE_CAUSAL_MARKS,
    OBS_LANGUAGE_TOKENS,
)
from lerobot.utils.device_utils import resolve_safetensors_device
from lerobot.utils.import_utils import (
    _kernels_available,
    _liger_kernel_available,
    _transformers_available,
    require_package,
)

from ..common.vla_utils import make_att_2d_masks
from ..pi05.modeling_pi05 import (
    PI05Policy,
    PI05Pytorch,
)
from ..pretrained import RTCActionSelectKwargs
from ..utils import log_model_loading_keys
from .configuration_fineart_vla import FineARTVLAConfig
from .processor_fineart_vla import (  # noqa: F401
    make_fineart_vla_pre_post_processors,
    register_paligemma_loc_tokens,
)

if TYPE_CHECKING or _transformers_available:
    from transformers import AutoTokenizer

    from ..pi_gemma import sdpa_attention_forward
else:
    AutoTokenizer = None
    sdpa_attention_forward = None

# Optional fused training kernels (lerobot[fineart_vla_kernels]); checked again where used.
if TYPE_CHECKING or _liger_kernel_available:
    from liger_kernel.transformers import apply_liger_kernel_to_paligemma
    from liger_kernel.transformers.fused_linear_cross_entropy import LigerFusedLinearCrossEntropyLoss
if TYPE_CHECKING or _kernels_available:
    from kernels import get_kernel

logger = logging.getLogger(__name__)

# Tokenizers with PaliGemma's <loc> tokens registered, cached per name for text generation.
_LOC_TOKENIZER_CACHE: dict[str, Any] = {}


def _get_loc_tokenizer(tok_name: str, auto_tokenizer_cls: Any, register_loc_fn: Any) -> Any:
    tokenizer = _LOC_TOKENIZER_CACHE.get(tok_name)
    if tokenizer is None:
        tokenizer = register_loc_fn(auto_tokenizer_cls.from_pretrained(tok_name))
        _LOC_TOKENIZER_CACHE[tok_name] = tokenizer
    return tokenizer


def _last_valid_prefix_hidden(hidden: Tensor, mask: Tensor) -> Tensor:
    """Select the prompt endpoint, including noncontiguous masked image/text padding."""
    positions = torch.arange(mask.shape[1], device=mask.device).expand_as(mask)
    indices = positions.masked_fill(~mask.bool(), -1).amax(dim=1).clamp_min(0)
    return hidden.gather(1, indices[:, None, None].expand(-1, 1, hidden.shape[-1]))


class FineARTVLAPytorch(PI05Pytorch):  # see openpi `PI0Pytorch`
    """FineART-VLA core model: PI0.5 with KI-aware joint layers and text supervision."""

    use_hf_vision_checkpointing_api = True
    checkpoint_vision_embeddings = False
    fp32_joint_attention = True
    use_typed_attention_masks = True
    use_on_device_suffix_mask = True
    precompute_denoise_times = True

    def forward(
        self,
        images,
        img_masks,
        tokens,
        masks,
        actions,
        noise,
        time,
        prefix_mask: Tensor | None = None,
        states=None,
        state_masks=None,
    ) -> Tensor:
        """Do a full training forward pass and compute the loss."""
        if prefix_mask is not None or states is not None or state_masks is not None:
            # Training-time RTC / proprioceptive memory inputs are handled by the PI05 path.
            return super().forward(
                images,
                img_masks,
                tokens,
                masks,
                actions,
                noise,
                time,
                prefix_mask=prefix_mask,
                states=states,
                state_masks=state_masks,
            )
        time_expanded = time[:, None, None]
        x_t = time_expanded * noise + (1 - time_expanded) * actions
        u_t = noise - actions

        # The prefix has no gradient path to a flow-only loss under KI.
        suppress_prefix = bool(getattr(self.config, "knowledge_insulation", False))
        with torch.no_grad() if suppress_prefix else nullcontext():
            prefix_embs, prefix_pad_masks, prefix_att_masks = self.embed_prefix(
                images, img_masks, tokens, masks
            )
        suffix_embs, suffix_pad_masks, suffix_att_masks, adarms_cond = self.embed_suffix(x_t, time)

        if (
            self.paligemma_with_expert.paligemma.model.language_model.layers[0].self_attn.q_proj.weight.dtype
            == torch.bfloat16
        ):
            suffix_embs = suffix_embs.to(dtype=torch.bfloat16)
            prefix_embs = prefix_embs.to(dtype=torch.bfloat16)

        pad_masks = torch.cat([prefix_pad_masks, suffix_pad_masks], dim=1)
        att_masks = torch.cat([prefix_att_masks, suffix_att_masks], dim=1)

        att_2d_masks = make_att_2d_masks(pad_masks, att_masks)
        position_ids = torch.cumsum(pad_masks, dim=1) - 1

        att_2d_masks_4d = self._prepare_attention_masks_4d(att_2d_masks, dtype=prefix_embs.dtype)

        # Transformer and vision layers own their checkpoint boundaries.
        ki_kwargs: dict[str, Any] = {"suppress_prefix_grads": True} if suppress_prefix else {}
        (_, suffix_out), _ = self.paligemma_with_expert.forward(
            attention_mask=att_2d_masks_4d,
            position_ids=position_ids,
            past_key_values=None,
            inputs_embeds=[prefix_embs, suffix_embs],
            use_cache=False,
            adarms_cond=[None, adarms_cond],
            **ki_kwargs,
        )

        suffix_out = suffix_out[:, -self.config.chunk_size :]
        suffix_out = suffix_out.to(dtype=torch.float32)

        def action_out_proj_func(suffix_out):
            return self.action_out_proj(suffix_out)

        v_t = self._apply_checkpoint(action_out_proj_func, suffix_out)

        return functional.mse_loss(u_t, v_t, reduction="none")

    def denoise_step(
        self,
        prefix_pad_masks,
        past_key_values,
        x_t,
        timestep,
    ):
        """Apply one denoising step of the noise `x_t` at a given timestep."""
        suffix_embs, suffix_pad_masks, suffix_att_masks, adarms_cond = self.embed_suffix(x_t, timestep)

        suffix_len = suffix_pad_masks.shape[1]
        batch_size = prefix_pad_masks.shape[0]
        prefix_len = prefix_pad_masks.shape[1]

        prefix_pad_2d_masks = prefix_pad_masks[:, None, :].expand(batch_size, suffix_len, prefix_len)
        suffix_att_2d_masks = make_att_2d_masks(suffix_pad_masks, suffix_att_masks)
        full_att_2d_masks = torch.cat([prefix_pad_2d_masks, suffix_att_2d_masks], dim=2)

        prefix_offsets = torch.sum(prefix_pad_masks, dim=-1)[:, None]
        position_ids = prefix_offsets + torch.cumsum(suffix_pad_masks, dim=1) - 1

        full_att_2d_masks_4d = self._prepare_attention_masks_4d(full_att_2d_masks, dtype=suffix_embs.dtype)
        self.paligemma_with_expert.gemma_expert.model.config._attn_implementation = "eager"  # noqa: SLF001

        # Crop appended suffix K/V after each step instead of copying the read-only prefix cache.
        outputs_embeds, _ = self.paligemma_with_expert.forward(
            attention_mask=full_att_2d_masks_4d,
            position_ids=position_ids,
            past_key_values=past_key_values,
            inputs_embeds=[None, suffix_embs],
            use_cache=False,
            adarms_cond=[None, adarms_cond],
        )
        past_key_values.crop(prefix_len)

        suffix_out = outputs_embeds[1]
        suffix_out = suffix_out[:, -self.config.chunk_size :]
        suffix_out = suffix_out.to(dtype=torch.float32)
        return self.action_out_proj(suffix_out)

    def sample_actions(
        self,
        images,
        img_masks,
        tokens,
        masks,
        noise=None,
        num_steps=None,
        lang_causal_marks=None,
        **kwargs,
    ) -> Tensor:
        """Sample actions, optionally marking trailing language positions causal.

        ``lang_causal_marks`` (B, L_lang bool) flags generated-subtask tokens so
        joint-sequence checkpoints see the same causal prefix layout at
        inference as during training (``_mark_target_span_causal``).
        """
        self._lang_causal_marks = lang_causal_marks
        try:
            return super().sample_actions(
                images, img_masks, tokens, masks, noise=noise, num_steps=num_steps, **kwargs
            )
        finally:
            self._lang_causal_marks = None

    def embed_prefix(self, images, img_masks, tokens, masks, states=None, state_masks=None):
        prefix_embs, prefix_pad, prefix_att = super().embed_prefix(
            images, img_masks, tokens, masks, states, state_masks
        )
        marks = getattr(self, "_lang_causal_marks", None)
        if marks is not None:
            prefix_att = _apply_causal_language_marks(prefix_att, marks.to(prefix_att.device))
        return prefix_embs, prefix_pad, prefix_att


# The universal `physical-intelligence/fast` tokenizer (and dataset refits of it)
# uses 1024 BPE codes; text generation must mask any that map below the <loc> range.
_FAST_ACTION_VOCAB_SIZE = 1024


_HF_KERNELS_ENABLED = False


def _enable_hf_kernels() -> None:
    """Patch supported PaliGemma operations before constructing the model."""
    global _HF_KERNELS_ENABLED
    if _HF_KERNELS_ENABLED:
        return
    if not _liger_kernel_available:
        logger.warning(
            "FineART-VLA: liger-kernel is not installed; skipping fused Triton "
            "kernels. Install with ``pip install 'lerobot[fineart_vla_kernels]'``."
        )
        return
    apply_liger_kernel_to_paligemma(
        rope=True,
        geglu=True,
        # Liger LayerNorm regresses at SigLIP shapes; RoPE and GeGLU remain enabled.
        layer_norm=False,
        rms_norm=False,
        cross_entropy=False,
        fused_linear_cross_entropy=False,
    )
    _HF_KERNELS_ENABLED = True
    logger.info("FineART-VLA: HF kernels (Liger) enabled — rope, geglu fused.")


def _reduce_action_loss(per_sample: Tensor, predict_actions_t: Tensor | None, reduction: str) -> Tensor:
    """Mask non-action samples and apply the requested batch reduction."""
    if predict_actions_t is None:
        return per_sample if reduction == "none" else per_sample.mean()
    mask = predict_actions_t.to(per_sample.dtype)
    if reduction == "none":
        return per_sample * mask
    return (per_sample * mask).sum() / mask.sum().clamp(min=1.0)


# Materialized logits win at VLA token counts; larger dense targets use Liger.
_LOGITS_CE_MAX_POSITIONS = 2048


def _lin_ce_small(
    flat_hidden: Tensor,
    lm_head_weight: Tensor,
    flat_labels: Tensor,
    z_loss_weight: float = 0.0,
) -> Tensor:
    """Small-N linear CE on materialized logits (see ``_lin_ce_flat``)."""
    logits = (flat_hidden @ lm_head_weight.t()).float()
    n_valid = (flat_labels != -100).sum().clamp(min=1)
    loss = functional.cross_entropy(logits, flat_labels, ignore_index=-100, reduction="sum") / n_valid
    if z_loss_weight > 0:
        lse = torch.logsumexp(logits, dim=-1)
        valid = (flat_labels != -100).to(lse.dtype)
        loss = loss + float(z_loss_weight) * (lse.square() * valid).sum() / n_valid
    return loss


# Built lazily so importing this module does not invoke Dynamo.
_compiled_lin_ce_small = None


def _get_compiled_lin_ce_small():
    global _compiled_lin_ce_small
    if _compiled_lin_ce_small is None:
        _compiled_lin_ce_small = torch.compile(_lin_ce_small, dynamic=False)
    return _compiled_lin_ce_small


def _lin_ce_flat(
    flat_hidden: Tensor,
    lm_head_weight: Tensor,
    flat_labels: Tensor,
    z_loss_weight: float = 0.0,
    compiled: bool = False,
) -> Tensor:
    """Dispatch sparse targets to fixed logits buckets and dense targets to Liger."""
    if flat_hidden.shape[0] > _LOGITS_CE_MAX_POSITIONS:
        valid = flat_labels != -100
        compact_hidden = flat_hidden[valid]
        compact_labels = flat_labels[valid]
        compact_rows = compact_hidden.shape[0]

        if compact_rows == 0:
            return _lin_ce_flat(
                functional.pad(compact_hidden, (0, 0, 0, 1)),
                lm_head_weight,
                functional.pad(compact_labels, (0, 1), value=-100),
                z_loss_weight,
                compiled=compiled,
            )

        # Fixed power-of-two buckets avoid shape churn while keeping sparse
        # supervision on the materialized-logits path.
        bucket_rows = 1 << (compact_rows - 1).bit_length()
        if bucket_rows < flat_hidden.shape[0]:
            weighted_losses = []
            for start in range(0, compact_rows, _LOGITS_CE_MAX_POSITIONS):
                end = min(start + _LOGITS_CE_MAX_POSITIONS, compact_rows)
                rows = end - start
                chunk_rows = 1 << (rows - 1).bit_length()
                hidden_chunk = compact_hidden[start:end]
                labels_chunk = compact_labels[start:end]
                pad_rows = chunk_rows - rows
                if pad_rows:
                    hidden_chunk = functional.pad(hidden_chunk, (0, 0, 0, pad_rows))
                    labels_chunk = functional.pad(labels_chunk, (0, pad_rows), value=-100)
                chunk_loss = _lin_ce_flat(
                    hidden_chunk,
                    lm_head_weight,
                    labels_chunk,
                    z_loss_weight,
                    compiled=compiled,
                )
                weighted_losses.append(chunk_loss * rows)
            return torch.stack(weighted_losses).sum() / compact_rows

    if flat_hidden.shape[0] <= _LOGITS_CE_MAX_POSITIONS:
        fn = _get_compiled_lin_ce_small() if compiled else _lin_ce_small
        return fn(flat_hidden, lm_head_weight, flat_labels, z_loss_weight)

    # Keep Liger optional for inference-only installations.
    require_package("liger-kernel", extra="fineart_vla_kernels", import_name="liger_kernel")

    loss_fn = LigerFusedLinearCrossEntropyLoss(
        ignore_index=-100,
        lse_square_scale=float(z_loss_weight),
        reduction="mean",
    )
    return loss_fn(lm_head_weight, flat_hidden, flat_labels)


def _shifted_lin_ce(
    hidden: Tensor,
    lm_head_weight: Tensor,
    labels: Tensor,
    z_loss_weight: float = 0.0,
    compiled: bool = False,
    reduction: str = "mean",
) -> Tensor:
    """Compute next-token CE through the shape-aware linear-CE dispatcher."""
    shift_hidden = hidden[:, :-1, :].contiguous()
    shift_labels = labels[:, 1:].contiguous().long()
    if reduction == "none":
        return torch.stack(
            [
                _lin_ce_flat(
                    sample_hidden.to(lm_head_weight.dtype),
                    lm_head_weight,
                    sample_labels,
                    z_loss_weight,
                    compiled=compiled,
                )
                for sample_hidden, sample_labels in zip(shift_hidden, shift_labels, strict=True)
            ]
        )
    batch_size, target_length, hidden_size = shift_hidden.shape
    flat_hidden = shift_hidden.reshape(batch_size * target_length, hidden_size)
    flat_labels = shift_labels.reshape(batch_size * target_length)
    # Match the dtype the eager path used: cast hidden to the lm_head's
    # weight dtype so bf16 weights see bf16 activations.
    flat_hidden = flat_hidden.to(lm_head_weight.dtype)
    return _lin_ce_flat(flat_hidden, lm_head_weight, flat_labels, z_loss_weight, compiled=compiled)


def _mark_target_span_causal(
    prefix_att_masks: Tensor, text_labels: Tensor, lang_start: int, lang_end: int
) -> Tensor:
    """Make supervised language targets causal while leaving prompts bidirectional."""
    att = prefix_att_masks.clone()
    n = min(text_labels.shape[1], lang_end - lang_start)
    if n <= 0:
        return att
    target = text_labels[:, :n] != -100  # (B, n) bool
    seg = att[:, lang_start : lang_start + n].bool()
    att[:, lang_start : lang_start + n] = seg | target
    return att


def _apply_causal_language_marks(prefix_att_masks: Tensor, marks: Tensor) -> Tensor:
    """OR per-token causal marks into the trailing language segment of a prefix."""
    att = prefix_att_masks.clone()
    n = min(marks.shape[1], att.shape[1])
    if n <= 0:
        return att
    seg = att[:, -n:].bool()
    att[:, -n:] = (seg | marks[:, -n:].bool()).to(att.dtype)
    return att


def _fast_lin_ce(
    hidden: Tensor,
    lm_head_weight: Tensor,
    action_tokens: Tensor,
    action_code_mask: Tensor,
    predict_actions_t: Tensor | None,
    compiled: bool = False,
    reduction: str = "mean",
) -> Tensor:
    """Compute FAST token CE over the enabled action-code positions."""
    shift_hidden = hidden[:, :-1, :].contiguous()
    shift_targets = action_tokens[:, 1:].contiguous().long()
    shift_valid = action_code_mask[:, 1:].contiguous().bool()
    if predict_actions_t is not None:
        sample_mask = predict_actions_t[:, None].expand_as(shift_valid)
        shift_valid = shift_valid & sample_mask
    # Encode the mask with ignore_index to avoid a host sync and preserve graph capture.
    shift_targets = torch.where(shift_valid, shift_targets, torch.full_like(shift_targets, -100))

    if reduction == "none":
        return torch.stack(
            [
                _lin_ce_flat(
                    sample_hidden.to(lm_head_weight.dtype),
                    lm_head_weight,
                    sample_labels,
                    compiled=compiled,
                )
                for sample_hidden, sample_labels in zip(shift_hidden, shift_targets, strict=True)
            ]
        )

    valid_counts = shift_valid.sum(dim=1)
    active_samples = valid_counts > 0
    if not bool(active_samples.any().item()):
        return shift_hidden.sum() * 0.0

    weighted_losses = []
    active_count = active_samples.sum()
    for token_count in torch.unique(valid_counts[active_samples]).tolist():
        group = active_samples & valid_counts.eq(token_count)
        group_size = group.sum()
        group_hidden = shift_hidden[group].reshape(-1, shift_hidden.shape[-1]).to(lm_head_weight.dtype)
        group_labels = shift_targets[group].reshape(-1)
        group_loss = _lin_ce_flat(
            group_hidden,
            lm_head_weight,
            group_labels,
            compiled=compiled,
        )
        weighted_losses.append(group_loss * group_size)
    return torch.stack(weighted_losses).sum() / active_count


# ----------------------------------------------------------------------
# Knowledge insulation helpers
# ----------------------------------------------------------------------
# Action queries consume detached VLM K/V. Flow-only callers may additionally
# suppress the now-dead prefix graph without changing forward values.


# Consumer GPUs need smaller FlexAttention backward tiles at head_dim=256.
_FLEX_SHRUNK_TILES = {"BLOCK_M1": 32, "BLOCK_N1": 64, "BLOCK_M2": 64, "BLOCK_N2": 32}
_flex_kernel_options: dict[int, dict | None] = {}
_flex_fns: tuple | None | bool = None


def _get_flex_kernel_options(device: torch.device) -> dict | None:
    if device.type != "cuda" or not torch.cuda.is_available():
        return None
    device_index = device.index if device.index is not None else torch.cuda.current_device()
    if device_index not in _flex_kernel_options:
        smem = torch.cuda.get_device_properties(
            device_index
        ).shared_memory_per_block_optin  # spellchecker:disable-line
        _flex_kernel_options[device_index] = _FLEX_SHRUNK_TILES if smem < 128 * 1024 else None
    return _flex_kernel_options[device_index]


def _get_flex_fns(device: torch.device):
    """Return compiled FlexAttention helpers when available."""
    global _flex_fns
    if device.type != "cuda" or not torch.cuda.is_available():
        return None
    if _flex_fns is None:
        try:
            _flex_fns = (
                torch.compile(flex_attention, dynamic=False),
                torch.compile(create_block_mask, dynamic=False),
            )
            _get_flex_kernel_options(device)
        except Exception as exc:
            logger.warning("FineART-VLA: FlexAttention unavailable (%s); using SDPA.", exc)
            _flex_fns = False
    return _flex_fns or None


class _FlexMaskBuilder:
    """Build KI masks while retaining stable compiled mask callables."""

    def __init__(self):
        self._key = None

    def build(self, prefix_pad, prefix_att, non_fast_prefix_len, k, chunk):
        _, create_bm = _get_flex_fns(prefix_pad.device)
        b, p = prefix_pad.shape
        a = k * chunk
        device = prefix_pad.device
        key = (b, p, a, int(non_fast_prefix_len), device)
        if self._key != key:
            self._key = key
            self._pad = torch.empty(b, p, dtype=torch.bool, device=device)
            self._cum = torch.empty(b, p, dtype=torch.long, device=device)
            pad, cum = self._pad, self._cum
            nf = int(non_fast_prefix_len)

            def vlm_rows(bi, h, q_idx, kv_idx):
                kv_p = kv_idx.clamp(max=p - 1)
                ok = (cum[bi, kv_p] <= cum[bi, q_idx]) & pad[bi, kv_p] & pad[bi, q_idx]
                return (kv_idx < p) & ok

            def action_rows(bi, h, q_idx, kv_idx):
                kv_p = kv_idx.clamp(max=p - 1)
                to_prefix = (kv_idx < nf) & pad[bi, kv_p]
                same_block = (q_idx // chunk) == ((kv_idx - p) // chunk)
                return to_prefix | ((kv_idx >= p) & same_block)

            self._vlm_mod, self._action_mod = vlm_rows, action_rows

        self._pad.copy_(prefix_pad)
        self._cum.copy_(torch.cumsum(prefix_att.to(torch.long), dim=1))
        s = p + a
        bm_vlm = create_bm(self._vlm_mod, B=b, H=None, Q_LEN=p, KV_LEN=s, device=device)
        bm_action = create_bm(self._action_mod, B=b, H=None, Q_LEN=a, KV_LEN=s, device=device)
        return bm_vlm, bm_action


# Lazily loaded FlashRT AdaRMS backend; unsupported cases use eager PyTorch.
_flashrt_adarms_cache = None


def _get_adarms_backend():
    global _flashrt_adarms_cache
    if _flashrt_adarms_cache is None:
        try:
            require_package("kernels", extra="fineart_vla_kernels")
            _flashrt_adarms_cache = get_kernel("flashrt/flashrt-adarms-train", revision="v1")
        except Exception as exc:
            logger.warning(
                "FineART-VLA: flashrt-adarms-train unavailable (%s); using the eager norm path.",
                exc,
            )
            _flashrt_adarms_cache = False
    return _flashrt_adarms_cache or None


class _FlashRTNorms:
    """FlashRT AdaRMS kernels behind PI0.5's ``norm_backend`` hooks for the joint layers."""

    def __init__(self, kernels: Any) -> None:
        self._kernels = kernels

    def input_norm(self, norm, x, cond):
        if cond is not None and norm.dense is not None:
            return self._kernels.adarms(x, norm.dense(cond), norm.eps, True)
        if norm.dense is None:
            return self._kernels.adarms(x, norm.weight, norm.eps, False)
        return norm(x, cond=cond)

    def residual_norm(self, norm, residual, out, gate, cond):
        if cond is not None and norm.dense is not None:
            return self._kernels.resgate_adarms(residual, out, gate, norm.dense(cond), norm.eps, True)
        return self._kernels.resgate_adarms(residual, out, gate, norm.weight, norm.eps, False)


def _manual_attention_part(qs, ks, vs, m, scale):
    """Materialized-logits GQA with an FP32 softmax."""
    batch_size, num_heads, query_length, head_dim = qs.shape
    num_kv_heads = ks.shape[1]
    if num_kv_heads != num_heads:
        groups = num_heads // num_kv_heads
        grouped_queries = qs.reshape(batch_size, num_kv_heads, groups * query_length, head_dim)
        logits = (grouped_queries @ ks.transpose(-1, -2)).reshape(batch_size, num_heads, query_length, -1)
    else:
        logits = qs @ ks.transpose(-1, -2)
    logits = logits * scale + m
    p = logits.float().softmax(dim=-1).to(qs.dtype)
    out = (
        (p.reshape(batch_size, num_kv_heads, groups * query_length, -1) @ vs).reshape(
            batch_size, num_heads, query_length, head_dim
        )
        if num_kv_heads != num_heads
        else p @ vs
    )
    return out.transpose(1, 2).contiguous()


# Knowledge insulation keeps the forward equivalent while detaching VLM K/V for action-query gradients.
_manual_attention = None


def _get_manual_attention():
    """Load the Hub implementation, with the inline function as fallback."""
    global _manual_attention
    if _manual_attention is None:
        part = _manual_attention_part
        try:
            require_package("kernels", extra="fineart_vla_kernels")
            _hub = getattr(
                get_kernel("flashrt/flashrt-flex-attention-train", revision="v1"),
                "manual_attention_part",
                None,
            )
            if _hub is not None:

                def part(qs, ks, vs, m, scale, _hub=_hub):
                    return _hub(qs, ks, vs, m, scale).transpose(1, 2).contiguous()

                logger.info("FineART-VLA: manual attention backed by flashrt-flex-attention-train (Hub).")
        except Exception as exc:
            logger.info(
                "FineART-VLA: flashrt-flex-attention-train unavailable (%s); using the inline manual-attention path.",
                exc,
            )
        _manual_attention = torch.compile(part, dynamic=False)
    return _manual_attention


def _flex_attention_fn(flex_masks):
    """Knowledge-insulation attention through FlexAttention block masks (one per query group)."""
    vlm_block_mask, action_block_mask = flex_masks

    def attend(self_attn, query, key, value, mask, scaling, part):
        flex_attn, _ = _get_flex_fns(query.device)
        return flex_attn(
            query,
            key,
            value,
            block_mask=vlm_block_mask if part == "vlm" else action_block_mask,
            scale=scaling,
            enable_gqa=self_attn.num_key_value_groups > 1,
            kernel_options=_get_flex_kernel_options(query.device),
        ).transpose(1, 2)

    return attend


def _manual_attention_fn(scope: str):
    """Knowledge-insulation attention with materialized FP32-softmax logits (``all`` or ``action`` rows)."""
    manual_fn = _get_manual_attention()

    def attend(self_attn, query, key, value, mask, scaling, part):
        # SDPA and the manual kernel need the additive bias in the query dtype.
        mask = mask.to(dtype=query.dtype)
        if part == "vlm" and scope == "action":
            return sdpa_attention_forward(self_attn, query, key, value, mask, scaling)[0]
        return manual_fn(query, key, value, mask, scaling)

    return attend


class FineARTVLAPolicy(PI05Policy):
    """FineART-VLA, extending Physical Intelligence's π0.5 with language supervision.

    It inherits unchanged PI0.5 policy behavior and replaces the core model with
    the joint flow/text implementation below.
    """

    config_class = FineARTVLAConfig
    name = "fineart_vla"
    model_class = FineARTVLAPytorch
    use_native_pretrained_loader = True

    def __init__(self, config: FineARTVLAConfig, **kwargs: Any) -> None:
        # Patch before constructing Gemma/SigLIP layers; the operation is optional and idempotent.
        if config.use_liger_kernels:
            _enable_hf_kernels()
        super().__init__(config, **kwargs)

        # Knowledge insulation runs inside PI0.5's joint layer; stock PI0.5 keeps it off.
        self.model.paligemma_with_expert.knowledge_insulation = config.knowledge_insulation
        if config.knowledge_insulation:
            logger.info(
                "FineART-VLA: knowledge insulation enabled — action→VLM K/V gradients are blocked in attention."
            )
            if config.use_flashrt_adarms:
                kernels = _get_adarms_backend()
                if kernels is not None:
                    self._flashrt_adarms = _FlashRTNorms(kernels)
                    logger.info("FineART-VLA: FlashRT adaRMS training kernels enabled.")

        if config.use_compiled_vision:
            _tower = self.model.paligemma_with_expert.paligemma.model.vision_tower
            _tower_eager_fwd = _tower.forward
            _tower_compiled_fwd = torch.compile(_tower_eager_fwd, dynamic=False)

            def _tower_dispatch(*args, _e=_tower_eager_fwd, _c=_tower_compiled_fwd, **kwargs):
                if torch.is_grad_enabled():
                    return _e(*args, **kwargs)
                return _c(*args, **kwargs)

            _tower.forward = _tower_dispatch
            logger.info("FineART-VLA: SigLIP vision tower compiled for no-grad passes.")

        # Cache the fixed K-repeat action mask outside the training step.
        if config.flow_num_repeats > 1:
            self.register_buffer(
                "_flow_block_diag",
                torch.block_diag(
                    *[
                        torch.ones(
                            config.chunk_size, config.chunk_size, dtype=torch.bool, device=config.device
                        )
                        for _ in range(config.flow_num_repeats)
                    ]
                ),
                persistent=False,
            )

    def forward(
        self,
        batch: dict[str, Tensor],
        reduction: str = "mean",
    ) -> tuple[Tensor, dict]:
        """Compute the enabled flow, text and FAST training losses."""
        if reduction not in {"mean", "none"}:
            raise ValueError(f"Unsupported loss reduction: {reduction!r}")
        text_labels = batch.get("text_labels")
        predict_actions_t = batch.get("predict_actions")

        # Delegate only unannotated batches; PI0.5 ignores recipe action-routing masks.
        if (
            text_labels is None
            and predict_actions_t is None
            and not getattr(self.config, "enable_fast_action_loss", False)
        ):
            return super().forward(batch, reduction=reduction)

        # Compute the host-side action-routing decision once for both flow and FAST.
        predict_any = predict_actions_t is None or bool(predict_actions_t.any().item())
        run_flow = self.config.flow_loss_weight > 0 and predict_any
        run_text = self.config.text_loss_weight > 0 and text_labels is not None

        loss_dict: dict[str, Any] = {}
        total: Tensor | None = None

        # Decide which losses fire this step.
        run_fast = (
            getattr(self.config, "enable_fast_action_loss", False)
            and self.config.fast_action_loss_weight > 0
            and predict_any
        )
        action_tokens = action_mask = action_code_mask = None
        if run_fast:
            action_tokens = batch.get(ACTION_TOKENS)
            action_mask = batch.get(ACTION_TOKEN_MASK)
            action_code_mask = batch.get(ACTION_CODE_TOKEN_MASK)
            if action_tokens is None or action_mask is None or action_code_mask is None:
                missing = [
                    key
                    for key, value in (
                        (ACTION_TOKENS, action_tokens),
                        (ACTION_TOKEN_MASK, action_mask),
                        (ACTION_CODE_TOKEN_MASK, action_code_mask),
                    )
                    if value is None
                ]
                raise ValueError(
                    "FineART-VLA FAST action loss is enabled, but the preprocessor did not produce "
                    f"required batch keys: {missing}."
                )

        # Flow uses one fused prefix/suffix pass; text-only batches skip the suffix.
        if run_flow:
            flow_loss, text_loss, fast_loss = self._compute_all_losses_fused(
                batch,
                text_labels=text_labels if run_text else None,
                action_tokens=action_tokens if run_fast else None,
                action_mask=action_mask if run_fast else None,
                action_code_mask=action_code_mask if run_fast else None,
                predict_actions_t=predict_actions_t,
                reduction=reduction,
            )
            loss_dict["flow_loss"] = flow_loss.detach().mean()
            total = self.config.flow_loss_weight * flow_loss
            if text_loss is not None:
                loss_dict["text_loss"] = text_loss.detach().mean()
                total = total + self.config.text_loss_weight * text_loss
            if fast_loss is not None:
                loss_dict["fast_action_loss"] = fast_loss.detach().mean()
                total = total + self.config.fast_action_loss_weight * fast_loss
        elif run_text or run_fast:
            text_loss, fast_loss = self._compute_text_and_fast_loss(
                batch,
                text_labels=text_labels if run_text else None,
                action_tokens=action_tokens if run_fast else None,
                action_mask=action_mask if run_fast else None,
                action_code_mask=action_code_mask if run_fast else None,
                predict_actions_t=predict_actions_t,
                reduction=reduction,
            )
            if text_loss is not None:
                loss_dict["text_loss"] = text_loss.detach().mean()
                weighted = self.config.text_loss_weight * text_loss
                total = weighted if total is None else total + weighted
            if fast_loss is not None:
                loss_dict["fast_action_loss"] = fast_loss.detach().mean()
                weighted = self.config.fast_action_loss_weight * fast_loss
                total = weighted if total is None else total + weighted

        if total is None:
            # Both flow and text disabled — make this an obvious bug
            # rather than a silent zero loss.
            raise RuntimeError(
                "FineARTVLAPolicy.forward: both flow_loss_weight and "
                "text_loss_weight are 0 (or text_labels missing) — "
                "nothing to train."
            )

        # The shared MetricsTracker accepts Python scalars, not detached tensors.
        # Transfer the small metric vector once so flow/text losses are not dropped.
        loss_dict["loss"] = total.detach().mean()
        values = torch.stack(list(loss_dict.values())).tolist()
        return total, dict(zip(loss_dict, values, strict=True))

    def _embed_supervised_prefix(
        self,
        batch: dict[str, Tensor],
        text_labels: Tensor | None,
        action_tokens: Tensor | None,
        action_mask: Tensor | None,
        *,
        suppress_prefix_grads: bool = False,
    ) -> tuple[Tensor, Tensor, Tensor, int, int]:
        """Embed images, language, and optional FAST supervision once."""
        images, img_masks = self._preprocess_images(batch)
        with torch.no_grad() if suppress_prefix_grads else nullcontext():
            prefix_embs, prefix_pad, prefix_att = self.model.embed_prefix(
                images,
                img_masks,
                batch[OBS_LANGUAGE_TOKENS],
                batch[OBS_LANGUAGE_ATTENTION_MASK],
            )
        non_fast_prefix_len = prefix_embs.shape[1]

        if text_labels is not None:
            lang_start = non_fast_prefix_len - text_labels.shape[1]
            if lang_start >= 0:
                prefix_att = _mark_target_span_causal(
                    prefix_att, text_labels, lang_start, non_fast_prefix_len
                )

        fast_len = 0
        if action_tokens is not None and action_mask is not None:
            fast_emb = self.model.paligemma_with_expert.embed_language_tokens(action_tokens)
            fast_len = action_tokens.shape[1]
            prefix_embs = torch.cat([prefix_embs, fast_emb], dim=1)
            prefix_pad = torch.cat([prefix_pad, action_mask.to(prefix_pad.dtype)], dim=1)
            prefix_att = torch.cat(
                [
                    prefix_att,
                    torch.ones(
                        (action_tokens.shape[0], fast_len),
                        dtype=torch.bool,
                        device=prefix_embs.device,
                    ),
                ],
                dim=1,
            )

        return prefix_embs, prefix_pad, prefix_att, non_fast_prefix_len, fast_len

    def _compute_all_losses_fused(
        self,
        batch: dict[str, Tensor],
        text_labels: Tensor | None,
        action_tokens: Tensor | None,
        action_mask: Tensor | None,
        action_code_mask: Tensor | None,
        predict_actions_t: Tensor | None = None,
        reduction: str = "mean",
    ) -> tuple[Tensor, Tensor | None, Tensor | None]:
        """Compute flow, text and FAST losses from one shared prefix."""
        # ---- preamble (mirrors PI05Pytorch.forward) ------------------
        actions = self.prepare_action(batch)

        # Flow-only KI steps have no live gradient path through the prefix.
        suppress_prefix_grads = (
            text_labels is None
            and action_tokens is None
            and getattr(self.config, "knowledge_insulation", False)
        )

        prefix_embs, prefix_pad, prefix_att, non_fast_prefix_len, fast_len = self._embed_supervised_prefix(
            batch,
            text_labels,
            action_tokens,
            action_mask,
            suppress_prefix_grads=suppress_prefix_grads,
        )

        # Amortized flow reuses one VLM prefix across fresh denoising targets.
        num_repeats = int(getattr(self.config, "flow_num_repeats", 1))
        if num_repeats > 1:
            prefix_out, flow_loss = self._amortized_prefix_and_flow(
                actions,
                prefix_embs,
                prefix_pad,
                prefix_att,
                non_fast_prefix_len,
                fast_len,
                predict_actions_t,
                num_repeats,
                suppress_prefix_grads=suppress_prefix_grads,
                reduction=reduction,
            )
        else:
            prefix_out, flow_loss = self._combined_prefix_and_flow(
                actions,
                prefix_embs,
                prefix_pad,
                prefix_att,
                non_fast_prefix_len,
                fast_len,
                predict_actions_t,
                suppress_prefix_grads=suppress_prefix_grads,
                reduction=reduction,
            )

        text_loss, fast_loss = self._prefix_ce_losses(
            prefix_out,
            text_labels,
            action_tokens,
            action_code_mask,
            fast_len,
            predict_actions_t,
            reduction,
        )
        return flow_loss, text_loss, fast_loss

    def _combined_prefix_and_flow(
        self,
        actions: Tensor,
        prefix_embs: Tensor,
        prefix_pad: Tensor,
        prefix_att: Tensor,
        non_fast_prefix_len: int,
        fast_len: int,
        predict_actions_t: Tensor | None,
        suppress_prefix_grads: bool = False,
        reduction: str = "mean",
    ) -> tuple[Tensor, Tensor]:
        """Run the single-repeat combined prefix and action path."""

        noise = self.model.sample_noise(actions.shape, actions.device)
        time = self.model.sample_time(actions.shape[0], actions.device)
        time_expanded = time[:, None, None]
        x_t = time_expanded * noise + (1 - time_expanded) * actions
        u_t = noise - actions

        # ---- suffix: noisy actions ----------------------------------
        suffix_embs, suffix_pad, suffix_att, adarms_cond = self.model.embed_suffix(x_t, time)

        # ---- bf16 alignment (mirrors PI05Pytorch.forward) -----------
        first_layer = self.model.paligemma_with_expert.paligemma.model.language_model.layers[0]
        if first_layer.self_attn.q_proj.weight.dtype == torch.bfloat16:
            suffix_embs = suffix_embs.to(dtype=torch.bfloat16)
            prefix_embs = prefix_embs.to(dtype=torch.bfloat16)

        pad_masks = torch.cat([prefix_pad, suffix_pad], dim=1)
        att_masks = torch.cat([prefix_att, suffix_att], dim=1)
        att_2d_masks = make_att_2d_masks(pad_masks, att_masks)

        # Block suffix-to-FAST attention to prevent trivial action leakage.
        if fast_len > 0:
            fast_start = non_fast_prefix_len
            fast_end = non_fast_prefix_len + fast_len  # = prefix_pad.shape[1]
            att_2d_masks[:, fast_end:, fast_start:fast_end] = False

        position_ids = torch.cumsum(pad_masks, dim=1) - 1
        if fast_len > 0:
            # Position flow parallel to FAST so its RoPE offsets match inference without FAST.
            non_fast_valid = prefix_pad[:, :non_fast_prefix_len].sum(dim=1, keepdim=True)
            suffix_pos = non_fast_valid + torch.cumsum(suffix_pad, dim=1) - 1
            position_ids = torch.cat([position_ids[:, : prefix_pad.shape[1]], suffix_pos], dim=1)
        att_2d_masks_4d = self.model._prepare_attention_masks_4d(att_2d_masks, dtype=prefix_embs.dtype)

        # ---- forward (capture BOTH expert outputs) ------------------
        ki_kwargs = self._ki_forward_kwargs(suppress_prefix_grads=suppress_prefix_grads)
        (prefix_out, suffix_out), _ = self.model.paligemma_with_expert.forward(
            attention_mask=att_2d_masks_4d,
            position_ids=position_ids,
            past_key_values=None,
            inputs_embeds=[prefix_embs, suffix_embs],
            use_cache=False,
            adarms_cond=[None, adarms_cond],
            **ki_kwargs,
        )

        # ---- flow loss (mirrors PI05Pytorch.forward) ----------------
        suffix_out_slice = suffix_out[:, -self.model.config.chunk_size :].to(dtype=torch.float32)
        v_t = self.model.action_out_proj(suffix_out_slice)
        flow_per_dim = functional.mse_loss(u_t, v_t, reduction="none")
        # Truncate to the actual action dimensionality (PI05 pads
        # internally to max_action_dim).
        if self.config.output_features is None:
            raise ValueError("output_features must be configured before computing actions")
        original_action_dim = self.config.output_features[ACTION].shape[0]
        flow_per_dim = flow_per_dim[:, :, :original_action_dim]
        per_sample_flow = flow_per_dim.mean(dim=(1, 2))
        flow_loss = _reduce_action_loss(per_sample_flow, predict_actions_t, reduction)
        return prefix_out, flow_loss

    def _ki_forward_kwargs(self, suppress_prefix_grads: bool = False, flex_masks=None) -> dict[str, Any]:
        """Joint-layer hooks for PaliGemmaWithExpertModel.forward (see ``compute_layer_complete``)."""
        kwargs: dict[str, Any] = {}
        if suppress_prefix_grads:
            kwargs["suppress_prefix_grads"] = True
        if flex_masks is not None:
            kwargs["attention_fn"] = _flex_attention_fn(flex_masks)
        elif self.config.use_manual_attention:
            kwargs["attention_fn"] = _manual_attention_fn(self.config.manual_attention_scope)
        norm_backend = getattr(self, "_flashrt_adarms", None)
        if norm_backend is not None:
            kwargs["norm_backend"] = norm_backend
        return kwargs

    def _amortized_prefix_and_flow(
        self,
        actions: Tensor,
        prefix_embs: Tensor,
        prefix_pad: Tensor,
        prefix_att: Tensor,
        non_fast_prefix_len: int,
        fast_len: int,
        predict_actions_t: Tensor | None,
        num_repeats: int,
        suppress_prefix_grads: bool = False,
        reduction: str = "mean",
    ) -> tuple[Tensor, Tensor]:
        """Run K independent action draws against one shared VLM prefix."""

        model = self.model
        k = num_repeats
        chunk = self.config.chunk_size
        batch_size, prefix_len = prefix_pad.shape

        first_layer = model.paligemma_with_expert.paligemma.model.language_model.layers[0]
        use_bf16 = first_layer.self_attn.q_proj.weight.dtype == torch.bfloat16
        if use_bf16:
            prefix_embs = prefix_embs.to(dtype=torch.bfloat16)

        # ---- K suffix blocks: independent noise/time draws ----------
        # Embed all independent K draws in one flattened batch.
        noise = model.sample_noise((k * batch_size, *actions.shape[1:]), actions.device)
        time = model.sample_time(k * batch_size, actions.device)
        actions_rep = actions.repeat(k, 1, 1)  # (k*B, chunk, motor_dim)
        time_expanded = time[:, None, None]
        x_t = time_expanded * noise + (1 - time_expanded) * actions_rep
        u_t = (noise - actions_rep).view(k, batch_size, chunk, -1).transpose(0, 1)  # (B, k, chunk, motor)
        s_embs, suffix_pad, suffix_att, adarms = model.embed_suffix(x_t, time)
        if use_bf16:
            s_embs = s_embs.to(dtype=torch.bfloat16)
        suffix_pad = suffix_pad[:batch_size]
        suffix_att = suffix_att[:batch_size]
        suffix_embs = (
            s_embs.view(k, batch_size, chunk, -1).transpose(0, 1).reshape(batch_size, k * chunk, -1)
        )  # (B, k*chunk, D)
        # Broadcast each draw's AdaRMS condition over its action chunk.
        adarms_cond = (
            adarms.view(k, batch_size, 1, adarms.shape[-1])
            .expand(k, batch_size, chunk, adarms.shape[-1])
            .transpose(0, 1)
            .reshape(batch_size, k * chunk, adarms.shape[-1])
        )  # (B, k*chunk, cond_dim)

        # Prefix rows cannot see action blocks; each action block sees only itself and the prefix.
        use_flex = (
            self.config.use_flex_attention
            and getattr(self.config, "knowledge_insulation", False)
            and not getattr(self, "_flex_attention_disabled", False)
            and _get_flex_fns(prefix_pad.device) is not None
        )
        flex_masks = None
        if use_flex:
            try:
                if not hasattr(self, "_flex_mask_builder"):
                    self._flex_mask_builder = _FlexMaskBuilder()
                flex_masks = self._flex_mask_builder.build(
                    prefix_pad, prefix_att, non_fast_prefix_len, k, chunk
                )
            except Exception as exc:
                logger.warning("FineART-VLA: FlexAttention initialization failed (%s); using SDPA.", exc)
                self._flex_attention_disabled = True
        if flex_masks is not None:
            att_2d_4d = None
        else:
            device = prefix_pad.device
            prefix_att_2d = make_att_2d_masks(prefix_pad, prefix_att)  # (B, P, P)
            prefix_rows = torch.cat(
                [
                    prefix_att_2d,
                    torch.zeros(batch_size, prefix_len, k * chunk, dtype=torch.bool, device=device),
                ],
                dim=2,
            )

            action_to_prefix = prefix_pad[:, None, :].expand(batch_size, k * chunk, prefix_len).clone()
            if fast_len > 0:
                action_to_prefix[:, :, non_fast_prefix_len:prefix_len] = False
            action_to_action = self._flow_block_diag[None].expand(batch_size, k * chunk, k * chunk)
            action_rows = torch.cat([action_to_prefix, action_to_action], dim=2)

            att_2d = torch.cat([prefix_rows, action_rows], dim=1)  # (B, P + k*chunk, P + k*chunk)
            att_2d_4d = model._prepare_attention_masks_4d(att_2d, dtype=prefix_embs.dtype)

        # Restart every independent flow block after the non-FAST prefix to match inference RoPE.
        if fast_len > 0:
            prefix_offsets = prefix_pad[:, :non_fast_prefix_len].sum(dim=-1)[:, None]
        else:
            prefix_offsets = torch.sum(prefix_pad, dim=-1)[:, None]
        block_positions = prefix_offsets + torch.cumsum(suffix_pad, dim=1) - 1  # (B, chunk)
        position_ids = torch.cat([torch.cumsum(prefix_pad, dim=1) - 1, block_positions.repeat(1, k)], dim=1)

        ki_kwargs = self._ki_forward_kwargs(suppress_prefix_grads, flex_masks)
        (prefix_out, suffix_out), _ = model.paligemma_with_expert.forward(
            attention_mask=att_2d_4d,
            position_ids=position_ids,
            past_key_values=None,
            inputs_embeds=[prefix_embs, suffix_embs],
            use_cache=False,
            adarms_cond=[None, adarms_cond],
            **ki_kwargs,
        )

        # ---- flow loss averaged over the K blocks -------------------
        # Project all blocks together before averaging their losses.
        if self.config.output_features is None:
            raise ValueError("output_features must be configured before computing actions")
        original_action_dim = self.config.output_features[ACTION].shape[0]
        v_t = model.action_out_proj(suffix_out.to(dtype=torch.float32))
        v_t = v_t.view(batch_size, k, chunk, -1)  # (B, k, chunk, motor)
        flow_per_dim = functional.mse_loss(u_t, v_t, reduction="none")[..., :original_action_dim]
        per_sample_flow = flow_per_dim.mean(dim=(1, 2, 3))
        flow_loss = _reduce_action_loss(per_sample_flow, predict_actions_t, reduction)
        return prefix_out, flow_loss

    def _prefix_ce_losses(
        self,
        prefix_out: Tensor | None,
        text_labels: Tensor | None,
        action_tokens: Tensor | None,
        action_code_mask: Tensor | None,
        fast_len: int,
        predict_actions_t: Tensor | None,
        reduction: str = "mean",
    ) -> tuple[Tensor | None, Tensor | None]:
        """Compute enabled text and FAST losses from the shared prefix output."""
        lm_head = self.model.paligemma_with_expert.paligemma.lm_head

        text_loss: Tensor | None = None
        if text_labels is not None and prefix_out is not None:
            lang_len = text_labels.shape[1]
            if fast_len > 0:
                text_hidden = prefix_out[:, -(fast_len + lang_len) : -fast_len, :]
            else:
                text_hidden = prefix_out[:, -lang_len:, :]
            # Liger avoids materializing the full vocabulary logits tensor.
            text_loss = _shifted_lin_ce(
                text_hidden,
                lm_head.weight,
                text_labels,
                z_loss_weight=getattr(self.config, "text_ce_z_loss_weight", 0.0),
                compiled=self.config.use_compiled_text_ce,
                reduction=reduction,
            )

        fast_loss: Tensor | None = None
        if fast_len > 0 and prefix_out is not None and action_code_mask is not None:
            fast_hidden = prefix_out[:, -fast_len:, :]
            fast_loss = _fast_lin_ce(
                fast_hidden,
                lm_head.weight,
                action_tokens,
                action_code_mask,
                predict_actions_t,
                compiled=self.config.use_compiled_text_ce,
                reduction=reduction,
            )

        return text_loss, fast_loss

    def _compute_text_and_fast_loss(
        self,
        batch: dict[str, Tensor],
        text_labels: Tensor | None,
        action_tokens: Tensor | None,
        action_mask: Tensor | None,
        action_code_mask: Tensor | None,
        predict_actions_t: Tensor | None = None,
        reduction: str = "mean",
    ) -> tuple[Tensor | None, Tensor | None]:
        """Single prefix forward → text CE + FAST CE.

        Embed [images, language] (and FAST when requested) once, run
        one backbone forward, then slice the resulting hidden states
        at the language and FAST positions to compute both CE losses.
        Bit-equivalent to running the two losses in separate forwards
        because the segment-aware ``make_att_2d_masks`` keeps FAST
        tokens invisible to language tokens, so adding FAST to the
        prefix doesn't perturb the hidden states at language positions.

        Returns ``(text_loss, fast_loss)``. Either can be ``None`` if
        the caller doesn't want that head.
        """

        prefix_embs, prefix_pad, prefix_att, _, fast_len = self._embed_supervised_prefix(
            batch,
            text_labels,
            action_tokens,
            action_mask,
        )
        att_2d = make_att_2d_masks(prefix_pad, prefix_att)
        position_ids = torch.cumsum(prefix_pad, dim=1) - 1
        att_2d_4d = self.model._prepare_attention_masks_4d(att_2d, dtype=prefix_embs.dtype)

        (vlm_out, _), _ = self.model.paligemma_with_expert.forward(
            attention_mask=att_2d_4d,
            position_ids=position_ids,
            past_key_values=None,
            inputs_embeds=[prefix_embs, None],
            use_cache=False,
        )
        if vlm_out is None:
            raise RuntimeError("FineART-VLA text+fast loss: VLM forward returned no hidden states.")
        return self._prefix_ce_losses(
            vlm_out,
            text_labels,
            action_tokens,
            action_code_mask,
            fast_len,
            predict_actions_t,
            reduction,
        )

    def supports_text_generation(self) -> bool:
        return self.config.text_loss_weight > 0

    @torch.no_grad()
    def generate_text(self, batch: dict[str, Tensor]) -> str:
        """Generate from the checkpoint recipe rendered by the shared runtime pipeline."""
        if self.config.memory_scratchpad:
            raise ValueError(
                "Combined memory/subtask checkpoints require a scratchpad-aware controller; "
                "the shared autosteer runtime accepts only a subtask response."
            )
        if self._batch_size_from_observation(batch) != 1:
            raise ValueError("FineART-VLA text generation requires one observation.")
        return self.select_message(batch)

    def select_message(
        self,
        batch: dict[str, Tensor],
        *,
        max_new_tokens: int = 128,
        min_new_tokens: int = 0,
        eos_token_id: int | None = None,
        temperature: float = 0.0,
        top_p: float = 1.0,
        tokenizer: Any = None,
        suppress_loc_tokens: bool = False,
        use_kv_cache: bool = True,
    ) -> str:
        """Generate text continuation from a multimodal prefix (used by the runtime CLI).

        ``suppress_loc_tokens=True`` masks PaliGemma's reserved ``<locDDDD>`` ids
        ([256000, 257024)) before sampling — the pretraining prior drifts back to
        them on small text-CE budgets. Pass ``True`` for subtask/memory/plan,
        ``False`` for VQA (spatial answers legitimately emit ``<loc>``).
        """
        self.eval()

        if tokenizer is None:
            tokenizer = _get_loc_tokenizer(
                self.config.text_tokenizer_name, AutoTokenizer, register_paligemma_loc_tokens
            )
        if eos_token_id is None:
            eos_token_id = tokenizer.eos_token_id

        special_ids: set[int] = set()
        try:
            for sid in tokenizer.all_special_ids or []:
                if sid is not None:
                    special_ids.add(int(sid))
        except Exception:  # noqa: BLE001  # nosec B110
            pass
        if eos_token_id is not None:
            special_ids.add(int(eos_token_id))

        images, img_masks = self._preprocess_images(batch)
        tokens = batch[OBS_LANGUAGE_TOKENS]
        masks = batch[OBS_LANGUAGE_ATTENTION_MASK]

        prefix_embs, prefix_pad_masks, prefix_att_masks = self.model.embed_prefix(
            images, img_masks, tokens, masks
        )

        device = prefix_embs.device
        bsize = prefix_embs.shape[0]
        ones_step = torch.ones((bsize, 1), dtype=torch.bool, device=device)

        current_embs = prefix_embs
        current_pad = prefix_pad_masks
        current_att = prefix_att_masks
        generated: list[int] = []
        new_emb = None

        # Cache the image-heavy prefix; disabling the cache retains the full-recompute parity path.
        cache = None

        backbone = self.model.paligemma_with_expert
        lm_head = backbone.paligemma.lm_head

        # Use q_proj's dtype because norms and embeddings may remain fp32 while SDPA queries are bf16.
        backbone_dtype = backbone.paligemma.model.language_model.layers[0].self_attn.q_proj.weight.dtype

        for _ in range(max_new_tokens):
            if cache is None:
                # Run the full bidirectional prefix initially or whenever caching is disabled.
                step_embs = current_embs
                att_2d = make_att_2d_masks(current_pad, current_att)
                position_ids = torch.cumsum(current_pad, dim=1) - 1
                att_2d_4d = self.model._prepare_attention_masks_4d(att_2d, dtype=backbone_dtype)
            else:
                # Incremental decoding feeds only the last token while retaining prefix padding masks.
                step_embs = new_emb
                att_2d = current_pad[:, None, :]
                att_2d_4d = self.model._prepare_attention_masks_4d(att_2d, dtype=backbone_dtype)
                position_ids = (torch.cumsum(current_pad, dim=1) - 1)[:, -1:]
            (vlm_out, _), new_cache = backbone.forward(
                attention_mask=att_2d_4d,
                position_ids=position_ids,
                past_key_values=cache,
                inputs_embeds=[step_embs, None],
                use_cache=use_kv_cache,
            )
            if use_kv_cache:
                cache = new_cache
            if vlm_out is None:
                break
            # Shared runtime tokenization right-pads prompts; generation must not
            # start from a padded position. Cached steps already contain one token.
            last = _last_valid_prefix_hidden(vlm_out, current_pad) if vlm_out.shape[1] > 1 else vlm_out
            last = last.to(lm_head.weight.dtype)
            logits_step = lm_head(last)[:, -1]  # (B, V)
            if special_ids and len(generated) < min_new_tokens:
                for sid in special_ids:
                    logits_step[..., sid] = float("-inf")
            # Suppress FAST-only vocabulary that otherwise leaks into generated text.
            vocab_size = logits_step.shape[-1]
            fast_skip = int(getattr(self.config, "fast_skip_tokens", 128))
            fast_lo = vocab_size - 1 - fast_skip - (_FAST_ACTION_VOCAB_SIZE - 1)
            if 0 < fast_lo < 256000:
                logits_step[..., fast_lo:256000] = float("-inf")
            if suppress_loc_tokens:
                logits_step[..., 256000:257024] = float("-inf")
            next_ids = self._sample_next_token(logits_step, temperature, top_p)
            tok_id = int(next_ids[0].item())
            generated.append(tok_id)
            if eos_token_id is not None and tok_id == eos_token_id:
                break

            # embed_language_tokens already applies the Gemma sqrt(hidden) scale (tf>=5.4.0).
            new_emb = backbone.embed_language_tokens(next_ids.unsqueeze(0))
            # Both paths track valid keys, but only recompute retains full embedding history.
            current_pad = torch.cat([current_pad, ones_step], dim=1)
            if not use_kv_cache:
                current_embs = torch.cat([current_embs, new_emb], dim=1)
                current_att = torch.cat([current_att, ones_step], dim=1)

        decoded = tokenizer.decode(generated, skip_special_tokens=True).strip()
        if not decoded and generated:
            logger.debug("FineART-VLA generated an empty text response; raw token IDs: %s", generated[:16])
        return decoded

    @staticmethod
    def _batch_size_from_observation(batch: dict[str, Any]) -> int:
        state = batch.get("observation.state")
        if isinstance(state, Tensor) and state.ndim > 0:
            return int(state.shape[0])
        for key, value in batch.items():
            if isinstance(key, str) and key.startswith("observation.images.") and torch.is_tensor(value):
                return int(value.shape[0])
        return 1

    @staticmethod
    def _sample_next_token(logits: Tensor, temperature: float, top_p: float) -> Tensor:
        if temperature <= 0.0:
            return logits.argmax(dim=-1)
        scaled = logits / max(temperature, 1e-6)
        probs = torch.softmax(scaled, dim=-1)
        if top_p < 1.0:
            sorted_p, sorted_ix = torch.sort(probs, descending=True, dim=-1)
            cum = torch.cumsum(sorted_p, dim=-1)
            mask = cum > top_p
            mask[..., 0] = False
            sorted_p = sorted_p.masked_fill(mask, 0.0)
            sorted_p = sorted_p / sorted_p.sum(dim=-1, keepdim=True).clamp_min(1e-8)
            choice = torch.multinomial(sorted_p, num_samples=1)
            return sorted_ix.gather(-1, choice).squeeze(-1)
        return torch.multinomial(probs, num_samples=1).squeeze(-1)

    @classmethod
    def _load_as_safetensor(cls, model, model_file: str, map_location: str, strict: bool):
        """Load native FineART-VLA weights, converting PI0.5/OpenPI ones such as ``lerobot/pi05_base``."""
        with safe_open(model_file, framework="pt") as checkpoint:
            # safe_open handles expose .keys() but are not iterable themselves.
            native = all(key.startswith("model.") for key in checkpoint.keys())  # noqa: SIM118
        if native:
            return super()._load_as_safetensor(model, model_file, map_location, strict)
        logger.info("FineART-VLA: converting PI0.5-format checkpoint %s", model_file)
        state_dict = load_file(model_file, device=resolve_safetensors_device(map_location))
        state_dict = model._convert_openpi_state_dict(state_dict)
        missing_keys, unexpected_keys = model.load_state_dict(state_dict, strict=strict)
        log_model_loading_keys(missing_keys, unexpected_keys)
        return model

    def _prepare_pretrained_state_dict(self, remapped_state_dict: dict[str, Tensor]) -> dict[str, Tensor]:
        remapped_state_dict = super()._prepare_pretrained_state_dict(remapped_state_dict)
        # The text loss trains the language head; start it from the embeddings when absent.
        lm_head_key = "model.paligemma_with_expert.paligemma.lm_head.weight"
        embed_tokens_key = "model.paligemma_with_expert.paligemma.model.language_model.embed_tokens.weight"
        if lm_head_key not in remapped_state_dict and embed_tokens_key in remapped_state_dict:
            remapped_state_dict[lm_head_key] = remapped_state_dict[embed_tokens_key].clone()
        return remapped_state_dict

    def get_optim_params(self):
        """Return policy parameters, optionally split into LR-scaled groups.

        Three orthogonal multipliers scale the base ``optimizer_lr``:
        ``lm_head_lr_scale`` (PaliGemma ``lm_head`` + tied ``embed_tokens``),
        ``backbone_lr_scale`` (the rest of the PaliGemma tower), and
        ``action_expert_lr_scale`` (the Gemma expert + action/time projection
        heads). The cosine scheduler multiplies every group by the same lambda
        each step so the ratios are preserved across decay. AdamW backend
        options stay in these policy-local parameter groups. Optional conditioning
        controls split time MLPs and adaptive normalization projections from the
        expert; their LR is additionally multiplied by ``conditioning_lr_scale``.
        """
        head_scale = float(getattr(self.config, "lm_head_lr_scale", 1.0))
        backbone_scale = float(getattr(self.config, "backbone_lr_scale", 1.0))
        expert_scale = float(getattr(self.config, "action_expert_lr_scale", 1.0))
        conditioning_scale = self.config.conditioning_lr_scale
        split_conditioning = conditioning_scale != 1.0
        backend = {
            "foreach": getattr(self.config, "optimizer_foreach", False),
            "fused": getattr(self.config, "optimizer_fused", True),
        }
        if head_scale == 1.0 and backbone_scale == 1.0 and expert_scale == 1.0 and not split_conditioning:
            return [{"params": self.parameters(), **backend}]

        # Keep the tied LM projection and embeddings in the same optimizer group.
        head_substrings = (
            "paligemma_with_expert.paligemma.lm_head.",
            "paligemma_with_expert.paligemma.model.language_model.embed_tokens.",
        )
        backbone_substring = "paligemma_with_expert.paligemma."
        head_params: list[torch.nn.Parameter] = []
        backbone_params: list[torch.nn.Parameter] = []
        expert_params: list[torch.nn.Parameter] = []
        conditioning_params: list[torch.nn.Parameter] = []
        for name, p in self.named_parameters():
            if not p.requires_grad:
                continue
            if any(s in name for s in head_substrings):
                head_params.append(p)
            elif backbone_substring in name:
                backbone_params.append(p)
            elif split_conditioning and (
                name.startswith(("model.time_mlp_in.", "model.time_mlp_out."))
                or (
                    "paligemma_with_expert.gemma_expert." in name
                    and any(
                        part in name
                        for part in (
                            ".input_layernorm.dense.",
                            ".post_attention_layernorm.dense.",
                            ".norm.dense.",
                        )
                    )
                )
            ):
                conditioning_params.append(p)
            else:
                expert_params.append(p)
        base_lr = float(self.config.optimizer_lr)
        groups: list[dict[str, object]] = []
        if backbone_params:
            groups.append({"params": backbone_params, "lr": base_lr * backbone_scale, "name": "backbone"})
        if expert_params:
            groups.append({"params": expert_params, "lr": base_lr * expert_scale, "name": "action_expert"})
        if conditioning_params:
            groups.append(
                {
                    "params": conditioning_params,
                    "lr": base_lr * expert_scale * conditioning_scale,
                    "name": "conditioning",
                }
            )
        elif split_conditioning:
            raise RuntimeError("Conditioning stabilization requested but no conditioning parameters matched")
        if head_params:
            groups.append({"params": head_params, "lr": base_lr * head_scale, "name": "lm_head"})
        # Sanity: a non-trivial head scale that matches no params would silently
        # do nothing — surface that fast.
        if head_scale != 1.0 and not head_params:
            raise RuntimeError(
                "lm_head_lr_scale != 1.0 but no parameters matched the LM-head "
                f"name patterns: {head_substrings!r}. Did the underlying PaliGemma "
                "module rename?"
            )
        logging.info(
            "FineARTVLAPolicy LR groups (base=%.3g): backbone=%.3g (×%.3g, n=%d), "
            "action_expert=%.3g (×%.3g, n=%d), lm_head=%.3g (×%.3g, n=%d)",
            base_lr,
            base_lr * backbone_scale,
            backbone_scale,
            len(backbone_params),
            base_lr * expert_scale,
            expert_scale,
            len(expert_params),
            base_lr * head_scale,
            head_scale,
            len(head_params),
        )
        for group in groups:
            group.update(backend)
        return groups

    @torch.no_grad()
    def predict_action_chunk(
        self, batch: dict[str, Tensor], **kwargs: Unpack[RTCActionSelectKwargs]
    ) -> Tensor:
        marks = batch.get(OBS_LANGUAGE_CAUSAL_MARKS)
        if marks is None:
            return super().predict_action_chunk(batch, **kwargs)
        return self._predict_action_chunk_with_marks(batch, marks, **kwargs)

    @torch.no_grad()
    def _predict_action_chunk_with_marks(
        self, batch: dict[str, Tensor], marks: Tensor, **kwargs: Unpack[RTCActionSelectKwargs]
    ) -> Tensor:
        """Base ``predict_action_chunk`` plus causal marks on the generated-subtask span."""
        self.eval()
        images, img_masks = self._preprocess_images(batch)
        tokens = batch[OBS_LANGUAGE_TOKENS]
        masks = batch[OBS_LANGUAGE_ATTENTION_MASK]
        actions = self.model.sample_actions(
            images, img_masks, tokens, masks, lang_causal_marks=marks, **kwargs
        )
        if self.config.output_features is None:
            raise ValueError("output_features must be configured before computing actions")
        original_action_dim = self.config.output_features[ACTION].shape[0]
        return actions[:, :, :original_action_dim]
