"""Sigmoid-choice head with loss-prior bias and near-zero initial scorer."""

from typing import override

import torch
from torch import nn

from .magnitude import MagnitudeNormalizer
from .model import MODEL_WIDTH, TemporalGradeHead, TemporalGradeInputError
from .model_inputs import FEATURE_WIDTH


class TemporalChoiceHead(nn.Module):
    """Apply learned per-grade biases after the shared temporal scorer."""

    class_bias: nn.Parameter

    @override
    def __init__(
        self,
        biases: tuple[float, float, float],
        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 temporal core and three class biases."""
        super().__init__()
        self.temporal_head = TemporalGradeHead(
            model_width,
            transformer_dropout=transformer_dropout,
            transformer_depth=transformer_depth,
            use_time_encoding=use_time_encoding,
            separate_conditional_scorer=separate_conditional_scorer,
            detach_conditional_features=detach_conditional_features,
            local_temporal_residual=local_temporal_residual,
            magnitude_normalizer=magnitude_normalizer,
            feature_width=feature_width,
        )
        self.class_bias = nn.Parameter(torch.tensor(biases, dtype=torch.float32))
        output = self.temporal_head.scorer[-1]
        if not isinstance(output, nn.Linear):
            raise TemporalGradeInputError(
                detail="temporal scorer output must be linear"
            )
        nn.init.normal_(output.weight, mean=0.0, std=0.001)
        if self.temporal_head.conditional_scorer is not None:
            self.temporal_head.conditional_scorer.load_state_dict(
                self.temporal_head.scorer.state_dict()
            )

    @property
    def transformer_layers(
        self,
    ) -> tuple[nn.TransformerEncoderLayer, ...]:
        """Expose the temporal head's distinct encoder layers."""
        return self.temporal_head.transformer_layers

    def forward(
        self,
        features: torch.Tensor,
        center_times: torch.Tensor,
        valid_ratios: torch.Tensor,
        padding_mask: torch.Tensor,
    ) -> torch.Tensor:
        """Return uncentered logits with the trainable loss-prior bias."""
        return (
            self.temporal_head(features, center_times, valid_ratios, padding_mask)
            + self.class_bias
        )
