import os
import math
import torch
import einsum
import numpy as np
from torch import nn, einsum
import torch.optim as optim
from tqdm import tqdm
from torch.nn.functional import mse_loss

from suno_utils.utils.text import read_jsonl


class CausalSelfAttention(nn.Module):
    def __init__(self, embed_size, num_heads):
        super().__init__()
        self.multihead_attn = nn.MultiheadAttention(
            embed_size, num_heads, batch_first=True
        )

    def forward(self, x):
        seq_length = x.size(1)
        # Create a causal mask
        mask = torch.triu(torch.ones(seq_length, seq_length), diagonal=1).bool()
        mask = mask.to(x.device)

        # Apply causal self-attention
        attn_output, _ = self.multihead_attn(x, x, x, attn_mask=mask)
        return attn_output


# 1. Model Architecture
class GPTModel(nn.Module):
    def __init__(
        self,
        acoustic_codebook_size: int,
        n_acoustic_codebooks: int,
        embed_size: int,
        num_heads: int,
        num_layers: int,
        max_seq_length: int,
    ):
        super(GPTModel, self).__init__()
        self.token_embedding = nn.Embedding(acoustic_codebook_size, embed_size)
        self.position_embedding = nn.Embedding(max_seq_length, embed_size)

        self.layers = nn.ModuleList(
            [
                nn.Sequential(
                    CausalSelfAttention(embed_size, num_heads),
                    nn.LayerNorm(embed_size),
                    nn.Linear(embed_size, embed_size * 4),
                    nn.GELU(),
                    nn.Linear(embed_size * 4, embed_size),
                    nn.LayerNorm(embed_size),
                )
                for _ in range(num_layers)
            ]
        )

        self.acoustic_heads = torch.nn.ModuleList()
        for n in range(n_acoustic_codebooks):
            self.acoustic_heads.append(nn.Linear(embed_size, acoustic_codebook_size))

    def forward(self, x: torch.Tensor):
        bs, seq_length, n_codebooks = x.size()
        position_ids = torch.arange(seq_length, device=x.device).unsqueeze(0)

        token_embeds = self.token_embedding(x)
        position_embeds = self.position_embedding(position_ids)
        x = token_embeds + position_embeds
        # x: (bs, seq_length, n_codebooks, embed_size)

        # sum across codebook dimension
        x = x.sum(dim=-1)  # (bs, seq_length, embed_size)

        for layer in self.layers:
            x = x + layer(x)  # Residual connection

        # Apply acoustic heads
        outputs = []
        for head in self.acoustic_heads:
            outputs.append(head(x))
        # stack outputs along the last dimension
        outputs = torch.stack(outputs, dim=-1)

        return outputs


def apply_delay_pattern(x, delay: int = 1, pad_token: int = 0):
    batch_size, seq_length, n_codebooks = x.size()

    # Calculate the maximum shift
    max_shift = delay * (n_codebooks - 1)

    # Create a new tensor filled with pad_token
    result = torch.full(
        (batch_size, seq_length + max_shift, n_codebooks),
        pad_token,
        dtype=x.dtype,
        device=x.device,
    )

    for i in range(n_codebooks):
        shift = delay * i
        result[:, shift : shift + seq_length, i] = x[:, :, i]

    return result


def restore_original_alignment(x, delay: int = 1, pad_token: int = 0):
    batch_size, extended_seq_length, n_codebooks = x.size()

    # Calculate the original sequence length
    original_seq_length = extended_seq_length - delay * (n_codebooks - 1)

    # Create a new tensor to store the result
    result = torch.zeros(
        batch_size, original_seq_length, n_codebooks, dtype=x.dtype, device=x.device
    )

    for i in range(n_codebooks):
        shift = delay * i
        result[:, :, i] = x[:, shift : shift + original_seq_length, i]

    return result


# 2. Dataset
class MemmapDataset(torch.utils.data.Dataset):
    def __init__(
        self,
        acoustic_tokens_memmap_path: str,
        n_tokens_memmap: int,
        n_acoustic_codebooks: int,
        acoustic_codebook_size: int,
    ):
        acoustic_tokens = np.memmap(
            acoustic_tokens_memmap_path, dtype=np.uint16, mode="r"
        )
        acoustic_tokens = acoustic_tokens.reshape(
            -1, n_tokens_memmap, n_acoustic_codebooks
        )

        print(f"Acoustic tokens shape: {acoustic_tokens.shape}")
        self.acoustic_tokens = acoustic_tokens
        self.acoustic_codebook_size = acoustic_codebook_size
        self.pad_token = acoustic_codebook_size + 1

    def __len__(self):
        return self.acoustic_tokens.shape[0]

    def __getitem__(self, idx):
        input_seq = torch.from_numpy(self.acoustic_tokens[idx, ...].copy())
        # input_seq: (n_tokens_memmap, n_acoustic_codebooks)

        # Apply delay pattern to input sequence
        input_seq = apply_delay_pattern(input_seq, delay=1, pad_token=self.pad_token)

        # create target sequence by shifting input sequence by 1
        target_seq = input_seq.clone()
        target_seq[:-1] = input_seq[1:]
        target_seq[-1] = self.pad_token

        return input_seq, target_seq


# 25 hz codec
# 10 sec chunks -> 250 codec time steps
# 12 codebooks means 250 x 12 = 3000 tokens seq length for flat
# delay pattern one should reduce that 250 + 1 = 251 tokens


# 3. Training Loop
def train(model, dataloader, optimizer, criterion, device):
    model.train()
    total_loss = 0
    pbar = tqdm(dataloader, total=len(dataloader))
    for input_seq, target_seq in pbar:
        input_seq, target_seq = input_seq.to(device), target_seq.to(device)

        optimizer.zero_grad()
        output = model(input_seq)
        loss = criterion(output.view(-1, output.size(-1)), target_seq.view(-1))
        loss.backward()
        optimizer.step()

        total_loss += loss.item()
        pbar.set_description(f"Loss: {loss.item():.4f}")

    return total_loss / len(dataloader)


if __name__ == "__main__":
    # Hyperparameters
    vocab_size = 1000  # Example value
    embed_size = 256
    num_heads = 8
    num_layers = 6
    max_seq_length = 100
    batch_size = 32
    num_epochs = 10
    learning_rate = 0.001

    # Device configuration
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    # Create model
    model = GPTModel(vocab_size, embed_size, num_heads, num_layers, max_seq_length).to(
        device
    )

    # print number of model parameters in millions
    num_params = sum(p.numel() for p in model.parameters()) / 1_000_000
    print(f"Number of GPT parameters: {num_params:.2f}M")

    dataset = MemmapDataset(coarse_memmap_path, max_seq_length)
    dataloader = torch.utils.data.DataLoader(
        dataset, batch_size=batch_size, shuffle=True
    )

    # Loss and optimizer
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=learning_rate)

    # Training loop
    for epoch in range(num_epochs):
        loss = train(model, dataloader, optimizer, criterion, device)
