"""可训练的微型 Decoder-only LM；用于理解链路，不用于生产。"""
import torch
from torch import nn


class TinyLM(nn.Module):
    def __init__(self, vocab_size=256, d_model=128, heads=4, layers=2, max_len=128):
        super().__init__()
        self.token = nn.Embedding(vocab_size, d_model)
        self.position = nn.Embedding(max_len, d_model)
        block = nn.TransformerEncoderLayer(d_model, heads, 4 * d_model, batch_first=True, norm_first=True)
        self.blocks = nn.TransformerEncoder(block, layers)
        self.norm = nn.LayerNorm(d_model)
        self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
        self.lm_head.weight = self.token.weight

    def forward(self, input_ids: torch.Tensor):
        _, length = input_ids.shape
        if length > self.position.num_embeddings:
            raise ValueError("sequence exceeds max_len")
        positions = torch.arange(length, device=input_ids.device)
        x = self.token(input_ids) + self.position(positions)[None]
        causal = torch.triu(torch.ones(length, length, device=x.device, dtype=torch.bool), diagonal=1)
        x = self.blocks(x, mask=causal)
        return self.lm_head(self.norm(x))


if __name__ == "__main__":
    model = TinyLM()
    ids = torch.randint(0, 256, (2, 16))
    logits = model(ids)
    loss = nn.functional.cross_entropy(logits[:, :-1].reshape(-1, 256), ids[:, 1:].reshape(-1))
    loss.backward()
    assert logits.shape == (2, 16, 256)
    print(float(loss))
