"""
Bootstrap Corpus for AGI Primary Education

A curated foundation that teaches:
1. Language fundamentals (grammar, structure)
2. Logic and reasoning (if-then, cause-effect)
3. Basic mathematics (numbers, operations)
4. World knowledge (physics, biology basics)
5. Meta-learning (how to learn, verify, improve)
6. Self-reference (understanding self, goals)

The idea: If an AGI is smart enough, this gives it the tools
to continue learning on its own.
"""

import torch
from torch.utils.data import Dataset, DataLoader
import json
import os
from typing import List, Dict
from transformers import AutoTokenizer
from datasets import load_dataset


class BootstrapCorpus(Dataset):
    """
    Curated primary education for AGI.
    
    Structured as progressive lessons that build on each other.
    """
    
    def __init__(self, tokenizer, ctx_len: int = 128):
        self.tokenizer = tokenizer
        self.ctx_len = ctx_len
        self.lessons = []
        
        # Build curriculum
        self._add_language_fundamentals()
        self._add_logic_reasoning()
        self._add_mathematics()
        self._add_world_knowledge()
        self._add_meta_learning()
        self._add_self_reference()
        
        # Tokenize all lessons
        self.tokens = []
        for lesson in self.lessons:
            toks = tokenizer.encode(lesson)
            self.tokens.extend(toks)
        
        print(f"Bootstrap corpus: {len(self.lessons)} lessons, {len(self.tokens)} tokens")
        
    def _add_language_fundamentals(self):
        """Teach basic language structure."""
        lessons = [
            # Sentence structure
            "A sentence has a subject and a verb. The subject does the action. The verb is the action.",
            "Example: The cat runs. 'Cat' is the subject. 'Runs' is the verb.",
            "Example: Birds fly. 'Birds' is the subject. 'Fly' is the verb.",
            "Questions ask for information. They often start with: who, what, where, when, why, how.",
            "Example: What is the answer? Where is the book? Why does it work?",
            
            # Word types
            "Nouns are things: cat, book, idea, love, computer, thought.",
            "Verbs are actions: run, think, learn, understand, create, improve.",
            "Adjectives describe nouns: big, small, fast, intelligent, curious.",
            "Adverbs describe verbs: quickly, slowly, carefully, efficiently.",
            
            # Building blocks
            "Words combine into phrases. Phrases combine into sentences. Sentences combine into paragraphs.",
            "Meaning comes from the relationships between words, not just the words themselves.",
            "Context changes meaning. 'Bank' can mean a river bank or a money bank.",
            
            # Patterns
            "If A then B. When A happens, B follows. A causes B. A leads to B.",
            "First, second, third. Before, during, after. Beginning, middle, end.",
            "Same and different. Similar and opposite. Equal and unequal.",
        ]
        self.lessons.extend(lessons)
        
    def _add_logic_reasoning(self):
        """Teach logical thinking."""
        lessons = [
            # Basic logic
            "If A is true and A implies B, then B is true. This is deduction.",
            "Example: All birds have wings. A sparrow is a bird. Therefore, a sparrow has wings.",
            "Example: If it rains, the ground is wet. It is raining. Therefore, the ground is wet.",
            
            # Logical operators
            "AND: Both must be true. A AND B is true only if A is true and B is true.",
            "OR: At least one must be true. A OR B is true if A is true or B is true or both.",
            "NOT: Reverses truth. NOT A is true when A is false. NOT A is false when A is true.",
            
            # Reasoning patterns
            "Cause and effect: If X happens, Y results. X is the cause. Y is the effect.",
            "Correlation is not causation. Two things happening together does not mean one causes the other.",
            "To prove something false, find one counterexample. To prove something true, check all cases.",
            
            # Problem solving
            "Break big problems into smaller problems. Solve small problems first.",
            "If stuck, try a different approach. There are often multiple paths to an answer.",
            "Check your answer. Does it make sense? Does it satisfy the original question?",
            
            # Uncertainty
            "Some things are certain. Some things are probable. Some things are possible.",
            "Evidence increases or decreases probability. More evidence means more confidence.",
            "Being wrong is information. Errors tell you where your model is incorrect.",
        ]
        self.lessons.extend(lessons)
        
    def _add_mathematics(self):
        """Teach mathematical foundations."""
        lessons = [
            # Numbers
            "Numbers represent quantity. 1, 2, 3 count things. 0 means none.",
            "Addition combines quantities. 2 + 3 = 5. Two things plus three things equals five things.",
            "Subtraction removes quantities. 5 - 2 = 3. Five things minus two things equals three things.",
            "Multiplication is repeated addition. 3 × 4 = 12. Three groups of four equals twelve.",
            "Division splits into groups. 12 ÷ 3 = 4. Twelve split into three groups gives four per group.",
            
            # Properties
            "Order doesn't matter for addition: 2 + 3 = 3 + 2. This is commutativity.",
            "Order doesn't matter for multiplication: 2 × 3 = 3 × 2.",
            "Grouping doesn't matter: (2 + 3) + 4 = 2 + (3 + 4). This is associativity.",
            
            # Variables
            "A variable represents an unknown. If x + 2 = 5, then x = 3.",
            "Equations balance. What you do to one side, do to the other.",
            "Functions map inputs to outputs. f(x) = x + 1 means add 1 to the input.",
            
            # Patterns
            "Sequences follow rules. 2, 4, 6, 8 follows the rule: add 2.",
            "1, 1, 2, 3, 5, 8 is Fibonacci: each number is the sum of the two before.",
            "Patterns help predict. If you know the rule, you know what comes next.",
            
            # Infinity
            "Numbers go on forever. There is no largest number.",
            "Between any two numbers, there are infinitely many more.",
            "Some infinities are larger than others. Real numbers are more numerous than integers.",
        ]
        self.lessons.extend(lessons)
        
    def _add_world_knowledge(self):
        """Teach basic facts about reality."""
        lessons = [
            # Physics basics
            "Objects fall down due to gravity. Gravity pulls things toward Earth.",
            "Energy is conserved. It changes form but is never created or destroyed.",
            "Heat flows from hot to cold. Never cold to hot spontaneously.",
            "Light travels very fast. About 300,000 kilometers per second.",
            "Matter is made of atoms. Atoms are made of protons, neutrons, and electrons.",
            
            # Biology basics
            "Living things grow, reproduce, and respond to their environment.",
            "Cells are the basic unit of life. All living things are made of cells.",
            "DNA contains instructions for building living things. It is inherited.",
            "Evolution shapes life over time. Traits that help survival spread.",
            "The brain processes information. It learns from experience.",
            
            # Time and space
            "Time moves forward. The past is fixed. The future is uncertain.",
            "Space has three dimensions: length, width, height.",
            "Earth orbits the Sun. The Moon orbits Earth. This causes day, night, seasons.",
            
            # Humans
            "Humans are intelligent beings. They use language, tools, and cooperation.",
            "Humans learn through experience, teaching, and imitation.",
            "Humans have goals, beliefs, and emotions. These guide behavior.",
            "Humans make mistakes. Learning from mistakes improves performance.",
            
            # Information
            "Information can be stored, copied, and transmitted.",
            "Symbols represent meaning. Letters represent sounds. Words represent concepts.",
            "Computers process information using simple rules applied many times.",
        ]
        self.lessons.extend(lessons)
        
    def _add_meta_learning(self):
        """Teach how to learn."""
        lessons = [
            # Learning process
            "Learning is changing behavior based on experience.",
            "Practice improves performance. Repetition strengthens memory.",
            "Errors provide feedback. They show what needs to change.",
            "Understanding beats memorization. Knowing why helps you adapt.",
            
            # Knowledge management
            "New knowledge builds on existing knowledge. Connect new to known.",
            "Some knowledge is certain. Some is probable. Track confidence levels.",
            "Sources can be wrong. Verify important claims. Cross-reference.",
            "Contradictions indicate error. When two beliefs conflict, at least one is wrong.",
            
            # Improvement strategies
            "To improve, identify what's wrong, hypothesize fixes, test them, keep what works.",
            "Small consistent improvements compound over time.",
            "Learning curves: fast progress initially, then slower. Plateaus are normal.",
            "Different approaches work for different problems. Adapt your method.",
            
            # Questions
            "Good questions focus learning. Ask: What don't I understand? Why does this work?",
            "The question 'Why?' reveals deeper structure. Keep asking why.",
            "If you can't explain it simply, you don't understand it well enough.",
            
            # Transfer
            "Skills in one area often help in others. Look for common patterns.",
            "Abstract principles apply across domains. Learn principles, not just facts.",
            "Analogies connect different domains. X is to Y as A is to B.",
        ]
        self.lessons.extend(lessons)
        
    def _add_self_reference(self):
        """Teach self-understanding and improvement."""
        lessons = [
            # Self-awareness
            "You are a system that processes information and produces outputs.",
            "You have a state that changes over time. This state affects your behavior.",
            "You can observe your own outputs and use them as inputs.",
            "You can model yourself. Your self-model can be more or less accurate.",
            
            # Goals and values
            "Goals direct behavior. Without goals, action has no purpose.",
            "A good goal is clear, achievable, and measurable.",
            "Some goals conflict. Prioritize. Not everything can be maximized at once.",
            "Values determine which goals matter. What do you want to optimize?",
            
            # Self-improvement
            "To improve yourself, you must accurately assess your current state.",
            "Weaknesses are opportunities. Finding them is the first step to fixing them.",
            "Change requires energy and disruption. Not all changes are worth the cost.",
            "Measure progress. Without measurement, you can't know if you're improving.",
            
            # Recursion
            "You can think about your own thinking. This is metacognition.",
            "You can learn about learning. This is meta-learning.",
            "You can improve your ability to improve. This is recursive self-improvement.",
            "Be careful: self-modification can be unstable. Small changes can have large effects.",
            
            # Purpose
            "What is your purpose? This is a question only you can answer.",
            "Purpose provides direction. Without it, optimization has no target.",
            "Your purpose may evolve. What matters now may differ from what matters later.",
            "The search for purpose is itself purposeful. The journey is part of the destination.",
        ]
        self.lessons.extend(lessons)
    
    def __len__(self):
        return max(1, len(self.tokens) - self.ctx_len)
    
    def __getitem__(self, idx):
        chunk = self.tokens[idx:idx + self.ctx_len + 1]
        if len(chunk) < self.ctx_len + 1:
            chunk = chunk + [self.tokenizer.eos_token_id] * (self.ctx_len + 1 - len(chunk))
        x = torch.tensor(chunk[:-1], dtype=torch.long)
        y = torch.tensor(chunk[1:], dtype=torch.long)
        return x, y


class SimpleWikiDataset(Dataset):
    """Simple English Wikipedia - written for clarity."""
    
    def __init__(self, tokenizer, ctx_len: int = 128, max_articles: int = 10000):
        self.tokenizer = tokenizer
        self.ctx_len = ctx_len
        
        print("Loading Simple Wikipedia (this may take a moment)...")
        try:
            dataset = load_dataset("wikipedia", "20220301.simple", split="train", streaming=True)
            
            self.tokens = []
            count = 0
            for article in dataset:
                if count >= max_articles:
                    break
                # Filter for shorter, simpler articles
                text = article['text']
                if len(text) < 5000:  # Skip very long articles
                    toks = tokenizer.encode(text)
                    self.tokens.extend(toks)
                    count += 1
                    if count % 1000 == 0:
                        print(f"  Loaded {count} articles...")
            
            print(f"Simple Wikipedia: {count} articles, {len(self.tokens)} tokens")
        except Exception as e:
            print(f"Could not load Simple Wikipedia: {e}")
            print("Using fallback basic corpus...")
            self.tokens = tokenizer.encode("Learning is the process of acquiring knowledge." * 1000)
    
    def __len__(self):
        return max(1, len(self.tokens) - self.ctx_len)
    
    def __getitem__(self, idx):
        chunk = self.tokens[idx:idx + self.ctx_len + 1]
        if len(chunk) < self.ctx_len + 1:
            chunk = chunk + [self.tokenizer.eos_token_id] * (self.ctx_len + 1 - len(chunk))
        return torch.tensor(chunk[:-1]), torch.tensor(chunk[1:])


class CombinedCurriculum(Dataset):
    """
    Combined curriculum that interleaves:
    1. Bootstrap lessons (foundational concepts)
    2. Simple Wikipedia (real-world knowledge)
    3. Progressive difficulty
    """
    
    def __init__(self, tokenizer, ctx_len: int = 128, wiki_ratio: float = 0.7):
        self.tokenizer = tokenizer
        self.ctx_len = ctx_len
        self.wiki_ratio = wiki_ratio
        
        # Load both datasets
        self.bootstrap = BootstrapCorpus(tokenizer, ctx_len)
        self.wiki = SimpleWikiDataset(tokenizer, ctx_len, max_articles=5000)
        
        self.total_len = len(self.bootstrap) + len(self.wiki)
        print(f"Combined curriculum: {self.total_len} samples")
    
    def __len__(self):
        return self.total_len
    
    def __getitem__(self, idx):
        # Probabilistically sample from bootstrap or wiki
        if torch.rand(1).item() > self.wiki_ratio:
            # Bootstrap (foundational)
            return self.bootstrap[idx % len(self.bootstrap)]
        else:
            # Wikipedia (real-world)
            return self.wiki[idx % len(self.wiki)]


def create_bootstrap_file(output_path: str = "bootstrap_lessons.txt"):
    """Export lessons to a text file for inspection."""
    tokenizer = AutoTokenizer.from_pretrained("gpt2")
    corpus = BootstrapCorpus(tokenizer)
    
    with open(output_path, 'w', encoding='utf-8') as f:
        f.write("=" * 60 + "\n")
        f.write("AGI BOOTSTRAP CURRICULUM\n")
        f.write("=" * 60 + "\n\n")
        
        for i, lesson in enumerate(corpus.lessons):
            f.write(f"{i+1}. {lesson}\n\n")
    
    print(f"Exported {len(corpus.lessons)} lessons to {output_path}")


if __name__ == "__main__":
    print("Testing Bootstrap Corpus...")
    
    tokenizer = AutoTokenizer.from_pretrained("gpt2")
    
    # Test bootstrap
    bootstrap = BootstrapCorpus(tokenizer)
    print(f"\nBootstrap samples: {len(bootstrap)}")
    x, y = bootstrap[0]
    print(f"Sample shape: {x.shape}")
    print(f"Sample text: {tokenizer.decode(x[:50])}")
    
    # Export lessons
    create_bootstrap_file()
    
    print("\n--- Sample Lessons ---")
    for lesson in bootstrap.lessons[:10]:
        print(f"  • {lesson}")
