import os
import re
import math
import time
import json
import boto3
import torch
import random
import tempfile
import numpy as np
import torch.nn as nn
import torch.nn.functional as F

from tqdm import tqdm
from torch import autocast
from typing import List
from tokenizers import Tokenizer
from contextlib import contextmanager

# -----------------------------------------------------------------------------
# helpers
# -----------------------------------------------------------------------------


def save_checkpoint(model, optimizer, global_step, checkpoint_dir, config):
    print(f"Saving checkpoint to {checkpoint_dir}")
    torch.save(
        {
            "model": model.state_dict(),
            "optimizer": optimizer.state_dict(),
            "global_step": global_step,
            "config": config,
        },
        os.path.join(checkpoint_dir, f"last_ckpt.pth"),
    )


def load_model(model_filepath):
    """
    Load a GPTModel from a checkpoint file.

    Args:
        model_filepath (str): Path to the model checkpoint file. Can be a local path or S3 path.

    Returns:
        tuple: (model, config) - The loaded model and its configuration
    """
    with _download_from_s3_if_needed(model_filepath) as tmp_fp:
        checkpoint = torch.load(tmp_fp, map_location="cpu")

    config = checkpoint.get("config", GPTConfig())
    model = GPTModel(config)

    # Load model state dict
    model.load_state_dict(checkpoint["model"])

    return model, config


def read_jsonl(filepath):
    data = []
    with open(filepath) as f:
        for line in f:
            line = line.strip()
            if len(line) == 0:
                continue
            m = json.loads(line)
            data.append(m)
    return data


def write_jsonl(data, filepath, do_append=False):
    openarg = "a" if do_append else "w"
    with open(filepath, openarg) as f:
        for d in data:
            f.write(json.dumps(d, ensure_ascii=False) + "\n")


def get_filename(filepath, keep_ext=True):
    if "http" in filepath:
        clean_filepath = filepath.split("?")[0]
    else:
        clean_filepath = filepath
    filename = clean_filepath.split("/")[-1]
    if "." not in filename:
        raise ValueError("filename does not seem to contain a period.")
    m = re.search(r"(.+)\.([^\.]+)$", filename)
    if not m:
        raise ValueError(f"filename could not be parsed for `{filepath}`")
    filename = m.group(1)
    file_ext = m.group(2).lower()
    if len(file_ext) > 10:
        raise ValueError(f"file extension suspiciously long for `{filepath}`")
    if keep_ext:
        filename = filename + "." + file_ext
    return filename


S3_BUCKET_PATH_RE = r"s3\:\/\/(.+?)\/"


def _parse_s3_filepath(s3_filepath):
    bucket_name = re.search(S3_BUCKET_PATH_RE, s3_filepath).group(1)
    rel_s3_filepath = re.sub(S3_BUCKET_PATH_RE, "", s3_filepath)
    return bucket_name, rel_s3_filepath


def download_s3_file(
    from_s3_filepath,
    to_local_filepath,
):
    bucket_name, from_rel_s3_filepath = _parse_s3_filepath(from_s3_filepath)
    client = boto3.client("s3")
    client.download_file(bucket_name, from_rel_s3_filepath, to_local_filepath)


@contextmanager
def _download_from_s3_if_needed(maybe_s3_filepath):
    tmp_filepath = maybe_s3_filepath
    if maybe_s3_filepath.startswith("s3://"):
        temp_dir = tempfile.TemporaryDirectory()
        filename = get_filename(maybe_s3_filepath, keep_ext=True)
        tmp_filepath = os.path.join(temp_dir.name, filename)
        download_s3_file(maybe_s3_filepath, tmp_filepath)
    yield tmp_filepath


def load_tokenizer(
    tokenizer_filepath="s3://suno-data/georg/models/tokenizers/tokenizer_60k.json",
):
    with _download_from_s3_if_needed(tokenizer_filepath) as tmp_fp:
        tokenizer = Tokenizer.from_file(tmp_fp)
    tokenizer.add_special_tokens(["\n"])
    tokenizer.pad_idx = tokenizer.token_to_id("[PAD]")
    return tokenizer


# -----------------------------------------------------------------------------
# dataset of text + semantic tokens
# -----------------------------------------------------------------------------``

MAX_TAG_LEN = 256
MAX_TOT_TAGS_LEN = 512


def _space_repl(m):
    s = m.group()
    n_newline = s.count("\n")
    if n_newline >= 2:
        return "\n\n"
    elif n_newline == 1:
        return "\n"
    return " "


def _clean_tag(s, retain_newlines=False):
    s = s.replace("[", " ").replace("]", " ")
    return _simplify_whitespace(s, retain_newlines=retain_newlines)


def _simplify_whitespace(text, retain_newlines=True):
    """simplify while respecting up to 2 newlines"""
    if retain_newlines:
        text = re.sub(r"\s+", _space_repl, text).strip()
    else:
        text = re.sub(r"\s+", " ", text).strip()
    return text


def prepare_text_inference(tags: List[str], lyrics: str):
    tags = [
        clean_tag
        for tag in tags
        if len(clean_tag := _clean_tag(tag, retain_newlines=False)) > 0
    ]
    tags_str = (
        f"{', '.join([tag[:MAX_TAG_LEN] for tag in tags])[:MAX_TOT_TAGS_LEN]}".strip()
    )
    lyrics = _simplify_whitespace(lyrics, retain_newlines=True)

    # Combine tags and lyrics
    text = ""
    if len(tags_str) > 0:
        text += f"[{tags_str}]"
    if len(lyrics) > 0:
        if len(text) > 0:
            text += "\n\n"
        text += lyrics

    return text


class SemanticDataset(torch.utils.data.Dataset):
    def __init__(
        self,
        dataset_dir: str,
        memmap_filename: str,
        metas_filename: str,
        semantic_n_tokens: int,
        cond_text_len: int = 2560,
    ):
        self.semantic_n_tokens = semantic_n_tokens
        self.cond_text_len = cond_text_len
        # open semantic memmap
        semantic_data = np.memmap(
            os.path.join(dataset_dir, memmap_filename),
            dtype=np.uint16,
            mode="r",
        )
        semantic_data = semantic_data.reshape(-1, semantic_n_tokens, 1)
        self.semantic_data = semantic_data[:, :, 0]

        # load metas
        metas = read_jsonl(os.path.join(dataset_dir, metas_filename))
        self.metas = metas

        assert len(self.semantic_data) == len(self.metas)
        print(f"Loaded {len(self.semantic_data)} samples")

        # load tokenizer
        self.tokenizer = load_tokenizer()

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

    def __getitem__(self, idx):
        # semantic codes
        semantic_codes = torch.from_numpy(self.semantic_data[idx].copy()).long()
        # append eos token
        semantic_codes = torch.cat(
            [
                torch.tensor([SEMANTIC_SOS_TOKEN]),
                semantic_codes,
                torch.tensor([SEMANTIC_EOS_TOKEN]),
            ]
        )

        lyrics = self.metas[idx].get("text_aligned", self.metas[idx].get("text", ""))
        tags = self.metas[idx].get("tags", [])
        text = prepare_text_inference(tags, lyrics)

        # Build condition tensors
        text_codes = self.tokenizer.encode(text).ids[: self.cond_text_len]
        text_codes = text_codes + [self.tokenizer.pad_idx] * max(
            0, self.cond_text_len - len(text_codes)
        )
        # append infer token
        text_codes = text_codes
        text_codes = torch.tensor(text_codes).long()
        # for now assume all tokens are valid
        attention_mask = torch.ones(len(text_codes) + len(semantic_codes)).bool()

        return text_codes, semantic_codes, attention_mask


def collate_fn(batch):
    text_input_ids, semantic_input_ids, attention_mask = zip(*batch)

    text_input_ids = torch.stack(text_input_ids)
    semantic_input_ids = torch.stack(semantic_input_ids)
    attention_mask = torch.stack(attention_mask)

    # Create input/target pairs for causal language modeling
    # the text will be pre-prompt and we only compute loss on the semantic tokens
    labels = semantic_input_ids[:, 1:]  # get labels BEFORE modifying input_ids
    semantic_input_ids = semantic_input_ids[:, :-1]  # then truncate input_ids
    attention_mask = attention_mask[:, :-1]  # match input_ids length

    return {
        "text_input_ids": text_input_ids,
        "semantic_input_ids": semantic_input_ids,
        "attention_mask": attention_mask,
        "labels": labels,
    }


# -----------------------------------------------------------------------------

# model
# -----------------------------------------------------------------------------


class MultiHeadAttention(nn.Module):
    def __init__(self, config):
        super().__init__()
        assert (
            config.hidden_size % config.num_heads == 0
        ), "hidden_size must be divisible by num_heads"
        self.num_heads = config.num_heads
        self.hidden_size = config.hidden_size
        self.head_size = config.hidden_size // config.num_heads

        self.query = nn.Linear(config.hidden_size, config.hidden_size)
        self.key = nn.Linear(config.hidden_size, config.hidden_size)
        self.value = nn.Linear(config.hidden_size, config.hidden_size)
        self.proj = nn.Linear(config.hidden_size, config.hidden_size)
        self.dropout = nn.Dropout(config.dropout)

    def forward(self, x, attention_mask=None):
        B, T, C = x.size()

        # Split into heads
        q = self.query(x).view(B, T, self.num_heads, self.head_size).transpose(1, 2)
        k = self.key(x).view(B, T, self.num_heads, self.head_size).transpose(1, 2)
        v = self.value(x).view(B, T, self.num_heads, self.head_size).transpose(1, 2)

        # Scaled dot-product attention
        scale = math.sqrt(self.head_size)
        scores = torch.matmul(q, k.transpose(-2, -1)) / scale

        # Causal mask - prevent attending to future tokens
        causal_mask = torch.triu(torch.ones(T, T), diagonal=1).bool()
        scores.masked_fill_(causal_mask.to(scores.device), float("-inf"))

        # Apply attention mask if provided
        if attention_mask is not None:
            scores = scores.masked_fill(
                ~attention_mask.unsqueeze(1).unsqueeze(2), float("-inf")
            )

        attn = F.softmax(scores, dim=-1)
        attn = self.dropout(attn)

        # Apply attention to values
        out = torch.matmul(attn, v)

        # Reshape and project
        out = out.transpose(1, 2).contiguous().view(B, T, C)
        out = self.proj(out)

        return out


class TransformerBlock(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.attention = MultiHeadAttention(config)
        self.mlp = nn.Sequential(
            nn.Linear(config.hidden_size, config.mlp_ratio * config.hidden_size),
            nn.GELU(),
            nn.Linear(config.mlp_ratio * config.hidden_size, config.hidden_size),
            nn.Dropout(config.dropout),
        )
        self.ln1 = nn.LayerNorm(config.hidden_size)
        self.ln2 = nn.LayerNorm(config.hidden_size)
        self.dropout = nn.Dropout(config.dropout)

    def forward(self, x, attention_mask=None):
        # Pre-LayerNorm architecture
        x = x + self.dropout(self.attention(self.ln1(x), attention_mask))
        x = x + self.dropout(self.mlp(self.ln2(x)))
        return x


class GPTModel(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.config = config

        self.token_embeddings = nn.Embedding(config.text_vocab_size, config.hidden_size)
        self.semantic_embeddings = nn.Embedding(
            config.semantic_vocab_size, config.hidden_size
        )
        self.position_embeddings = nn.Embedding(
            config.max_position_embeddings, config.hidden_size
        )
        self.dropout = nn.Dropout(config.dropout)

        self.blocks = nn.ModuleList(
            [TransformerBlock(config) for _ in range(config.num_layers)]
        )

        self.ln_f = nn.LayerNorm(config.hidden_size)
        self.head = nn.Linear(
            config.hidden_size, config.semantic_vocab_size, bias=False
        )

        # Initialize weights
        self.apply(self._init_weights)

    def _init_weights(self, module):
        if isinstance(module, nn.Linear):
            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
            if module.bias is not None:
                torch.nn.init.zeros_(module.bias)
        elif isinstance(module, nn.Embedding):
            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
        elif isinstance(module, nn.LayerNorm):
            torch.nn.init.zeros_(module.bias)
            torch.nn.init.ones_(module.weight)

    def forward(self, text_input_ids, semantic_input_ids, attention_mask=None):
        B, T = text_input_ids.size()
        B, S = semantic_input_ids.size()

        # Get token embeddings
        token_emb = self.token_embeddings(text_input_ids)
        semantic_emb = self.semantic_embeddings(semantic_input_ids)

        # concatenate token and semantic embeddings
        input_emb = torch.cat([token_emb, semantic_emb], dim=1)

        # Add positional embeddings
        pos = torch.arange(0, T + S, dtype=torch.long, device=input_emb.device)
        pos_emb = self.position_embeddings(pos)

        x = self.dropout(input_emb + pos_emb)

        # Apply transformer blocks
        for block in self.blocks:
            x = block(x, attention_mask)

        x = self.ln_f(x)
        logits = self.head(x)

        return logits


# Example configuration

SEMANTIC_PAD_TOKEN = 4000
SEMANTIC_SOS_TOKEN = 4001
SEMANTIC_EOS_TOKEN = 4002


class GPTConfig:
    def __init__(self):
        self.text_vocab_size = 60004
        self.semantic_vocab_size = 4003
        self.max_position_embeddings = 751 + 751
        self.hidden_size = 1536
        self.num_layers = 12
        self.num_heads = 12
        self.mlp_ratio = 4
        self.dropout = 0.1
        self.semantic_pad_token = SEMANTIC_PAD_TOKEN
        self.semantic_sos_token = SEMANTIC_SOS_TOKEN
        self.semantic_eos_token = SEMANTIC_EOS_TOKEN


def train_step(model, batch, optimizer):
    optimizer.zero_grad()
    text_input_ids = batch["text_input_ids"].cuda()
    semantic_input_ids = batch["semantic_input_ids"].cuda()
    attention_mask = batch["attention_mask"].cuda()
    labels = batch["labels"].cuda()

    # forward pass
    with autocast(device_type="cuda", dtype=torch.bfloat16):
        logits = model(
            text_input_ids=text_input_ids,
            semantic_input_ids=semantic_input_ids,
            attention_mask=attention_mask,
        )
        # crop logits to the semantic tokens
        logits = logits[:, -semantic_input_ids.shape[1] :, :]

        # Reshape logits to [batch_size * sequence_length, vocab_size]
        logits = logits.reshape(-1, 4003)
        # Reshape labels to [batch_size * sequence_length]
        labels = labels.view(-1)

        loss = F.cross_entropy(logits, labels)

    loss.backward()
    optimizer.step()
    return loss


if __name__ == "__main__":
    config = GPTConfig()
    model = GPTModel(config)
    # Convert model to bfloat16
    model = model.to(dtype=torch.bfloat16)
    # model = torch.nn.DataParallel(model)

    # count parameters
    print(f"GPTModel: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M params")
    model.cuda()
    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)

    # setup dataset
    train_dataset = SemanticDataset(
        dataset_dir="/app/suno/data/diffusion_mix/dac_vae_fixed_25hz",
        memmap_filename="data_semantic_val.bin",
        metas_filename="metas_val.jsonl",
        semantic_n_tokens=750,
        cond_text_len=750,
    )
    train_dataloader = torch.utils.data.DataLoader(
        train_dataset,
        batch_size=20,
        shuffle=True,
        collate_fn=collate_fn,
        num_workers=8,
    )

    # train
    global_step = 0
    max_steps = 1_000_000

    run_start_time = time.strftime("%Y-%m-%d_%H-%M-%S")
    checkpoint_dir = f"/app/suno/christian/checkpoints/gpt/{run_start_time}_s{random.randint(0, 9999)}"
    os.makedirs(checkpoint_dir, exist_ok=False)

    # start training
    while global_step < max_steps:
        pbar = tqdm(train_dataloader)
        for batch in pbar:
            loss = train_step(model, batch, optimizer)
            pbar.set_description(f"Loss: {loss.item()}")
            global_step += 1

            if global_step % 100 == 0:
                save_checkpoint(model, optimizer, global_step, checkpoint_dir, config)

            if global_step >= max_steps:
                break
