"""Frozen Theia features plus a small trainable temporal adapter.

The released Theia checkpoint is RGB-pretrained.  This module deliberately
does not claim depth pretraining: simulator depth is mapped to a uint8
grayscale image, repeated over RGB channels, and passed through the pinned
Theia backbone.  The backbone itself remains frozen.
"""

from __future__ import annotations

from pathlib import Path
import hashlib
import json
import torch
from torch import nn
from torch.nn import functional as F

from . import THEIA_REPOSITORY, THEIA_REVISION


def depth_to_uint8(depth: torch.Tensor, near: float, far: float) -> torch.Tensor:
    """Map optical-Z depth ``[B,T,H,W]`` to Theia's uint8 image contract."""
    depth = depth.float()
    valid = torch.isfinite(depth) & (depth >= near) & (depth <= far)
    # 255 is near, 1 is far; zero is reserved for invalid/out-of-range.
    value = 1.0 + 254.0 * (far - depth) / max(far - near, 1e-6)
    value = value.clamp(1.0, 255.0).round()
    return torch.where(valid, value, torch.zeros_like(value)).to(torch.uint8)


class FrozenTheia(nn.Module):
    """Load the pinned HF Theia checkpoint and expose raw 197 DeiT tokens."""

    def __init__(self, repository=THEIA_REPOSITORY, revision=THEIA_REVISION,
                 device=None):
        super().__init__()
        try:
            from transformers import AutoModel
        except ImportError as exc:  # pragma: no cover - installation diagnostic
            raise RuntimeError("Install transformers before loading Theia") from exc
        kwargs = dict(trust_remote_code=True, revision=revision)
        # The repo bundle includes both pinned snapshots. The upstream Theia
        # constructor otherwise fetches its DeiT config/processor from a
        # moving Hub main branch, even when Theia itself is revision-pinned.
        weights = Path(__file__).resolve().parents[2] / "weights"
        snapshot = weights / "theia-pretrained"
        if repository == THEIA_REPOSITORY and revision == THEIA_REVISION and snapshot.is_dir():
            sources = json.loads((weights / "pretrained-sources.json").read_text())
            for name, expected in sources["files"].items():
                if hashlib.sha256((weights / name).read_bytes()).hexdigest() != expected:
                    raise ValueError(f"Pretrained snapshot checksum mismatch: {name}")
            from transformers import AutoConfig
            local_cfg = AutoConfig.from_pretrained(str(snapshot), trust_remote_code=True)
            local_cfg.backbone = str(weights / "deit-tiny-patch16-224")
            self.model = AutoModel.from_pretrained(str(snapshot), config=local_cfg, trust_remote_code=True)
        else:
            self.model = AutoModel.from_pretrained(repository, **kwargs)
        # Theia's public model wraps the DeiT backbone.  The wrapper's forward
        # returns the unmodified [CLS + 196 patch, 192] sequence.
        if not hasattr(self.model, "backbone"):
            raise RuntimeError("Pinned Theia model has no .backbone attribute")
        self.backbone = self.model.backbone
        self.model.eval().requires_grad_(False)
        if device is not None:
            self.model.to(device)
        self.repository = repository
        self.revision = revision
        self.hidden_dim = int(getattr(getattr(self.backbone, "model", None),
                                     "config", getattr(self.model, "config", None)).hidden_size)
        if self.hidden_dim != 192:
            raise RuntimeError(f"Expected Theia-Tiny hidden size 192, got {self.hidden_dim}")

    @property
    def parameter_count(self):
        return sum(p.numel() for p in self.model.parameters())

    def forward(self, images: torch.Tensor) -> torch.Tensor:
        # The remote DeiT wrapper accepts uint8 NCHW/HWC images and applies its
        # pinned resize/rescale/normalization processor.  Keep this call under
        # inference_mode: only the temporal adapter receives PPO gradients.
        with torch.inference_mode():
            result = self.backbone(images)
            if hasattr(result, "last_hidden_state"):
                result = result.last_hidden_state
            elif isinstance(result, dict) and "last_hidden_state" in result:
                result = result["last_hidden_state"]
            if result.ndim != 3 or result.shape[1:] != (197, 192):
                raise RuntimeError(f"Unexpected Theia token shape {tuple(result.shape)}")
            return result


def tokenize(theia: FrozenTheia, depth: torch.Tensor, cfg) -> torch.Tensor:
    """Frozen Theia tokens for depth frames ``[N,H,W]`` -> ``[N,197,192]``.

    The single tokenization path shared by the actor's fallback and the
    environment's per-frame cache, so both produce bit-identical features.
    """
    gray = depth_to_uint8(depth, cfg.depth_near, cfg.depth_far)
    x = F.interpolate(gray[:, None].float(), size=(224, 224), mode="bilinear", align_corners=False)
    images = x.round().clamp(0, 255).to(torch.uint8).repeat(1, 3, 1, 1)
    chunks = [theia(images[start:start + cfg.theia_batch_size])
              for start in range(0, images.shape[0], cfg.theia_batch_size)]
    tokens = torch.cat(chunks, dim=0) if chunks else images.new_zeros((0, 197, 192), dtype=torch.float32)
    return tokens.half() if cfg.feature_dtype == "float16" else tokens


class TheiaPerception(nn.Module):
    """Patch attention + temporal MLP producing the 128-dim SONIC condition.

    ``forward`` accepts either raw depth history ``[B,T,H,W]`` or precomputed
    frozen tokens ``[B,T,197,192]`` from the environment's cache.  Theia has
    no trainable parameters and runs under inference_mode, so the two paths
    are equivalent; the cache only removes redundant recomputation.
    """

    def __init__(self, cfg):
        super().__init__()
        self.cfg = cfg
        self.theia = FrozenTheia(cfg.theia_repository, cfg.theia_revision)
        self.token_norm = nn.LayerNorm(192) if cfg.theia_token_norm else nn.Identity()
        self.patch_score = nn.Linear(192, 1)
        # A small learned time code preserves ordering without a recurrent state.
        self.time_embedding = nn.Parameter(torch.zeros(cfg.depth_history, 192))
        nn.init.normal_(self.time_embedding, std=0.01)
        self.temporal = nn.Sequential(
            nn.LayerNorm(cfg.depth_history * 192),
            nn.Linear(cfg.depth_history * 192, 192),
            nn.SiLU(),
            nn.Linear(192, 128),
            nn.LayerNorm(128),
        )

    def raw_tokens(self, depth):
        b, t = depth.shape[:2]
        return tokenize(self.theia, depth.flatten(0, 1), self.cfg).reshape(b, t, 197, 192)

    def forward(self, depth=None, *, tokens=None):
        if tokens is None:
            tokens = self.raw_tokens(depth)
        tokens = tokens.float()
        tokens = self.token_norm(tokens)
        tokens = tokens + self.time_embedding[None, :, None, :]
        weights = torch.softmax(self.patch_score(tokens).squeeze(-1), dim=-1)
        self.last_attention = weights.detach()  # [B,T,197], display only
        frame = (weights[..., None] * tokens).sum(dim=2)
        return self.temporal(frame.reshape(frame.shape[0], -1))

    def trainable_parameters(self):
        return [p for p in self.parameters() if p.requires_grad]
