from data import Vocab
from model import (
    Transformer,
    ModelArgs,
    SimpleLossCompute,
    LabelSmoothing,
    MusicalPositionEmbedTransformer,
)
import torch
from data_gen import Batch
from embedder import Embedder
import pytest
import torch.nn as nn
from typing import Optional


class LittleTrainer:
    def __init__(self, model):
        super().__init__()
        self.model = model
        self.loss_compute = SimpleLossCompute(LabelSmoothing(0, 0.01))
        self.optimizer = torch.optim.Adam(
            self.model.parameters(),
            lr=0.001,
            betas=(0.1, 0.15),
            eps=1e-7,
        )
        self.scaler = torch.cuda.amp.GradScaler()

    def step(self, input, target=None, **kwargs):
        with torch.cuda.amp.autocast(dtype=torch.float16, enabled=True):
            self.optimizer.zero_grad(set_to_none=True)

            if isinstance(input, Batch):
                pred = self.model.forward(
                    input.tgt,
                    start_pos=0,
                    **kwargs,
                )
                loss_node = self.loss_compute(pred, input.tgt_y)
            else:
                pred = self.model.forward(
                    input,
                    start_pos=0,
                    **kwargs,
                )
                loss_node = self.loss_compute(pred, target)

            loss = loss_node.item()
            print(loss)

            self.scaler.scale(loss_node).backward()
            self.scaler.unscale_(self.optimizer)
            torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
            self.scaler.step(self.optimizer)
            self.scaler.update()

            del loss_node

            return loss, pred


@pytest.fixture()
def test_vocab():
    return Vocab(1, 127, 1, 127, 20, 127, 2, 1, 127, 8, 24 * 4 * 8, 8, 24)


@pytest.fixture()
def untrained_model(test_vocab):
    smallargs = ModelArgs(
        256,
        2,
        4,
        test_vocab.N,
        cache=True,
        enable_flash=False,
        max_seq_len=384,
        max_inference_seq_len=384,
    )
    model = MusicalPositionEmbedTransformer(test_vocab, smallargs)
    model.eval()
    return model


class TestEmbedder(nn.Module):
    def __init__(self, max_length=64, embed_dim=64):
        super().__init__()
        self.max_length = max_length
        self.embed_dim = embed_dim
        self.dummy_param = nn.Parameter(torch.empty(0))

    def tokenize(self, sentences):
        device = next(self.parameters(recurse=True)).device
        z = torch.zeros(
            len(sentences), self.max_length, dtype=torch.long, device=device
        )
        for i in range(len(sentences)):
            s = sentences[i].encode("utf-8")
            for j in range(min(len(s), self.max_length)):
                z[i, j] = s[j]
        return z, torch.ones(len(sentences), self.max_length, device=device)

    def forward(
        self,
        input_ids: torch.LongTensor,
        attention_mask: Optional[torch.FloatTensor] = None,
    ):
        assert input_ids.shape == attention_mask.shape
        return torch.normal(
            0,
            0.2,
            (input_ids.shape[0], input_ids.shape[1], self.embed_dim),
            device=input_ids.device,
        )


@pytest.fixture(params=[True, False])
def untrained_cross_model(test_vocab, request):
    cross_embed_dim = 64
    smallargs = ModelArgs(
        256,
        4,
        4,
        test_vocab.N,
        cache=True,
        enable_flash=request.param,
        enable_cross_attention=True,
        max_seq_len=384,
        max_inference_seq_len=384,
        cross_attention_embedding_dim=cross_embed_dim,
    )
    model = MusicalPositionEmbedTransformer(
        test_vocab, smallargs, encoder=TestEmbedder(embed_dim=cross_embed_dim)
    )
    model.eval()
    if request.param:
        model = model.cuda()
    return model


@pytest.fixture(params=[True, False])
def untrained_embed_nocross_model(test_vocab, request):
    smallargs = ModelArgs(
        256,
        4,
        4,
        test_vocab.N,
        cache=True,
        enable_flash=request.param,
        enable_cross_attention=False,
        max_seq_len=384,
        max_inference_seq_len=383,
    )
    model = MusicalPositionEmbedTransformer(
        test_vocab, smallargs, encoder=TestEmbedder(embed_dim=256)
    )
    model.eval()
    if request.param:
        model = model.cuda()
    return model
