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

def run_distillation():
    # 1. Setup
    device = "cuda" if torch.cuda.is_available() else "cpu"
    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)

    # 2. Collect Data
    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)

    # 3. Fit FEM Mesh
    num_nodes = 128
    per_layer_data = []

    for layer in tqdm(range(num_layers), desc="Fitting FEM"):
        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), torch.tensor(b_node), kmeans))
        ffn_inputs[layer] = None
        ffn_outputs[layer] = None

    return model_orig, per_layer_data

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)
        ids = self.km.predict(x_f.detach().cpu().numpy())
        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])

if __name__ == "__main__":
    model_orig, data = run_distillation()
    fem_model = FEMGPT2(model_orig, data)
    print("FEM-LLM setup complete.")
