import torch
import torch.nn as nn
import numpy as np
from transformers import GPT2LMHeadModel, GPT2Tokenizer
from datasets import load_dataset
from sklearn.cluster import MiniBatchKMeans
from tqdm import tqdm
from collections import defaultdict
import time
import gc

def run_distillation():
    device = "cuda" if torch.cuda.is_available() else "cpu"
    print(f"Using device: {device}")

    model_orig = GPT2LMHeadModel.from_pretrained("gpt2").to(device)
    model_orig.eval()
    config = model_orig.config
    hidden_size = config.n_embd
    num_layers = config.n_layer

    tokenizer = GPT2Tokenizer.from_pretrained("gpt2")
    dataset = load_dataset("wikitext", "wikitext-2-raw-v1", split="train")
    text = " ".join(dataset["text"][:50])
    inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=256)
    input_ids = inputs["input_ids"].to(device)

    ffn_inputs = defaultdict(list)
    ffn_outputs = defaultdict(list)

    def make_pre_hook(layer_idx):
        def pre_hook(module, inp):
            x = inp[0].detach().cpu().numpy()
            ffn_inputs[layer_idx].append(x.reshape(-1, hidden_size))
        return pre_hook

    def make_post_hook(layer_idx):
        def post_hook(module, inp, out):
            y = out.detach().cpu().numpy()
            ffn_outputs[layer_idx].append(y.reshape(-1, hidden_size))
        return post_hook

    for i, block in enumerate(model_orig.transformer.h):
        block.mlp.register_forward_pre_hook(make_pre_hook(i))
        block.mlp.register_forward_hook(make_post_hook(i))

    with torch.no_grad():
        _ = model_orig(input_ids)

    num_nodes = 128
    per_layer_data = []

    for layer in tqdm(range(num_layers), desc="Fitting FEM Layers"):
        X = np.concatenate(ffn_inputs[layer], axis=0)
        Y = np.concatenate(ffn_outputs[layer], axis=0)

        kmeans = MiniBatchKMeans(n_clusters=num_nodes, batch_size=1024, random_state=42, n_init=3)
        labels = kmeans.fit_predict(X)

        a_node = np.ones((num_nodes, hidden_size), dtype=np.float16)
        b_node = np.zeros((num_nodes, hidden_size), dtype=np.float16)

        for n in range(num_nodes):
            idxs = np.where(labels == n)[0]
            if len(idxs) > 5:
                X_n, Y_n = X[idxs], Y[idxs]
                for d in range(hidden_size):
                    a = np.cov(X_n[:, d], Y_n[:, d])[0,1] / (np.var(X_n[:, d]) + 1e-6)
                    b = np.mean(Y_n[:, d]) - a * np.mean(X_n[:, d])
                    a_node[n, d], b_node[n, d] = a, b

        per_layer_data.append((torch.tensor(a_node, dtype=torch.float16), torch.tensor(b_node, dtype=torch.float16), kmeans))
        ffn_inputs[layer] = None
        ffn_outputs[layer] = None

    return model_orig, per_layer_data, input_ids

class FEMLinearFFN(nn.Module):
    def __init__(self, a, b, km):
        super().__init__()
        self.register_buffer("a", a)
        self.register_buffer("b", b)
        self.km = km
    def forward(self, x):
        s, h = x.shape[:-1], x.shape[-1]
        x_f = x.view(-1, h)
        pred = self.km.predict(x_f.detach().cpu().numpy())
        ids = torch.from_numpy(pred).to(x.device)
        # self.a and self.b are now buffers, so they move with .to(device)
        return (self.a[ids] * x_f + self.b[ids]).view(*s, h)

class FEMGPT2(nn.Module):
    def __init__(self, orig, data):
        super().__init__()
        self.transformer = orig.transformer
        self.lm_head = orig.lm_head
        for i, block in enumerate(self.transformer.h):
            a, b, km = data[i]
            block.mlp = FEMLinearFFN(a, b, km)
    def forward(self, input_ids, attention_mask=None):
        return self.lm_head(self.transformer(input_ids, attention_mask=attention_mask)[0])

def benchmark(model, ids):
    model.eval()
    torch.cuda.empty_cache()
    torch.cuda.reset_peak_memory_stats()
    with torch.no_grad():
        for _ in range(5): _ = model(ids)
    mem = torch.cuda.max_memory_allocated() / (1024**2)
    print(f"Peak Memory: {mem:.2f} MB")


def generate_text(model, tokenizer, prompt, max_length=20):
    model.eval()
    device = next(model.parameters()).device
    inputs = tokenizer(prompt, return_tensors='pt').to(device)
    input_ids = inputs['input_ids']
    
    print(f'Prompt: {prompt}')
    
    with torch.no_grad():
        for _ in range(max_length):
            outputs = model(input_ids)
            next_token_logits = outputs[:, -1, :]
            next_token = torch.argmax(next_token_logits, dim=-1).unsqueeze(-1)
            input_ids = torch.cat([input_ids, next_token], dim=-1)
            
            if next_token.item() == tokenizer.eos_token_id:
                break
                
    return tokenizer.decode(input_ids[0], skip_special_tokens=True)


if __name__ == "__main__":
    model_orig, data, input_ids = run_distillation()
    fem_model = FEMGPT2(model_orig, data).to('cpu')
    input_ids = input_ids.to('cpu')
    del model_orig
    gc.collect()
    torch.cuda.empty_cache()
    benchmark(fem_model, input_ids[:, :128])

    # Test the FEM model
    if 'fem_model' in globals() and 'tokenizer' in globals():
        test_prompt = 'The quick brown fox'
        generated = generate_text(fem_model, tokenizer, test_prompt)
        print(f'Generated: {generated}')
    else:
        print('Model or tokenizer not found. Please run the distillation cell first.')

    
"""
 Using device: cuda
Loading weights: 100% 148/148 [00:00<00:00, 608.79it/s, Materializing param=transformer.wte.weight]GPT2LMHeadModel LOAD REPORT from: gpt2
Key                  | Status     |  | 
---------------------+------------+--+-
h.{0...11}.attn.bias | UNEXPECTED |  | 

Notes:
- UNEXPECTED	:can be ignored when loading from different task/architecture; not ok if you expect identical arch.
Fitting FEM Layers: 100%|██████████| 12/12 [00:17<00:00,  1.45s/it]
Peak Memory: 363.31 MB
"""
