"""
Resume QLoRA AGI Training

Loads from checkpoint and continues training.
"""

import torch
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
import time
import os
import json
from dataclasses import dataclass
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
from peft import PeftModel, LoraConfig, get_peft_model
from datasets import load_dataset

DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")


@dataclass
class Config:
    base_model: str = "TinyLlama/TinyLlama-1.1B-Chat-v1.0"
    checkpoint_path: str = "agi_lora_checkpoint"
    
    lora_r: int = 64
    lora_alpha: int = 128
    lora_dropout: float = 0.05
    
    batch_size: int = 2
    grad_accum: int = 8
    lr: float = 1e-4  # Lower LR for fine-tuning continuation
    max_steps: int = 2000  # Additional steps
    ctx_len: int = 256
    
    use_4bit: bool = True


class BootstrapDataset(Dataset):
    """Bootstrap curriculum."""
    
    def __init__(self, tokenizer, ctx_len: int = 256):
        self.tokenizer = tokenizer
        self.ctx_len = ctx_len
        
        self.lessons = [
            # Logic
            "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. 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.",
            
            # Math - more examples
            "Question: What is 7 × 8?\nAnswer: 7 × 8 = 56. Seven times eight equals fifty-six.",
            "Question: What is 6 × 9?\nAnswer: 6 × 9 = 54. Six times nine equals fifty-four.",
            "Question: What is 12 × 12?\nAnswer: 12 × 12 = 144. Twelve times twelve equals one hundred forty-four.",
            "Question: What is 15 + 27?\nAnswer: 15 + 27 = 42. Fifteen plus twenty-seven equals forty-two.",
            "Question: What is 100 - 37?\nAnswer: 100 - 37 = 63. One hundred minus thirty-seven equals sixty-three.",
            "Question: Solve for x: 2x + 4 = 10\nAnswer: 2x + 4 = 10. Subtract 4: 2x = 6. Divide by 2: x = 3.",
            
            # Meta-learning
            "Question: How do you learn 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.",
            
            # Self-improvement
            "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 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.",
            
            # World knowledge
            "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 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.",
            
            # Reasoning
            "Question: If all cats have tails, and Whiskers is a cat, does Whiskers have a tail?\nAnswer: Yes. Since all cats have tails, and Whiskers is a cat, therefore Whiskers has a tail. This is deductive reasoning.",
            "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.",
        ]
        
        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 __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 = 3000):
        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:
                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:
    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:
                self.lr_mult = min(2.0, self.lr_mult * 1.1)
            elif improvement > 0.1:
                self.lr_mult = max(0.5, self.lr_mult * 0.95)
        return -loss


def load_model(cfg: Config, resume: bool = True):
    """Load model, optionally from checkpoint."""
    print(f"Loading base model: {cfg.base_model}")
    
    tokenizer = AutoTokenizer.from_pretrained(cfg.base_model)
    if tokenizer.pad_token is None:
        tokenizer.pad_token = tokenizer.eos_token
    
    bnb_config = BitsAndBytesConfig(
        load_in_4bit=True,
        bnb_4bit_quant_type="nf4",
        bnb_4bit_compute_dtype=torch.float16,
        bnb_4bit_use_double_quant=True,
    )
    
    model = AutoModelForCausalLM.from_pretrained(
        cfg.base_model,
        quantization_config=bnb_config,
        device_map="auto",
        trust_remote_code=True,
    )
    
    if resume and os.path.exists(cfg.checkpoint_path):
        print(f"Loading checkpoint: {cfg.checkpoint_path}")
        model = PeftModel.from_pretrained(model, cfg.checkpoint_path)
        
        # Load metadata
        meta_path = f"{cfg.checkpoint_path}/training_meta.json"
        if os.path.exists(meta_path):
            with open(meta_path) as f:
                meta = json.load(f)
            start_step = meta.get('step', 0)
            print(f"Resuming from step {start_step}")
        else:
            start_step = 500  # Default
    else:
        print("Starting fresh training")
        from peft import prepare_model_for_kbit_training
        model = prepare_model_for_kbit_training(model)
        
        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"],
        )
        model = get_peft_model(model, lora_config)
        start_step = 0
    
    model.print_trainable_parameters()
    return model, tokenizer, start_step


def train():
    print("\n" + "="*60)
    print("RESUME QLoRA AGI Training")
    print("="*60)
    
    cfg = Config()
    model, tokenizer, start_step = load_model(cfg, resume=True)
    
    print("\nLoading datasets...")
    bootstrap_data = BootstrapDataset(tokenizer, cfg.ctx_len)
    wiki_data = WikiTextDataset(tokenizer, cfg.ctx_len, max_samples=3000)
    
    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 = torch.optim.AdamW(model.parameters(), lr=cfg.lr)
    salience = SalienceTracker()
    
    print(f"\nResuming from step {start_step}")
    print(f"Training for {cfg.max_steps} more steps")
    print(f"LR: {cfg.lr}")
    print("-"*60)
    
    model.train()
    step = start_step
    accum = 0
    running_loss = 0
    start = time.time()
    
    while step < start_step + cfg.max_steps:
        use_bootstrap = torch.rand(1).item() < 0.4  # 40% bootstrap
        
        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)
        
        outputs = model(input_ids=input_ids, attention_mask=attention_mask, labels=input_ids)
        loss = outputs.loss / cfg.grad_accum
        
        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)
            
            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
                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} | {mem:.1f}GB")
            
            if step % 200 == 0:
                model.eval()
                print("\n--- Test ---")
                test_prompts = [
                    "Question: What is 7 × 8?\nAnswer:",
                    "Question: If all dogs bark and Fido is a dog, does Fido bark?\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 > start_step:
                save_path = f"agi_lora_step{step}"
                model.save_pretrained(save_path)
                
                meta = {
                    'base_model': cfg.base_model,
                    'step': step,
                    'lora_r': cfg.lora_r,
                    'lora_alpha': cfg.lora_alpha,
                }
                with open(f"{save_path}/training_meta.json", 'w') as f:
                    json.dump(meta, f, indent=2)
                
                # Also update main checkpoint
                model.save_pretrained("agi_lora_checkpoint")
                with open("agi_lora_checkpoint/training_meta.json", 'w') as f:
                    json.dump(meta, f, indent=2)
                
                print(f"[Saved checkpoint at step {step}]")
            
            running_loss = 0
            accum = 0
            step += 1
    
    # Final save
    print("\n" + "="*60)
    print("Training complete!")
    
    final_path = "agi_lora_final"
    model.save_pretrained(final_path)
    meta = {'base_model': cfg.base_model, 'step': step, 'lora_r': cfg.lora_r, 'lora_alpha': cfg.lora_alpha}
    with open(f"{final_path}/training_meta.json", 'w') as f:
        json.dump(meta, f, indent=2)
    
    model.save_pretrained("agi_lora_checkpoint")
    with open("agi_lora_checkpoint/training_meta.json", 'w') as f:
        json.dump(meta, f, indent=2)
    
    print(f"Saved final model to {final_path}/")
    return model


if __name__ == "__main__":
    train()
