"""Original reduced BERT pretraining example with MLM and NSP heads.

One post-LN bidirectional Transformer, width 8, two heads, vocabulary 16,
fixed length 6. MLM decoder weights are tied to token embeddings. Output packs
96 token logits followed by 2 sentence-pair logits for the portable CPU recipe.
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F


class BidirectionalAttention(nn.Module):
    def __init__(self):
        super().__init__()
        self.query = nn.Linear(8, 8)
        self.key = nn.Linear(8, 8)
        self.value = nn.Linear(8, 8)
        self.softmax = nn.Softmax(dim=-1)
        self.project = nn.Linear(8, 8)

    def forward(self, x):
        b, t, _ = x.shape
        q = self.query(x).reshape(b, t, 2, 4).transpose(1, 2)
        k = self.key(x).reshape(b, t, 2, 4).transpose(1, 2)
        v = self.value(x).reshape(b, t, 2, 4).transpose(1, 2)
        weights = self.softmax(q @ k.transpose(-2, -1) / 2)
        return self.project((weights @ v).transpose(1, 2).reshape(b, t, 8))


class BertBlock(nn.Module):
    def __init__(self):
        super().__init__()
        self.attention = BidirectionalAttention()
        self.attention_norm = nn.LayerNorm(8)
        self.up = nn.Linear(8, 32)
        self.gelu = nn.GELU()
        self.down = nn.Linear(32, 8)
        self.output_norm = nn.LayerNorm(8)

    def forward(self, x):
        hidden = self.attention_norm(x + self.attention(x))
        return self.output_norm(hidden + self.down(self.gelu(self.up(hidden))))


class TinyBertPretraining(nn.Module):
    def __init__(self):
        super().__init__()
        self.word_embeddings = nn.Embedding(16, 8)
        self.position_embeddings = nn.Embedding(6, 8)
        self.segment_embeddings = nn.Embedding(2, 8)
        self.embedding_norm = nn.LayerNorm(8)
        self.encoder = BertBlock()
        self.mlm_transform = nn.Linear(8, 8)
        self.mlm_gelu = nn.GELU()
        self.mlm_norm = nn.LayerNorm(8)
        self.mlm_decoder = nn.Linear(8, 16)
        self.mlm_decoder.weight = self.word_embeddings.weight
        self.pooler = nn.Linear(8, 8)
        self.pooler_tanh = nn.Tanh()
        self.nsp_head = nn.Linear(8, 2)

    def forward(self, tokens, segments):
        positions = torch.arange(tokens.shape[1], device=tokens.device)
        x = (
            self.word_embeddings(tokens)
            + self.position_embeddings(positions)
            + self.segment_embeddings(segments)
        )
        hidden = self.encoder(self.embedding_norm(x))
        mlm = self.mlm_decoder(self.mlm_norm(self.mlm_gelu(self.mlm_transform(hidden))))
        nsp = self.nsp_head(self.pooler_tanh(self.pooler(hidden[:, 0])))
        return torch.cat([mlm.flatten(1), nsp], dim=1)


def unpack(packed):
    return packed[:, :96].reshape(-1, 6, 16), packed[:, 96:]


def masked_loss(logits, positions, original):
    labels = torch.full_like(original, -100)
    labels[:, positions] = original[:, positions]
    return F.cross_entropy(logits.reshape(-1, 16), labels.flatten())


def verify():
    torch.manual_seed(23)
    model = TinyBertPretraining().double().eval()
    tokens = torch.tensor([[1, 4, 5, 2, 6, 2]])
    segments = torch.tensor([[0, 0, 0, 0, 1, 1]])
    with torch.no_grad():
        first, nsp = unpack(model(tokens, segments))
        assert first.shape == (1, 6, 16) and nsp.shape == (1, 2)
        changed = tokens.clone()
        changed[0, 4] = 7
        second, _ = unpack(model(changed, segments))
        assert not torch.allclose(first[:, 1], second[:, 1])
        loss = masked_loss(first, [2], tokens)
        perturb = first.clone()
        perturb[:, [0, 1, 3, 4, 5]] += torch.randn(1, 5, 16, dtype=torch.float64) * 100
        torch.testing.assert_close(
            loss, masked_loss(perturb, [2], tokens), rtol=0, atol=0
        )
        assert (
            model.mlm_decoder.weight.data_ptr()
            == model.word_embeddings.weight.data_ptr()
        )
        probabilities = nsp.softmax(-1)
        torch.testing.assert_close(
            probabilities.sum(-1), torch.ones(1, dtype=torch.float64)
        )
    loss = masked_loss(unpack(model(tokens, segments))[0], [2], tokens)
    loss.backward()
    assert model.word_embeddings.weight.grad.norm() > 0
    return [
        "A later token can change earlier token logits through bidirectional attention",
        "MLM cross entropy ignores all unselected token positions",
        "The MLM decoder shares the token-embedding weight and receives a learning signal",
        "Packed MLM/NSP logits have the declared shapes and sentence-pair probabilities sum to one",
    ]


def experiment():
    torch.manual_seed(23)
    model = TinyBertPretraining().double().eval()
    original = torch.tensor([[1, 4, 5, 2, 6, 2]])
    segments = torch.tensor([[0, 0, 0, 0, 1, 1]])
    cases = []
    with torch.no_grad():
        for identity, label, selected, nsp_label in [
            ("mask-a", "Mask one token in segment A", [2], 0),
            ("mask-b", "Mask one token in segment B", [4], 0),
            ("both", "Mask a token in each segment", [2, 4], 0),
            ("not-next", "Same tokens with a negative sentence-pair label", [2], 1),
        ]:
            tokens = original.clone()
            tokens[:, selected] = 3
            logits, nsp = unpack(model(tokens, segments))
            word = model.word_embeddings(tokens)
            positions = model.position_embeddings(torch.arange(6))
            seg = model.segment_embeddings(segments)
            hidden = model.embedding_norm(word + positions + seg)
            attention = model.encoder.attention
            q = attention.query(hidden).reshape(1, 6, 2, 4).transpose(1, 2)
            k = attention.key(hidden).reshape(1, 6, 2, 4).transpose(1, 2)
            weights = (q @ k.transpose(-2, -1) / 2).softmax(-1)
            mlm_loss = masked_loss(logits, selected, original)
            nsp_loss = F.cross_entropy(nsp, torch.tensor([nsp_label]))
            cases.append(
                {
                    "id": identity,
                    "label": label,
                    "target": "mlm_decoder",
                    "note": "Toy token IDs: CLS=1, SEP=2 and MASK=3. Segments are A,A,A,A,B,B. Only selected positions contribute to MLM loss. This controlled example always replaces selected words with MASK; original BERT selected 15% of tokens with an 80/10/10 mask/random/unchanged rule. The NSP labels here are illustrative, without a real sentence corpus.",
                    "matrices": [
                        {
                            "label": "Bidirectional attention · head 0",
                            "values": weights[0, 0].tolist(),
                        }
                    ],
                    "vectors": [
                        {"label": "Original token IDs", "values": original[0].tolist()},
                        {"label": "Corrupted token IDs", "values": tokens[0].tolist()},
                        {
                            "label": "MLM loss positions · 1 means selected",
                            "values": [int(i in selected) for i in range(6)],
                        },
                        {
                            "label": "Prediction at first masked position · 16 token probabilities",
                            "values": logits[0, selected[0]].softmax(-1).tolist(),
                        },
                        {
                            "label": "Sentence pair · IsNext / NotNext probabilities",
                            "values": nsp[0].softmax(-1).tolist(),
                        },
                    ],
                    "metrics": [
                        {
                            "label": "Masked-token cross entropy",
                            "value": mlm_loss.item(),
                        },
                        {
                            "label": "Sentence-pair cross entropy",
                            "value": nsp_loss.item(),
                        },
                        {
                            "label": "Combined pretraining objective",
                            "value": (mlm_loss + nsp_loss).item(),
                        },
                    ],
                }
            )
    return {
        "kind": "matrices",
        "title": "Use both sides of context, predict selected targets.",
        "description": "Recorded MLM and NSP objectives on synthetic token IDs. The encoder can attend across both segments. Target selection defines the loss; unselected token logits are ignored by MLM. No language understanding is claimed.",
        "controlLabel": "Pretraining example",
        "cases": cases,
    }
