"""Binary-first Pass/Fail gate and conditional A/S probability contracts."""

import math
from enum import StrEnum
from typing import Final

import torch
from torch.nn import functional

from .losses import GradeLossInputError

S_INDEX: Final = 2


class ChoiceModelFamily(StrEnum):
    """Select output semantics without changing the temporal scorer architecture."""

    LEGACY = "raina_laya_choice_siglip_v4"
    HIERARCHICAL = "raina_laya_hierarchical_v5"


def hierarchical_log_probs(logits: torch.Tensor) -> torch.Tensor:
    """Factor final Fail/A/S probabilities from raw [Fail, Pass, S|Pass] scores."""
    gate = functional.log_softmax(logits[:, :2], dim=1)
    h = logits[:, 2]
    return torch.stack(
        (
            gate[:, 0],
            gate[:, 1] + functional.logsigmoid(-h),
            gate[:, 1] + functional.logsigmoid(h),
        ),
        dim=1,
    )


def hierarchical_predictions(logits: torch.Tensor) -> torch.Tensor:
    """Resolve gate ties to Fail and conditional ties to A, never joint argmax."""
    return torch.where(
        logits[:, 1] > logits[:, 0],
        1 + (logits[:, 2] > 0).long(),
        0,
    )


def hierarchical_gate_margin(logits: torch.Tensor) -> torch.Tensor:
    """Return Pass-minus-Fail gate scores: above zero is Pass, ties stay Fail."""
    return logits[:, 1] - logits[:, 0]


def mix_gate_targets(
    targets: torch.Tensor, teacher_margins: torch.Tensor, mix: float
) -> torch.Tensor:
    """Mix the hard Pass indicator with the teacher gate's Pass probability per song."""
    return (1 - mix) * (targets != 0).to(teacher_margins.dtype) + (
        mix * teacher_margins.sigmoid()
    )


def hierarchical_loss(
    logits: torch.Tensor,
    targets: torch.Tensor,
    *,
    fail_weight: float = 1.0,
    conditional_weight: float = 1.0,
    gate_target: torch.Tensor | None = None,
) -> torch.Tensor:
    """Return per-song gate CE plus actual-Pass-masked conditional BCE.

    `gate_target` replaces only the gate's hard Pass indicator with a per-song soft
    Pass probability; Fail weights and the conditional term still follow `targets`.
    """
    actual_pass = targets != 0
    gate = functional.cross_entropy(
        logits[:, :2],
        actual_pass.long()
        if gate_target is None
        else torch.stack((1 - gate_target, gate_target), dim=1).to(logits.dtype),
        reduction="none",
    )
    gate = torch.where(actual_pass, gate, gate * fail_weight)
    conditional = functional.binary_cross_entropy_with_logits(
        logits[:, 2],
        (targets == S_INDEX).to(logits.dtype),
        reduction="none",
    )
    return gate + conditional_weight * conditional * actual_pass.to(logits.dtype)


def hierarchical_biases(
    loss_mass: tuple[float, float, float],
    fail_weight: float = 1.0,
) -> tuple[float, float, float]:
    """Initialize the gate and conditional priors from final weighted song mass."""
    if any(not math.isfinite(mass) or mass <= 0 for mass in loss_mass):
        raise GradeLossInputError(
            detail="hierarchical priors require all three classes"
        )
    fail, a, s = loss_mass
    return math.log(fail * fail_weight), math.log(a + s), math.log(s / a)
