"""Temporal Transformer head for Fail, A, and S grade logits."""

from copy import deepcopy
from math import log
from typing import Final, final, override

import torch
from torch import nn
from torch.nn import functional

from .magnitude import MagnitudeNormalizer, MagnitudeResidual
from .model_inputs import (
    FEATURE_WIDTH,
    INPUT_RANK,
    TemporalGradeInputError,
    validate_temporal_inputs,
)
from .temporal_local import TemporalLocalResidual

__all__ = [
    "GRADE_ORDER",
    "INPUT_RANK",
    "TemporalGradeHead",
    "TemporalGradeInputError",
]

MODEL_WIDTH: Final = 256
GRADE_COUNT: Final = 3
MUSIC_TYPE_INDEX: Final = 0
GRADE_TYPE_INDEX: Final = 1
GRADE_ORDER: Final[tuple[str, str, str]] = ("Fail", "A", "S")
SHALLOW_TRANSFORMER_DEPTH: Final = 1
DEEP_TRANSFORMER_DEPTH: Final = 2
THIRD_TRANSFORMER_DEPTH: Final = 3


def sinusoidal_time_encoding(
    center_times: torch.Tensor, model_width: int
) -> torch.Tensor:
    """Encode absolute centers with the original sinusoidal operations."""
    exponents = (
        torch.arange(
            0,
            model_width,
            2,
            dtype=center_times.dtype,
            device=center_times.device,
        )
        / model_width
    )
    frequencies = torch.exp(-log(10_000.0) * exponents)
    angles = center_times.unsqueeze(-1) * frequencies
    return torch.stack((angles.sin(), angles.cos()), dim=-1).flatten(-2)


@final
class TemporalGradeHead(nn.Module):
    """Score all temporal feature tokens with three learned grade queries."""

    grade_tokens: nn.Parameter

    @override
    def __init__(
        self,
        model_width: int = MODEL_WIDTH,
        *,
        transformer_dropout: float = 0.1,
        transformer_depth: int = 2,
        use_time_encoding: bool = True,
        separate_conditional_scorer: bool = False,
        detach_conditional_features: bool = False,
        local_temporal_residual: bool = False,
        magnitude_normalizer: MagnitudeNormalizer | None = None,
        feature_width: int = FEATURE_WIDTH,
    ) -> None:
        """Initialize the projection, temporal encoder, and shared scorer."""
        super().__init__()
        if magnitude_normalizer is not None and (
            local_temporal_residual
            or separate_conditional_scorer
            or detach_conditional_features
        ):
            raise TemporalGradeInputError(
                detail="magnitude residual requires shared scorer and no local residual"
            )
        if local_temporal_residual and (
            separate_conditional_scorer or detach_conditional_features
        ):
            raise TemporalGradeInputError(
                detail="local temporal residual requires shared scorer"
            )
        if type(detach_conditional_features) is not bool:
            raise TemporalGradeInputError(
                detail="detach_conditional_features must be boolean"
            )
        if detach_conditional_features and separate_conditional_scorer is not True:
            raise TemporalGradeInputError(
                detail=(
                    "detach_conditional_features requires separate_conditional_scorer"
                )
            )
        self.detach_conditional_features = detach_conditional_features
        if type(transformer_depth) is not int or transformer_depth not in (
            SHALLOW_TRANSFORMER_DEPTH,
            DEEP_TRANSFORMER_DEPTH,
            THIRD_TRANSFORMER_DEPTH,
        ):
            raise TemporalGradeInputError(
                detail="transformer_depth must be integer 1, 2, or 3"
            )
        self.model_width = model_width
        self.feature_width = feature_width
        self.transformer_depth = transformer_depth
        self.use_time_encoding = use_time_encoding
        self.feature_projection: nn.Linear = nn.Linear(feature_width, model_width)
        self.projection_norm: nn.LayerNorm = nn.LayerNorm(model_width)
        self.ratio_embedding: nn.Linear = nn.Linear(1, model_width)
        self.grade_tokens = nn.Parameter(torch.empty(GRADE_COUNT, model_width))
        self.type_embedding: nn.Embedding = nn.Embedding(2, model_width)
        self.first_transformer_layer: nn.TransformerEncoderLayer = (
            nn.TransformerEncoderLayer(
                d_model=model_width,
                nhead=4,
                dim_feedforward=model_width * 4,
                dropout=transformer_dropout,
                activation="gelu",
                batch_first=True,
                norm_first=True,
            )
        )
        self.second_transformer_layer: nn.TransformerEncoderLayer | None = (
            nn.TransformerEncoderLayer(
                d_model=model_width,
                nhead=4,
                dim_feedforward=model_width * 4,
                dropout=transformer_dropout,
                activation="gelu",
                batch_first=True,
                norm_first=True,
            )
            if transformer_depth >= DEEP_TRANSFORMER_DEPTH
            else None
        )
        self.third_transformer_layer: nn.TransformerEncoderLayer | None = (
            nn.TransformerEncoderLayer(
                d_model=model_width,
                nhead=4,
                dim_feedforward=model_width * 4,
                dropout=transformer_dropout,
                activation="gelu",
                batch_first=True,
                norm_first=True,
            )
            if transformer_depth == THIRD_TRANSFORMER_DEPTH
            else None
        )
        self.scorer: nn.Sequential = nn.Sequential(
            nn.LayerNorm(model_width),
            nn.Linear(model_width, model_width),
            nn.GELU(),
            nn.Linear(model_width, 1, bias=False),
        )
        _ = nn.init.normal_(self.grade_tokens, mean=0.0, std=0.02)
        self.conditional_scorer: nn.Sequential | None = (
            deepcopy(self.scorer) if separate_conditional_scorer else None
        )
        self.local_temporal_residual: TemporalLocalResidual | None = (
            TemporalLocalResidual(model_width) if local_temporal_residual else None
        )
        self.magnitude_residual: MagnitudeResidual | None = (
            MagnitudeResidual(model_width, magnitude_normalizer)
            if magnitude_normalizer is not None
            else None
        )

    @property
    def transformer_layers(
        self,
    ) -> tuple[nn.TransformerEncoderLayer, ...]:
        """Return the temporal encoder layers in execution order."""
        if self.second_transformer_layer is None:
            return (self.first_transformer_layer,)
        if self.third_transformer_layer is None:
            return self.first_transformer_layer, self.second_transformer_layer
        return (
            self.first_transformer_layer,
            self.second_transformer_layer,
            self.third_transformer_layer,
        )

    @override
    def forward(
        self,
        features: torch.Tensor,
        center_times: torch.Tensor,
        valid_ratios: torch.Tensor,
        padding_mask: torch.Tensor,
    ) -> torch.Tensor:
        """Return three raw query scores, interpreted by the model family.

        Optional detachment blocks only direct conditional-loss encoder gradients.
        Global gradient clipping can still couple the subsequent optimizer updates.
        """
        self._validate_inputs(
            features,
            center_times,
            valid_ratios,
            padding_mask,
            self.feature_width,
        )
        batch_size = features.shape[0]
        normalized: torch.Tensor = functional.normalize(features, p=2.0, dim=-1)
        music_tokens: torch.Tensor = self.projection_norm(
            self.feature_projection(normalized),
        )
        if self.magnitude_residual is not None:
            music_tokens = self.magnitude_residual(features, music_tokens, padding_mask)
        if self.local_temporal_residual is not None:
            music_tokens = self.local_temporal_residual(music_tokens, padding_mask)
        if self.use_time_encoding:
            music_tokens = music_tokens + self._sinusoidal_time_encoding(
                center_times, self.model_width
            )
        music_tokens = (
            music_tokens
            + self.ratio_embedding(valid_ratios.unsqueeze(-1))
            + self.type_embedding.weight[MUSIC_TYPE_INDEX]
        )
        grade_tokens = self.grade_tokens.unsqueeze(0).expand(batch_size, -1, -1)
        grade_tokens = grade_tokens + self.type_embedding.weight[GRADE_TYPE_INDEX]
        sequence: torch.Tensor = torch.cat((grade_tokens, music_tokens), dim=1)
        grade_mask = torch.zeros(
            (batch_size, GRADE_COUNT),
            dtype=torch.bool,
            device=padding_mask.device,
        )
        key_padding_mask = torch.cat((grade_mask, padding_mask), dim=1)
        for layer in self.transformer_layers:
            sequence = layer(
                sequence,
                src_key_padding_mask=key_padding_mask,
            )
        logits: torch.Tensor = self.scorer(sequence[:, :GRADE_COUNT])
        if self.conditional_scorer is not None:
            conditional_features = sequence[:, 2:3]
            if self.detach_conditional_features:
                conditional_features = conditional_features.detach()
            conditional: torch.Tensor = self.conditional_scorer(conditional_features)
            logits = torch.cat((logits[:, :2], conditional), dim=1)
        return logits.squeeze(-1)

    _sinusoidal_time_encoding = staticmethod(sinusoidal_time_encoding)
    _validate_inputs = staticmethod(validate_temporal_inputs)
