from dataclasses import dataclass
import math
from typing import Optional

import torch
import torch.nn as nn
from torch.nn import functional as F

from .base import (
    Block,
    NormFunc,
    configure_optimizers,
    estimate_mfu,
    get_init_fn,
    init_weights_simple,
)


TIE_WEIGHTS = False
USE_SIN_POS_EMB = False
SIMPLE_INIT = True


def create_sin_embedding(
    positions: torch.Tensor,
    dim: int,
    max_period: float = 10_000,
    dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
    """Create sinusoidal positional embedding, with shape `[B, T, C]`"""
    assert dim % 2 == 0
    half_dim = dim // 2
    positions = positions.to(dtype)
    adim = torch.arange(half_dim, device=positions.device, dtype=dtype).view(1, 1, -1)
    max_period_tensor = torch.full(
        [], max_period, device=positions.device, dtype=dtype
    )  # avoid sync point
    phase = positions / (max_period_tensor ** (adim / (half_dim - 1)))
    return torch.cat([torch.cos(phase), torch.sin(phase)], dim=-1)


@dataclass
class FineConfig:
    n_layer: int = 24
    n_head: int = 16  # query heads
    n_kv_head: Optional[int] = None
    d_head: int = 64
    bias: bool = False
    dropout: float = 0.0

    coarse_vocab_size: int = 4160
    coarse_codebook_size: int = 4096
    coarse_n_codebooks: int = 8
    coarse_pad_token: int = 4096
    coarse_infer_token: int = 4097
    coarse_rate_hz: int = 25
    coarse_shift_factor: int = 0
    coarse_samples: int = 50
    # strided mask
    coarse_mask_period: int = 1
    coarse_masked_samples: int = 0

    fine_vocab_size: int = 1152
    fine_codebook_size: int = 1024
    fine_n_codebooks: int = 16
    fine_pad_token: int = 1024
    fine_infer_token: int = 1025
    fine_rate_hz: int = 100
    fine_shift_factor: int = 1
    fine_samples: int = 200

    t_memmap: int = 3375  # not sure what this is yet

    def __post_init__(self):
        # default to multi head attention
        if self.n_kv_head is None:
            self.n_kv_head = self.n_head

        assert self.coarse_masked_samples <= self.coarse_mask_period

    @property
    def n_embd(self):
        """The width of the residual stream"""
        return self.n_head * self.d_head

    @property
    def t_fine(self):
        """Includes infer token"""
        return self.fine_samples + self.fine_shift_factor * (self.fine_n_codebooks - 1) + 1

    @property
    def t_coarse(self):
        return self.coarse_samples

    @property
    def block_size(self):
        return self.t_coarse + self.t_fine - 1


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

        model_dict = dict(
            wte_fine=nn.ModuleList(
                [
                    nn.Embedding(config.fine_vocab_size, config.n_embd)
                    for _ in range(config.fine_n_codebooks)
                ]
            ),
            ln_fine=NormFunc(config.n_embd),
            wte_coarse=nn.ModuleList(
                [
                    nn.Embedding(config.coarse_vocab_size, config.n_embd)
                    for _ in range(config.coarse_n_codebooks)
                ]
            ),
            ln_coarse=NormFunc(config.n_embd),
            drop=nn.Dropout(config.dropout),
            h=nn.ModuleList([Block(config) for _ in range(config.n_layer)]),
            ln_f=NormFunc(config.n_embd),
        )
        if not USE_SIN_POS_EMB:
            model_dict["wpe"] = nn.Embedding(config.block_size, config.n_embd)
        self.transformer = nn.ModuleDict(model_dict)
        self.lm_heads = nn.ModuleList(
            [
                nn.Linear(config.n_embd, config.fine_vocab_size, bias=False)
                for _ in range(config.fine_n_codebooks)
            ]
        )
        if TIE_WEIGHTS:
            for n in range(config.fine_n_codebooks):
                self.transformer.wte_fine[n].weight = self.lm_heads[n].weight

        # init all weights
        if SIMPLE_INIT:
            self.apply(self._init_weights_simple)
            for pn, p in self.named_parameters():
                if pn.endswith("c_proj.weight"):
                    torch.nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer))
        else:
            self._init_weights()

        print(f"number of parameters: {self.get_num_params()/1e6:.0f}M")

    def forward(self, x, y=None, coarse_offset=0, return_logits=False, last_only=True):
        device = x.device
        b, ns, t = x.size()
        assert ns == self.config.fine_n_codebooks + self.config.coarse_n_codebooks

        if y is not None:
            assert t == self.config.block_size
            _, _, t2 = y.size()
            assert t2 == self.config.t_fine - 1, (t2, self.config.t_fine)

        x_emb = 0
        # embed coarse
        for n in range(self.config.coarse_n_codebooks):
            x_emb += self.transformer.ln_coarse(self.transformer.wte_coarse[n](x[:, n, :]))
        # embed fine
        for n in range(self.config.fine_n_codebooks):
            n2 = n + self.config.coarse_n_codebooks
            x_emb += self.transformer.ln_fine(self.transformer.wte_fine[n](x[:, n2, :]))

        # x_emb (b, t, n_embd)

        if USE_SIN_POS_EMB:
            pos = torch.arange(t, device=x.device).view(1, -1, 1)
            # pos = pos + offsets.view(-1, 1, 1)
            pos_emb = create_sin_embedding(pos, self.config.n_embd, dtype=x_emb.dtype)
        else:
            pos = torch.arange(t, dtype=torch.long, device=device).unsqueeze(0)  # shape (1, t)
            pos_emb = self.transformer.wpe(pos)  # (1, t, n_embd)

        x = self.transformer.drop(x_emb + pos_emb)

        for block in self.transformer.h:
            x = block(x)
        x = self.transformer.ln_f(x)

        x = x[:, coarse_offset : coarse_offset + self.config.t_fine - 1, :]
        if return_logits:
            if last_only:
                x = x[:, -1, :]
            fine_logits_list = []
            for n in range(self.config.fine_n_codebooks):
                fine_logits_list.append(self.lm_heads[n](x))
            fine_logits = torch.stack(fine_logits_list).swapaxes(0, 1)
            return fine_logits

        loss_dict = {}
        for n in range(self.config.fine_n_codebooks):
            logits = self.lm_heads[n](x)
            loss_dict[f"fine_{n}"] = F.cross_entropy(
                logits.reshape(-1, logits.size(-1)),
                y[:, n, :].reshape(-1),
                ignore_index=-1,
            )

        return loss_dict

    def get_num_params(self, non_embedding=True):
        n_params = sum(p.numel() for p in self.parameters())
        if non_embedding:
            for m in self.transformer.wte_coarse:
                n_params -= m.weight.numel()
            for m in self.transformer.wte_fine:
                n_params -= m.weight.numel()
            if not USE_SIN_POS_EMB:
                n_params -= self.transformer.wpe.weight.numel()
        return n_params

    def _init_weights_simple(self, module):
        init_weights_simple(self, module)

    def _init_weights(self):
        # embeddings
        get_init_fn(self.config.n_embd, init_depth=None)(self.transformer.wte_text.weight)
        for module in self.transformer.wte_coarse:
            get_init_fn(self.config.n_embd, init_depth=None)(module.weight)
        for module in self.transformer.wte_fine:
            get_init_fn(self.config.n_embd, init_depth=None)(module.weight)
        if not USE_SIN_POS_EMB:
            get_init_fn(self.config.n_embd, init_depth=None)(self.transformer.wpe.weight)
        # heads
        for module in self.lm_heads:
            get_init_fn(self.config.n_embd, init_depth=None)(module.weight)
            if module.bias is not None:
                torch.nn.init.zeros_(module.bias)
        # attention blocks
        for layer_idx, block in enumerate(self.transformer.h):
            # mlp
            module = block.mlp.c_fc
            get_init_fn(self.config.n_embd, init_depth=layer_idx + 1)(module.weight)
            if module.bias is not None:
                torch.nn.init.zeros_(module.bias)
            module = block.mlp.c_proj
            get_init_fn(block.mlp.embd_inner, init_depth=layer_idx + 1)(module.weight)
            if module.bias is not None:
                torch.nn.init.zeros_(module.bias)
            # attention
            for module in [block.attn.c_attn, block.attn.c_proj]:
                get_init_fn(self.config.n_embd, init_depth=layer_idx + 1)(module.weight)
                if module.bias is not None:
                    torch.nn.init.zeros_(module.bias)

    def configure_optimizers(self, weight_decay, learning_rate, betas, device_type):
        return configure_optimizers(self, weight_decay, learning_rate, betas, device_type)

    def estimate_mfu(self, fwdbwd_per_iter, dt):
        return estimate_mfu(self, fwdbwd_per_iter, dt)
