"""
Unified Character-Level Predictor
Conv2D processes 28x28 letter images → LSTM encodes sequence → Predict next character
On-the-fly generation during training
"""

import os
import re
import random
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from collections import defaultdict
from pathlib import Path
from PIL import Image, ImageDraw, ImageFont

# ═══════════════════════════════════════════════════════════════════
# CONFIGURATION
# ═══════════════════════════════════════════════════════════════════

class Config:
    book_folder = "./books"
    
    # Image (28x28 letter rendering)
    img_size = 28
    font_size = 20
    
    # Model
    embed_dim = 128
    hidden_dim = 256
    num_layers = 2
    
    # Sequence
    seq_len = 8  # Number of letter images in context
    
    # Training
    batch_size = 64
    epochs = 20
    lr = 0.001
    
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

cfg = Config()

# ═══════════════════════════════════════════════════════════════════
# TEXT LOADING
# ═══════════════════════════════════════════════════════════════════

def load_text_files(folder_path):
    """Load raw text from markdown files."""
    text = ""
    folder = Path(folder_path)
    
    if not folder.exists():
        print(f"⚠️ Folder '{folder_path}' not found. Using sample text.")
        return sample_text()
    
    for md_file in folder.glob("*.md"):
        print(f"📖 Loading: {md_file.name}")
        with open(md_file, 'r', encoding='utf-8', errors='ignore') as f:
            content = f.read()
            # Clean markdown but keep structure
            content = re.sub(r'#+', ' ', content)
            content = re.sub(r'\[([^\]]+)\]\([^\)]+\)', r'\1', content)
            content = re.sub(r'[*_`~]', '', content)
            text += content + " "
    
    if len(text.strip()) == 0:
        return sample_text()
    
    # Keep only lowercase letters and spaces for simplicity
    text = re.sub(r'[^a-z\s]', ' ', text.lower())
    text = re.sub(r'\s+', ' ', text).strip()
    
    print(f"✅ Loaded {len(text)} characters")
    return text

def sample_text():
    return "in the beginning god created the heavens and the earth now the earth was formless and empty darkness was over the surface of the deep"

# ═══════════════════════════════════════════════════════════════════
# CHARACTER TO IMAGE RENDERING
# ═══════════════════════════════════════════════════════════════════

def render_letter(char, size=28, font_size=20):
    """Render a single letter as 28x28 grayscale image."""
    img = Image.new('L', (size, size), color=255)
    draw = ImageDraw.Draw(img)
    
    try:
        font = ImageFont.truetype("/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf", font_size)
    except:
        try:
            font = ImageFont.truetype("arial.ttf", font_size)
        except:
            font = ImageFont.load_default()
    
    bbox = draw.textbbox((0, 0), char, font=font)
    text_width = bbox[2] - bbox[0]
    text_height = bbox[3] - bbox[1]
    
    x = (size - text_width) // 2
    y = (size - text_height) // 2
    
    draw.text((x, y), char, fill=0, font=font)
    
    img_array = np.array(img, dtype=np.float32) / 255.0
    img_array = 1.0 - img_array  # Invert: black text on white
    
    return img_array

def render_word(word, size=28, font_size=20):
    """Render a word as 28x28 image (for testing)."""
    return render_letter(word[0] if word else 'a', size, font_size)

# ═══════════════════════════════════════════════════════════════════
# ON-THE-FLY LETTER IMAGE GENERATOR
# ═══════════════════════════════════════════════════════════════════

class LetterImageGenerator:
    """
    Generates letter images on-the-fly for training.
    No precomputed dataset - yields batches during training.
    """
    
    def __init__(self, text, seq_len=8, stride=1):
        self.text = text
        self.seq_len = seq_len
        self.stride = stride
        
        # Pre-render all needed letters (cache)
        self.letter_cache = {}
        chars = set(text)
        for char in chars:
            if char.isalpha():
                self.letter_cache[char] = render_letter(char)
        
        # Get valid starting positions (need seq_len + 1 chars for target)
        self.valid_positions = [
            i for i in range(0, len(text) - seq_len, stride)
            if all(text[i+j].isalpha() for j in range(seq_len + 1))
        ]
        
        if len(self.valid_positions) == 0:
            # Fallback: use all positions
            self.valid_positions = list(range(0, len(text) - seq_len, 1))
        
        print(f"🖼️ Letter image generator: {len(self.valid_positions)} valid sequences")
    
    def get_batch(self, batch_size):
        """Get a batch of letter image sequences."""
        # Random sample positions
        positions = random.sample(self.valid_positions, min(batch_size, len(self.valid_positions)))
        
        # Batch tensors
        batch_images = []  # [batch, seq_len, 1, 28, 28]
        batch_targets = []  # [batch]
        
        for start_pos in positions:
            # Get sequence of letters
            input_letters = self.text[start_pos:start_pos + self.seq_len]
            target_letter = self.text[start_pos + self.seq_len]
            
            # Render letter images
            images = []
            for letter in input_letters:
                if letter in self.letter_cache:
                    images.append(self.letter_cache[letter])
                else:
                    images.append(np.zeros((28, 28), dtype=np.float32))
            
            # Stack into sequence
            images = np.stack(images)  # [seq_len, 28, 28]
            batch_images.append(images)
            
            # Target: ord('a') = 0, ord('b') = 1, etc.
            batch_targets.append(ord(target_letter) - ord('a'))
        
        # Convert to tensors
        images_tensor = torch.FloatTensor(np.stack(batch_images)).unsqueeze(2)  # [batch, seq_len, 1, 28, 28]
        targets_tensor = torch.LongTensor(batch_targets)
        
        return images_tensor, targets_tensor
    
    def get_epoch_batches(self, batch_size, num_batches=None):
        """Yield all batches for one epoch."""
        if num_batches is None:
            num_batches = len(self.valid_positions) // batch_size
        
        for _ in range(num_batches):
            yield self.get_batch(batch_size)

# ═══════════════════════════════════════════════════════════════════
# UNIFIED MODEL: Conv2D Encoder + LSTM Sequence
# ═══════════════════════════════════════════════════════════════════

class LetterImageEncoder(nn.Module):
    """
    Conv2D encoder for 28x28 letter images.
    Shared weights for all letters in the sequence.
    """
    
    def __init__(self, embed_dim=128):
        super().__init__()
        
        # Convolutional feature extraction (same as Word-MLP)
        self.conv = nn.Sequential(
            # Block 1: 28 -> 14
            nn.Conv2d(1, 32, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),
            
            # Block 2: 14 -> 7
            nn.Conv2d(32, 64, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),
            
            # Block 3: 7 -> 3
            nn.Conv2d(64, 128, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),
        )
        
        self.flatten = nn.Flatten()
        # 128 * 3 * 3 = 1152
        self.projection = nn.Linear(128 * 3 * 3, embed_dim)
    
    def forward(self, x):
        """
        x: [batch, seq_len, 1, 28, 28]
        Returns: [batch, seq_len, embed_dim]
        """
        batch_size, seq_len = x.size(0), x.size(1)
        
        # Merge batch and seq for parallel conv processing
        x = x.view(batch_size * seq_len, 1, 28, 28)
        
        # Conv
        x = self.conv(x)
        x = self.flatten(x)
        x = self.projection(x)
        
        # Reshape back
        x = x.view(batch_size, seq_len, -1)  # [batch, seq_len, embed_dim]
        
        return x


class LetterSequencePredictor(nn.Module):
    """
    Unified model:
    1. Conv2D encodes 28x28 letter images → embeddings
    2. LSTM encodes sequence trajectory
    3. Predicts next character
    """
    
    def __init__(self, num_classes=26, embed_dim=128, hidden_dim=256, num_layers=2):
        super().__init__()
        
        self.num_classes = num_classes
        self.hidden_dim = hidden_dim
        
        # Letter image encoder (Conv2D)
        self.encoder = LetterImageEncoder(embed_dim)
        
        # Sequence encoder (LSTM)
        self.lstm = nn.LSTM(
            input_size=embed_dim,
            hidden_size=hidden_dim,
            num_layers=num_layers,
            batch_first=True,
            dropout=0.2 if num_layers > 1 else 0
        )
        
        # Prediction head
        self.predictor = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim),
            nn.ReLU(),
            nn.Dropout(0.1),
            nn.Linear(hidden_dim, num_classes)
        )
        
        # Entropy tracking for ODE-CCT
        self.register_buffer('entropy_history', torch.zeros(100))
        self.entropy_idx = 0
    
    def forward(self, images, return_entropy=False):
        """
        images: [batch, seq_len, 1, 28, 28] - letter images
        Returns: logits [batch, num_classes], entropy (optional)
        """
        # Step 1: Encode letter images with Conv2D
        embeddings = self.encoder(images)  # [batch, seq_len, embed_dim]
        
        # Step 2: LSTM encodes sequence trajectory
        lstm_out, (hidden, cell) = self.lstm(embeddings)
        
        # Step 3: Use last hidden state for prediction
        last_hidden = hidden[-1]  # [batch, hidden_dim]
        
        # Step 4: Predict next character
        logits = self.predictor(last_hidden)  # [batch, num_classes]
        
        if return_entropy:
            probs = F.softmax(logits, dim=1)
            entropy = -(probs * torch.log(probs + 1e-8)).sum(dim=1).mean()
            
            # Track entropy for ODE-CCT
            idx = self.entropy_idx % 100
            self.entropy_history[idx] = entropy
            self.entropy_idx += 1
            
            return logits, entropy.item()
        
        return logits
    
    def get_trajectory(self, images):
        """Get LSTM hidden state as trajectory representation."""
        embeddings = self.encoder(images)
        _, (hidden, cell) = self.lstm(embeddings)
        return hidden, cell


class CCTLetterPredictor(LetterSequencePredictor):
    """
    Conditional Collapse Theory enhanced predictor.
    - Detects periodic patterns in letter sequences
    - Caches common n-grams
    - Adaptive compute allocation
    """
    
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        
        # N-gram cache for periodic patterns
        self.ngram_cache = {}
        self.ngram_order = 3
        
        # Periodic phrase cache
        self.phrase_cache = {}
        
    def check_cache(self, images):
        """Check if sequence matches cached pattern."""
        # Use first letter as quick key
        first_letter_idx = images[0, 0, 0, 14, 14].item()  # Simplified
        seq_key = tuple(images[:, 0, 0, 14, 14].tolist())
        
        return self.phrase_cache.get(seq_key)
    
    def update_cache(self, images, prediction):
        """Cache prediction for sequence."""
        seq_key = tuple(images[:, 0, 0, 14, 14].tolist())
        self.phrase_cache[seq_key] = prediction
    
    def forward_with_cct(self, images, use_cache=True):
        """
        Forward with Conditional Collapse Theory.
        """
        if use_cache:
            cached = self.check_cache(images[0])
            if cached is not None and random.random() < 0.3:
                return cached
        
        logits = self.forward(images)
        pred = logits.argmax(dim=1)
        
        if use_cache:
            self.update_cache(images[0], pred[0].item())
        
        return pred

# ═══════════════════════════════════════════════════════════════════
# TRAINING (ON-THE-FLY)
# ═══════════════════════════════════════════════════════════════════

def train_on_the_fly(model, text, epochs=20, lr=0.001):
    """
    Train model with on-the-fly letter image generation.
    """
    
    optimizer = torch.optim.Adam(model.parameters(), lr=lr)
    criterion = nn.CrossEntropyLoss()
    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=3, factor=0.5)
    
    model.to(cfg.device)
    
    # Create streaming generator
    generator = LetterImageGenerator(text, seq_len=cfg.seq_len, stride=1)
    
    for epoch in range(epochs):
        model.train()
        
        epoch_loss = 0.0
        epoch_correct = 0
        epoch_total = 0
        
        entropy_sum = 0.0
        batches_done = 0
        
        num_batches = len(generator.valid_positions) // cfg.batch_size
        
        for batch_idx in range(num_batches):
            # Generate batch on-the-fly (letter images)
            images, targets = generator.get_batch(cfg.batch_size)
            
            images = images.to(cfg.device)
            targets = targets.to(cfg.device)
            
            optimizer.zero_grad()
            
            logits, entropy = model(images, return_entropy=True)
            loss = criterion(logits, targets)
            
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
            optimizer.step()
            
            epoch_loss += loss.item()
            preds = logits.argmax(dim=1)
            epoch_correct += (preds == targets).sum().item()
            epoch_total += targets.size(0)
            
            entropy_sum += entropy
            batches_done += 1
            
            if (batch_idx + 1) % 500 == 0:
                print(f"  Batch {batch_idx+1}/{num_batches} | "
                      f"Loss: {epoch_loss/batches_done:.4f} | "
                      f"Acc: {100*epoch_correct/epoch_total:.2f}%")
        
        train_acc = 100 * epoch_correct / epoch_total
        avg_loss = epoch_loss / batches_done
        avg_entropy = entropy_sum / batches_done
        perplexity = np.exp(avg_entropy)
        
        scheduler.step(avg_loss)
        
        print(f"Epoch {epoch+1}/{epochs} | "
              f"Loss: {avg_loss:.4f} | Acc: {train_acc:.2f}% | "
              f"Entropy: {avg_entropy:.4f} | PPL: {perplexity:.2f}")
    
    return model

# ═══════════════════════════════════════════════════════════════════
# TESTING
# ═══════════════════════════════════════════════════════════════════

def test_model(model, text):
    """Test model on held-out text."""
    
    model.eval()
    criterion = nn.CrossEntropyLoss()
    
    # Use last 10% for testing
    test_text = text[int(len(text) * 0.9):]
    
    test_gen = LetterImageGenerator(test_text, seq_len=cfg.seq_len, stride=5)
    
    correct = 0
    total = 0
    entropies = []
    
    with torch.no_grad():
        for images, targets in test_gen.get_epoch_batches(cfg.batch_size, num_batches=50):
            images = images.to(cfg.device)
            targets = targets.to(cfg.device)
            
            logits, entropy = model(images, return_entropy=True)
            loss = criterion(logits, targets)
            
            preds = logits.argmax(dim=1)
            correct += (preds == targets).sum().item()
            total += targets.size(0)
            
            probs = F.softmax(logits, dim=1)
            max_probs = probs.max(dim=1)[0]
            entropies.extend(max_probs.tolist())
    
    accuracy = 100 * correct / total
    avg_conf = np.mean(entropies)
    
    print("\n" + "="*60)
    print("📊 LETTER IMAGE SEQUENCE PREDICTION RESULTS")
    print("="*60)
    print(f"Test Accuracy: {accuracy:.2f}% ({correct}/{total})")
    print(f"Average Confidence: {avg_conf:.4f}")
    
    # Per-letter accuracy
    print("\n🔤 Per-Letter Accuracy Analysis:")
    letter_correct = defaultdict(int)
    letter_total = defaultdict(int)
    
    with torch.no_grad():
        for images, targets in test_gen.get_epoch_batches(cfg.batch_size, num_batches=50):
            images = images.to(cfg.device)
            targets = targets.to(cfg.device)
            
            logits, _ = model(images)
            preds = logits.argmax(dim=1)
            
            for pred, true in zip(preds.cpu(), targets.cpu()):
                letter = chr(true.item() + ord('a'))
                letter_total[letter] += 1
                if pred == true:
                    letter_correct[letter] += 1
    
    letter_acc = {l: letter_correct.get(l, 0) / max(letter_total.get(l, 1), 1) for l in letter_total}
    sorted_letters = sorted(letter_acc.items(), key=lambda x: -x[1])
    
    print(f"  Best predicted:")
    for letter, acc in sorted_letters[:5]:
        print(f"    '{letter}': {acc*100:.1f}%")
    
    print(f"  Worst predicted:")
    for letter, acc in sorted_letters[-5:]:
        print(f"    '{letter}': {acc*100:.1f}%")
    
    # Sample predictions
    print("\n🔍 Sample Predictions:")
    model.eval()
    sample_count = 0
    
    with torch.no_grad():
        for images, targets in test_gen.get_epoch_batches(cfg.batch_size, num_batches=5):
            for i in range(min(3, len(targets))):
                # Get input letters
                input_letters = []
                for j in range(cfg.seq_len):
                    idx = images[i, j, 0, 14, 14].item()
                    # Find which letter has this position
                    letter = '?'  # Simplified
                    input_letters.append(letter)
                
                true_letter = chr(targets[i].item() + ord('a'))
                
                logits, _ = model(images[i:i+1], return_entropy=True)
                pred_letter = chr(logits.argmax(dim=1).item() + ord('a'))
                
                status = "✅" if pred_letter == true_letter else "❌"
                print(f"  {status} Input: {input_letters[-4:]} → | True: '{true_letter}' | Pred: '{pred_letter}'")
                
                sample_count += 1
                if sample_count >= 8:
                    break
            if sample_count >= 8:
                break
    
    print("="*60)
    
    return accuracy

# ═══════════════════════════════════════════════════════════════════
# TEXT GENERATION
# ═══════════════════════════════════════════════════════════════════

def generate_text(model, seed_text, length=50, temperature=1.0):
    """Generate text given a seed."""
    
    model.eval()
    
    # Render seed letters
    letters = list(seed_text.lower().replace(' ', ''))[-cfg.seq_len:]
    
    while len(letters) < cfg.seq_len:
        letters.insert(0, 'a')
    
    # Render as images
    images = []
    for letter in letters:
        if letter.isalpha():
            img = render_letter(letter)
            images.append(img)
        else:
            images.append(np.zeros((28, 28), dtype=np.float32))
    
    generated = seed_text
    
    with torch.no_grad():
        for _ in range(length):
            # Build input tensor
            input_images = np.stack(images[-cfg.seq_len:])
            input_tensor = torch.FloatTensor(input_images).unsqueeze(0).unsqueeze(1).to(cfg.device)
            
            logits, _ = model(input_tensor, return_entropy=False)
            
            # Temperature sampling
            probs = F.softmax(logits / temperature, dim=1)
            pred_idx = torch.multinomial(probs, 1).item()
            
            pred_letter = chr(pred_idx + ord('a'))
            generated += pred_letter
            
            # Add to history
            if pred_letter.isalpha():
                images.append(render_letter(pred_letter))
            else:
                images.append(np.zeros((28, 28), dtype=np.float32))
    
    return generated

# ═══════════════════════════════════════════════════════════════════
# MAIN
# ═══════════════════════════════════════════════════════════════════

def main():
    print("🧠 Unified Letter Image Sequence Predictor")
    print("   Conv2D (28x28) → LSTM → Next Character")
    print("="*60)
    
    # Step 1: Load text
    print("\n📂 Step 1: Loading text files...")
    text = load_text_files(cfg.book_folder)
    
    # Step 2: Build model
    print("\n🏗️ Step 2: Building unified model...")
    model = LetterSequencePredictor(
        num_classes=26,  # a-z
        embed_dim=cfg.embed_dim,
        hidden_dim=cfg.hidden_dim,
        num_layers=cfg.num_layers
    )
    print(model)
    
    # Step 3: Train (on-the-fly)
    print("\n🚀 Step 3: Training (on-the-fly letter image generation)...")
    print(f"   Sequence length: {cfg.seq_len} letter images")
    print(f"   Each letter: 28x28 grayscale")
    model = train_on_the_fly(model, text, epochs=cfg.epochs, lr=cfg.lr)
    
    # Step 4: Test
    print("\n🧪 Step 4: Testing...")
    accuracy = test_model(model, text)
    
    # Step 5: Generate
    print("\n🎨 Text Generation:")
    seed_texts = ["the", "god", "and", "light"]
    
    for seed in seed_texts:
        generated = generate_text(model, seed, length=50)
        print(f"\n  Seed: '{seed}'")
        print(f"  Generated: '{generated[:100]}...'")
    
    # Save
    torch.save({
        'model_state_dict': model.state_dict(),
        'config': {
            'num_classes': 26,
            'embed_dim': cfg.embed_dim,
            'hidden_dim': cfg.hidden_dim,
            'num_layers': cfg.num_layers,
            'seq_len': cfg.seq_len
        }
    }, 'letter_sequence_model.pt')
    
    print("\n💾 Model saved to 'letter_sequence_model.pt'")
    print(f"Final Test Accuracy: {accuracy:.2f}%")

if __name__ == "__main__":
    main()
