import torch
import numpy as np
import zlib
import pickle
import os
import gc
from transformers import GPT2LMHeadModel, GPT2Config, AutoTokenizer

# -------------------------------
# Compression/decompression functions
# -------------------------------
def quantize_tensor(tensor, num_levels=256):
    min_val = tensor.min().item()
    max_val = tensor.max().item()
    if min_val == max_val:
        q = torch.zeros_like(tensor, dtype=torch.uint8)
        scale = 0.0
        return q, min_val, max_val, scale
    scale = (max_val - min_val) / (num_levels - 1)
    q = torch.round((tensor - min_val) / scale).clamp(0, num_levels-1).to(torch.uint8)
    return q, min_val, max_val, scale

def compress_weights(model, num_levels=256):
    compressed_data = {}
    total_original_bytes = 0
    total_compressed_bytes = 0
    seen_tensors = set()
    for name, param in model.named_parameters():
        if param.requires_grad and id(param) not in seen_tensors:
            seen_tensors.add(id(param))
            q, min_val, max_val, scale = quantize_tensor(param.detach().cpu(), num_levels)
            q_np = q.numpy().tobytes()
            comp = zlib.compress(q_np, level=9)
            compressed_data[name] = {
                'compressed': comp,
                'shape': param.shape,
                'min_val': min_val,
                'max_val': max_val,
                'num_levels': num_levels,
                'dtype': param.dtype
            }
            total_original_bytes += param.numel() * param.element_size()
            total_compressed_bytes += len(comp)
    ratio = total_compressed_bytes / total_original_bytes
    print(f"Compression ratio: {ratio:.4f}")
    return compressed_data

def decompress_state_dict(compressed_data):
    state_dict = {}
    for name, info in compressed_data.items():
        q_bytes = zlib.decompress(info['compressed'])
        q_np = np.frombuffer(q_bytes, dtype=np.uint8).reshape(info['shape']).copy()
        q = torch.from_numpy(q_np).float()
        min_val = info['min_val']
        max_val = info['max_val']
        num_levels = info['num_levels']
        if num_levels > 1:
            scale = (max_val - min_val) / (num_levels - 1)
            dequantized = min_val + scale * q
        else:
            dequantized = torch.full_like(q, min_val)
        state_dict[name] = dequantized.to(info['dtype'])

    # Handle tied weights: lm_head shares weights with wte in GPT-2
    if 'transformer.wte.weight' in state_dict:
        state_dict['lm_head.weight'] = state_dict['transformer.wte.weight'].clone()

    return state_dict

# -------------------------------
# 1. Check if compressed model exists, otherwise compress
# -------------------------------
model_name = "distilgpt2"
compressed_file = f"{model_name}_compressed.pkl"

if os.path.exists(compressed_file):
    print(f"✓ Found existing compressed model: {compressed_file}")
    print("Skipping compression, loading directly...\n")
else:
    print(f"No compressed model found. Loading and compressing {model_name}...\n")
    
    # Load original model
    model_orig = GPT2LMHeadModel.from_pretrained(model_name)
    model_orig.eval()
    
    # Compress weights
    print("\nCompressing weights...")
    compressed = compress_weights(model_orig, num_levels=256)
    
    # Save compressed dict to disk
    with open(compressed_file, "wb") as f:
        pickle.dump(compressed, f)
    print(f"Compressed model saved to {compressed_file}")
    
    # Free original model memory
    del model_orig, compressed
    gc.collect()

# -------------------------------
# 2. Load compressed weights and decompress into model
# -------------------------------
print(f"\nLoading compressed weights from {compressed_file}...")
with open(compressed_file, "rb") as f:
    compressed_loaded = pickle.load(f)

print("Decompressing weights...")
decompressed_state_dict = decompress_state_dict(compressed_loaded)

# Create model with decompressed weights
config = GPT2Config.from_pretrained(model_name)
model = GPT2LMHeadModel(config)
model.load_state_dict(decompressed_state_dict, strict=True)
model.eval()

# Free memory
del decompressed_state_dict, compressed_loaded
gc.collect()

print("\n✓ Compressed model loaded successfully!")

# -------------------------------
# 4. Interactive prompt
# -------------------------------
tokenizer = AutoTokenizer.from_pretrained(model_name)
tokenizer.pad_token = tokenizer.eos_token

print("\n" + "="*60)
print("Compressed Model Interactive Chat")
print("="*60)
print("Type your prompts and press Enter. Type 'quit' or 'exit' to stop.")
print("Press Enter with empty input to quit.\n")

while True:
    user_input = input("You: ").strip()
    
    if not user_input or user_input.lower() in ['quit', 'exit']:
        print("\nGoodbye!")
        break
    
    # Build prompt
    prompt = f"{user_input}"
    input_ids = tokenizer.encode(prompt, return_tensors="pt")
    attention_mask = torch.ones_like(input_ids)
    
    # Generate response
    with torch.no_grad():
        output = model.generate(
            input_ids,
            attention_mask=attention_mask,
            max_length=len(input_ids[0]) + 100,
            do_sample=True,
            temperature=0.8,
            top_k=50,
            top_p=0.95,
            pad_token_id=tokenizer.eos_token_id,
            eos_token_id=tokenizer.eos_token_id
        )
    
    response = tokenizer.decode(output[0], skip_special_tokens=True)
    print(f"\nModel: {response}\n")
    print("-" * 40 + "\n")
