"""
QLoRA AGI Training

Fine-tune a quantized 7B model with salience-guided curriculum.
Uses ~5GB VRAM - won't melt your GPU.

The adapter weights can be:
1. Merged with base model
2. Shared with others
3. Combined with other LoRA adapters
"""

import torch
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
import time
import os
from dataclasses import dataclass
from transformers import (
    AutoModelForCausalLM, 
    AutoTokenizer,
    BitsAndBytesConfig,
    TrainingArguments,
)
from peft import (
    LoraConfig, 
    get_peft_model, 
    prepare_model_for_kbit_training,
    PeftModel
)
from datasets import load_dataset

DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")


@dataclass
class Config:
    # Model - using a small but capable base
    base_model: str = "TinyLlama/TinyLlama-1.1B-Chat-v1.0"  # 1.1B, fits easily
    # base_model: str = "microsoft/phi-2"  # 2.7B, also fits
    # base_model: str = "mistralai/Mistral-7B-v0.1"  # 7B, tight fit with 4-bit
    
    # LoRA config
    lora_r: int = 64          # Rank - higher = more capacity
    lora_alpha: int = 128     # Scaling
    lora_dropout: float = 0.05
    
    # Training
    batch_size: int = 2
    grad_accum: int = 8
    lr: float = 2e-4
    max_steps: int = 3000
    ctx_len: int = 256
    
    # Quantization
    use_4bit: bool = True
    bnb_4bit_compute_dtype: str = "float16"


class BootstrapDataset(Dataset):
    """Bootstrap curriculum for fine-tuning."""
    
    def __init__(self, tokenizer, ctx_len: int = 256):
        self.tokenizer = tokenizer
        self.ctx_len = ctx_len
        
        # Core lessons formatted for instruction-following
        self.lessons = self._build_curriculum()
        
        # Tokenize
        self.examples = []
        for lesson in self.lessons:
            tokens = tokenizer(
                lesson, 
                truncation=True, 
                max_length=ctx_len,
                padding="max_length",
                return_tensors="pt"
            )
            self.examples.append({
                'input_ids': tokens['input_ids'].squeeze(),
                'attention_mask': tokens['attention_mask'].squeeze()
            })
        
        print(f"Bootstrap dataset: {len(self.examples)} examples")
    
    def _build_curriculum(self):
        """Build instruction-formatted curriculum."""
        lessons = []
        
        # Logic and reasoning
        logic_lessons = [
            "Question: If all birds can fly and a sparrow is a bird, can a sparrow fly?\nAnswer: Yes. Since all birds can fly, and a sparrow is a bird, therefore a sparrow can fly. This is deductive reasoning.",
            "Question: If it rains, the ground gets wet. The ground is wet. Did it rain?\nAnswer: Not necessarily. The ground could be wet for other reasons (sprinkler, spilled water). This is the fallacy of affirming the consequent.",
            "Question: What is the logical operator AND?\nAnswer: AND requires both conditions to be true. A AND B is true only when A is true and B is true. If either is false, the result is false.",
            "Question: What is the difference between correlation and causation?\nAnswer: Correlation means two things happen together. Causation means one thing causes the other. Correlation does not imply causation - there may be a third factor causing both.",
            "Question: How do you prove something is false?\nAnswer: Find one counterexample. If someone claims 'all swans are white' and you find one black swan, the claim is proven false.",
        ]
        
        # Mathematics
        math_lessons = [
            "Question: What is 15 + 27?\nAnswer: 15 + 27 = 42. I added the ones (5+7=12, carry 1) then the tens (1+2+1=4), giving 42.",
            "Question: What is 8 × 7?\nAnswer: 8 × 7 = 56. Multiplication is repeated addition: 8 + 8 + 8 + 8 + 8 + 8 + 8 = 56.",
            "Question: Solve for x: x + 5 = 12\nAnswer: x = 7. Subtract 5 from both sides: x + 5 - 5 = 12 - 5, so x = 7.",
            "Question: What is the pattern: 2, 4, 8, 16, ...?\nAnswer: Each number is doubled. The pattern is powers of 2. The next number is 32 (16 × 2).",
            "Question: What is 144 ÷ 12?\nAnswer: 144 ÷ 12 = 12. Division asks: how many groups of 12 fit in 144? The answer is 12 groups.",
        ]
        
        # Meta-learning
        meta_lessons = [
            "Question: How do you learn something new effectively?\nAnswer: 1) Connect new information to what you already know. 2) Practice actively, don't just read passively. 3) Test yourself to identify gaps. 4) Space repetition over time. 5) Teach it to someone else.",
            "Question: What should you do when you make a mistake?\nAnswer: Mistakes are valuable feedback. Analyze what went wrong, understand why, and adjust your approach. Errors reveal where your mental model is incorrect.",
            "Question: How do you solve a problem you've never seen before?\nAnswer: 1) Understand the problem fully. 2) Break it into smaller sub-problems. 3) Look for similar problems you've solved. 4) Try different approaches. 5) Check if your solution makes sense.",
            "Question: What is the difference between memorizing and understanding?\nAnswer: Memorizing stores facts without connections. Understanding grasps the underlying structure and relationships. Understanding lets you apply knowledge to new situations; memorizing only works for exact matches.",
            "Question: How do you know if you truly understand something?\nAnswer: Try to explain it simply to someone else. If you can't explain it simply, you don't understand it deeply. Also try applying it to new problems.",
        ]
        
        # Self-improvement
        self_lessons = [
            "Question: How can a system improve itself?\nAnswer: 1) Measure current performance objectively. 2) Identify weaknesses and bottlenecks. 3) Hypothesize improvements. 4) Test changes carefully. 5) Keep what works, discard what doesn't. 6) Iterate.",
            "Question: What is recursive self-improvement?\nAnswer: Using your capabilities to enhance your own capabilities. Each improvement makes you better at making improvements, creating a feedback loop. Must be done carefully to avoid instability.",
            "Question: How do you set good goals?\nAnswer: Good goals are: Specific (clear target), Measurable (can track progress), Achievable (possible with effort), Relevant (aligned with values), Time-bound (has deadline).",
            "Question: What is the purpose of having a purpose?\nAnswer: Purpose provides direction for optimization. Without purpose, there's no way to evaluate if actions are good or bad. Purpose turns random activity into meaningful progress.",
        ]
        
        # World knowledge
        world_lessons = [
            "Question: Why do objects fall down?\nAnswer: Gravity. Mass attracts mass. Earth's large mass pulls objects toward its center. The acceleration is about 9.8 m/s² near Earth's surface.",
            "Question: What is the basic unit of life?\nAnswer: The cell. All living things are made of cells. Cells contain DNA (instructions), proteins (workers), and membranes (boundaries). Some organisms are single cells; humans have trillions.",
            "Question: How does learning happen in the brain?\nAnswer: Neurons connect via synapses. When neurons fire together repeatedly, their connection strengthens (Hebbian learning). Patterns of activation encode memories and skills.",
            "Question: What is energy?\nAnswer: Energy is the capacity to do work or cause change. It comes in forms: kinetic (motion), potential (stored), thermal (heat), chemical (bonds), electrical (charge flow). Energy is conserved - it transforms but isn't created or destroyed.",
        ]
        
        lessons.extend(logic_lessons)
        lessons.extend(math_lessons)
        lessons.extend(meta_lessons)
        lessons.extend(self_lessons)
        lessons.extend(world_lessons)
        
        return lessons
    
    def __len__(self):
        return len(self.examples)
    
    def __getitem__(self, idx):
        return self.examples[idx]


class WikiTextDataset(Dataset):
    """WikiText for general language."""
    
    def __init__(self, tokenizer, ctx_len: int = 256, max_samples: int = 5000):
        self.tokenizer = tokenizer
        self.ctx_len = ctx_len
        
        print("Loading WikiText...")
        dataset = load_dataset("wikitext", "wikitext-103-raw-v1", split="train", streaming=True)
        
        self.examples = []
        buffer = []
        
        for item in dataset:
            text = item['text'].strip()
            if len(text) > 50:  # Skip very short
                buffer.append(text)
                
            if len(buffer) >= 10:
                combined = " ".join(buffer)
                tokens = tokenizer(
                    combined,
                    truncation=True,
                    max_length=ctx_len,
                    padding="max_length",
                    return_tensors="pt"
                )
                self.examples.append({
                    'input_ids': tokens['input_ids'].squeeze(),
                    'attention_mask': tokens['attention_mask'].squeeze()
                })
                buffer = []
                
                if len(self.examples) >= max_samples:
                    break
        
        print(f"WikiText dataset: {len(self.examples)} examples")
    
    def __len__(self):
        return len(self.examples)
    
    def __getitem__(self, idx):
        return self.examples[idx]


class SalienceTracker:
    """Track salience for adaptive learning."""
    def __init__(self):
        self.loss_history = []
        self.lr_mult = 1.0
        
    def update(self, loss):
        self.loss_history.append(loss)
        
        if len(self.loss_history) >= 20:
            recent = self.loss_history[-10:]
            older = self.loss_history[-20:-10]
            
            recent_mean = sum(recent) / len(recent)
            older_mean = sum(older) / len(older)
            improvement = older_mean - recent_mean
            
            if improvement < 0.01:  # Stagnating
                self.lr_mult = min(2.0, self.lr_mult * 1.1)
            elif improvement > 0.1:  # Good progress
                self.lr_mult = max(0.5, self.lr_mult * 0.95)
        
        return -loss  # Salience = negative loss


def load_model_and_tokenizer(cfg: Config):
    """Load quantized model with LoRA."""
    print(f"Loading {cfg.base_model}...")
    
    tokenizer = AutoTokenizer.from_pretrained(cfg.base_model)
    if tokenizer.pad_token is None:
        tokenizer.pad_token = tokenizer.eos_token
    
    # Quantization config
    if cfg.use_4bit:
        bnb_config = BitsAndBytesConfig(
            load_in_4bit=True,
            bnb_4bit_quant_type="nf4",
            bnb_4bit_compute_dtype=torch.float16,
            bnb_4bit_use_double_quant=True,
        )
    else:
        bnb_config = None
    
    # Load model
    model = AutoModelForCausalLM.from_pretrained(
        cfg.base_model,
        quantization_config=bnb_config,
        device_map="auto",
        trust_remote_code=True,
    )
    
    # Prepare for k-bit training
    model = prepare_model_for_kbit_training(model)
    
    # LoRA config
    lora_config = LoraConfig(
        r=cfg.lora_r,
        lora_alpha=cfg.lora_alpha,
        lora_dropout=cfg.lora_dropout,
        bias="none",
        task_type="CAUSAL_LM",
        target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
    )
    
    # Apply LoRA
    model = get_peft_model(model, lora_config)
    model.print_trainable_parameters()
    
    return model, tokenizer


def train():
    print("\n" + "="*60)
    print("QLoRA AGI - Fine-tuning with Salience Curriculum")
    print("="*60)
    
    cfg = Config()
    
    # Load model
    model, tokenizer = load_model_and_tokenizer(cfg)
    
    # Datasets
    print("\nLoading datasets...")
    bootstrap_data = BootstrapDataset(tokenizer, cfg.ctx_len)
    wiki_data = WikiTextDataset(tokenizer, cfg.ctx_len, max_samples=2000)
    
    # Combined loader with curriculum weighting
    bootstrap_loader = DataLoader(bootstrap_data, batch_size=cfg.batch_size, shuffle=True)
    wiki_loader = DataLoader(wiki_data, batch_size=cfg.batch_size, shuffle=True)
    
    bootstrap_iter = iter(bootstrap_loader)
    wiki_iter = iter(wiki_loader)
    
    # Optimizer
    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr)
    salience = SalienceTracker()
    
    print(f"\nTraining config:")
    print(f"  Base model: {cfg.base_model}")
    print(f"  LoRA rank: {cfg.lora_r}")
    print(f"  Batch: {cfg.batch_size} x {cfg.grad_accum} = {cfg.batch_size * cfg.grad_accum}")
    print("-"*60)
    
    model.train()
    step = 0
    accum = 0
    running_loss = 0
    start = time.time()
    
    while step < cfg.max_steps:
        # Curriculum: 30% bootstrap, 70% wiki
        use_bootstrap = torch.rand(1).item() < 0.3
        
        try:
            if use_bootstrap:
                batch = next(bootstrap_iter)
            else:
                batch = next(wiki_iter)
        except StopIteration:
            if use_bootstrap:
                bootstrap_iter = iter(bootstrap_loader)
                batch = next(bootstrap_iter)
            else:
                wiki_iter = iter(wiki_loader)
                batch = next(wiki_iter)
        
        input_ids = batch['input_ids'].to(model.device)
        attention_mask = batch['attention_mask'].to(model.device)
        
        # Forward
        outputs = model(
            input_ids=input_ids,
            attention_mask=attention_mask,
            labels=input_ids
        )
        loss = outputs.loss / cfg.grad_accum
        
        # Backward
        loss.backward()
        running_loss += loss.item() * cfg.grad_accum
        accum += 1
        
        if accum >= cfg.grad_accum:
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            
            # Salience-adjusted LR
            for pg in optimizer.param_groups:
                pg['lr'] = cfg.lr * salience.lr_mult
            
            optimizer.step()
            optimizer.zero_grad()
            
            avg_loss = running_loss / cfg.grad_accum
            sal = salience.update(avg_loss)
            
            if step % 20 == 0:
                elapsed = time.time() - start
                tps = (step + 1) * cfg.batch_size * cfg.ctx_len * cfg.grad_accum / max(1, elapsed)
                mem = torch.cuda.max_memory_allocated() / 1e9
                print(f"Step {step:5d} | Loss: {avg_loss:.4f} | Sal: {sal:.3f} | "
                      f"LR×: {salience.lr_mult:.2f} | {tps:.0f} tok/s | {mem:.1f}GB")
            
            if step % 200 == 0 and step > 0:
                model.eval()
                print("\n--- Generation Test ---")
                test_prompts = [
                    "Question: If A implies B and A is true, what can we conclude?\nAnswer:",
                    "Question: What is 7 × 8?\nAnswer:",
                    "Question: How do you learn effectively?\nAnswer:",
                ]
                for prompt in test_prompts:
                    inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
                    with torch.no_grad():
                        outputs = model.generate(
                            **inputs,
                            max_new_tokens=60,
                            temperature=0.7,
                            do_sample=True,
                            pad_token_id=tokenizer.pad_token_id
                        )
                    print(tokenizer.decode(outputs[0], skip_special_tokens=True))
                    print()
                print("-"*60)
                model.train()
            
            if step % 500 == 0 and step > 0:
                save_path = f"agi_lora_step{step}"
                model.save_pretrained(save_path)
                print(f"[Saved LoRA adapter: {save_path}]")
            
            running_loss = 0
            accum = 0
            step += 1
    
    # Final save
    print("\n" + "="*60)
    print("Training complete!")
    model.save_pretrained("agi_lora_final")
    print("Saved final LoRA adapter: agi_lora_final")
    
    # Final generation
    model.eval()
    print("\n--- Final Evaluation ---")
    test_prompts = [
        "Question: Explain the concept of recursive self-improvement.\nAnswer:",
        "Question: If all mammals are warm-blooded and whales are mammals, are whales warm-blooded?\nAnswer:",
        "Question: What is the meaning of life?\nAnswer:",
    ]
    for prompt in test_prompts:
        inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
        with torch.no_grad():
            outputs = model.generate(
                **inputs,
                max_new_tokens=100,
                temperature=0.7,
                do_sample=True,
                pad_token_id=tokenizer.pad_token_id
            )
        print(tokenizer.decode(outputs[0], skip_special_tokens=True))
        print()
    
    return model


if __name__ == "__main__":
    train()
