"""
Prefix-to-word completion with exact-match reward.

Task:
- Input: rendered prefix of a word, e.g. "the_" or "consci_"
- Target: the full word identity

Training:
- Supervised cross-entropy warmup
- Reward fine-tuning with exact-match reward

This is closer to a masked-word completion task than next-word prediction.
"""

import re
import random
from collections import defaultdict
from pathlib import Path

import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from PIL import Image, ImageDraw, ImageFont
from torch.utils.data import Dataset, DataLoader
from torch.nn.utils.rnn import pad_sequence


class Config:
    book_folder = "./books"

    # Image
    img_size = 28
    font_size = 20

    # Model
    embed_dim = 128
    num_classes = None

    # Prefix task
    min_prefix_len = 1
    max_prefix_len = 4  # set to 1 for first-letter-only

    # Training
    batch_size = 64
    epochs = 20
    lr = 0.001
    warmup_epochs = 3
    cls_weight = 0.7
    reward_weight = 0.3
    baseline_momentum = 0.9

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")


cfg = Config()


def load_markdown_files(folder_path):
    words = []
    folder = Path(folder_path)
    if not folder.exists():
        print(f"⚠️ Folder '{folder_path}' not found. Using sample text.")
        return sample_words()

    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()
        extracted = re.findall(r"[a-zA-Z]{2,}", content)
        words.extend([w.lower() for w in extracted])

    if len(words) == 0:
        print(f"⚠️ No words found in {folder_path}. Using sample text.")
        return sample_words()

    print(f"✅ Extracted {len(words)} words from markdown files.")
    return words


def sample_words():
    sample_text = """
    intelligence artificial machine learning neural network deep
    consciousness awareness thinking reasoning planning algorithm
    data information knowledge wisdom understanding insight logic
    mathematics physics chemistry biology science research study
    """
    return re.findall(r"[a-zA-Z]{2,}", sample_text.lower())


def build_vocabulary(words, max_vocab_size=5000):
    word_counts = defaultdict(int)
    for word in words:
        word_counts[word] += 1

    sorted_words = sorted(word_counts.items(), key=lambda x: -x[1])
    vocab_words = [w for w, _ in sorted_words[:max_vocab_size]]

    word2idx = {word: idx for idx, word in enumerate(vocab_words)}
    idx2word = {idx: word for word, idx in word2idx.items()}

    print(f"📚 Vocabulary size: {len(vocab_words)}")
    return word2idx, idx2word, vocab_words


def render_text_image(text, size=28, font_size=20):
    img = Image.new("L", (size, size), color=255)
    draw = ImageDraw.Draw(img)

    try:
        font = ImageFont.truetype("/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", font_size)
    except Exception:
        try:
            font = ImageFont.truetype("arial.ttf", font_size)
        except Exception:
            font = ImageFont.load_default()

    bbox = draw.textbbox((0, 0), text, 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), text, fill=0, font=font)

    img_array = np.array(img, dtype=np.float32) / 255.0
    return 1.0 - img_array


def render_prefix_image(word, prefix_len, size=28, font_size=20):
    prefix = word[:prefix_len]
    # Trailing underscore makes the "missing remainder" explicit.
    return render_text_image(prefix + "_", size=size, font_size=font_size)


class PrefixDataset(Dataset):
    """Prefix-to-word dataset with exact-match target labels."""

    def __init__(self, words, word2idx, max_samples=50000):
        self.samples = []
        self.word2idx = word2idx

        vocab_set = set(word2idx.keys())
        valid_words = [w for w in words if w in vocab_set and len(w) > 1]

        if len(valid_words) > max_samples:
            valid_words = valid_words[:max_samples]

        for word in valid_words:
            max_prefix = min(cfg.max_prefix_len, len(word) - 1)
            min_prefix = min(cfg.min_prefix_len, max_prefix)
            prefix_len = random.randint(min_prefix, max_prefix)
            prefix = word[:prefix_len]

            try:
                img = render_prefix_image(word, prefix_len, size=cfg.img_size, font_size=cfg.font_size)
            except Exception:
                img = np.zeros((cfg.img_size, cfg.img_size), dtype=np.float32)

            self.samples.append(
                {
                    "image": img,
                    "target_idx": word2idx[word],
                    "prefix": prefix,
                    "word": word,
                }
            )

        print(f"🖼️ Prefix dataset: {len(self.samples)} samples")

    def __len__(self):
        return len(self.samples)

    def __getitem__(self, idx):
        sample = self.samples[idx]
        return (
            torch.FloatTensor(sample["image"]).unsqueeze(0),
            torch.LongTensor([sample["target_idx"]]),
            sample["prefix"],
            sample["word"],
        )


def collate_prefix_batch(batch):
    images, target_idx, prefix_text, full_word = zip(*batch)
    images = torch.stack(images, dim=0)
    target_idx = torch.stack(target_idx, dim=0)
    return images, target_idx, list(prefix_text), list(full_word)


class PrefixWordMLP(nn.Module):
    """Visual encoder + classifier for prefix completion."""

    def __init__(self, num_classes, embed_dim=128):
        super().__init__()
        self.features = nn.Sequential(
            nn.Conv2d(1, 32, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),
            nn.Conv2d(32, 64, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),
            nn.Conv2d(64, 128, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),
        )
        self.flatten = nn.Flatten()
        self.embedding = nn.Linear(128 * 3 * 3, embed_dim)
        self.classifier = nn.Linear(embed_dim, num_classes)

    def forward(self, x):
        x = self.features(x)
        x = self.flatten(x)
        embed = self.embedding(x)
        logits = self.classifier(embed)
        return logits, embed


def train_model(model, train_loader, val_loader, epochs=20, lr=0.001):
    optimizer = torch.optim.Adam(model.parameters(), lr=lr)
    criterion = nn.CrossEntropyLoss()
    baseline = 0.0

    model.to(cfg.device)

    for epoch in range(epochs):
        model.train()
        train_loss = 0.0
        train_ce = 0.0
        train_reward = 0.0
        train_correct = 0
        train_total = 0

        for images, target_idx, _, _ in train_loader:
            images = images.to(cfg.device)
            target_idx = target_idx.squeeze(-1).to(cfg.device)

            optimizer.zero_grad()
            logits, _ = model(images)

            ce_loss = criterion(logits, target_idx)
            probs = F.softmax(logits, dim=1)
            dist = torch.distributions.Categorical(probs)
            sampled = dist.sample()
            reward = (sampled == target_idx).float()
            reward_mean = reward.mean().item()

            if epoch < cfg.warmup_epochs:
                loss = ce_loss
            else:
                baseline = cfg.baseline_momentum * baseline + (1.0 - cfg.baseline_momentum) * reward_mean
                advantage = reward - baseline
                rl_loss = -(advantage.detach() * dist.log_prob(sampled)).mean()
                loss = cfg.cls_weight * ce_loss + cfg.reward_weight * rl_loss

            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
            optimizer.step()

            preds = logits.argmax(dim=1)
            train_correct += (preds == target_idx).sum().item()
            train_total += target_idx.size(0)
            train_loss += loss.item()
            train_ce += ce_loss.item()
            train_reward += reward_mean

        train_acc = 100.0 * train_correct / train_total
        avg_loss = train_loss / max(len(train_loader), 1)
        avg_ce = train_ce / max(len(train_loader), 1)
        avg_reward = train_reward / max(len(train_loader), 1)

        model.eval()
        val_loss = 0.0
        val_correct = 0
        val_total = 0
        val_reward = 0.0

        with torch.no_grad():
            for images, target_idx, _, _ in val_loader:
                images = images.to(cfg.device)
                target_idx = target_idx.squeeze(-1).to(cfg.device)

                logits, _ = model(images)
                loss = criterion(logits, target_idx)
                preds = logits.argmax(dim=1)

                val_loss += loss.item()
                val_correct += (preds == target_idx).sum().item()
                val_total += target_idx.size(0)
                val_reward += (preds == target_idx).float().mean().item()

        val_acc = 100.0 * val_correct / val_total
        avg_val_loss = val_loss / max(len(val_loader), 1)
        avg_val_reward = val_reward / max(len(val_loader), 1)

        print(
            f"Epoch {epoch+1}/{epochs} | "
            f"Loss: {avg_loss:.4f} | CE: {avg_ce:.4f} | "
            f"Train: {train_acc:.2f}% | Reward: {avg_reward:.3f} | "
            f"Val Loss: {avg_val_loss:.4f} | Val: {val_acc:.2f}% | "
            f"Val Reward: {avg_val_reward:.3f}"
        )

    return model


def test_model(model, test_loader, idx2word):
    model.eval()
    criterion = nn.CrossEntropyLoss()

    test_loss = 0.0
    test_correct = 0
    test_total = 0
    exact_rewards = []

    with torch.no_grad():
        for images, target_idx, prefix_text, full_word in test_loader:
            images = images.to(cfg.device)
            target_idx = target_idx.squeeze(-1).to(cfg.device)

            logits, _ = model(images)
            loss = criterion(logits, target_idx)
            preds = logits.argmax(dim=1)

            test_loss += loss.item()
            test_correct += (preds == target_idx).sum().item()
            test_total += target_idx.size(0)
            exact_rewards.extend((preds == target_idx).float().tolist())

    test_acc = 100.0 * test_correct / test_total
    avg_loss = test_loss / max(len(test_loader), 1)
    avg_reward = float(np.mean(exact_rewards)) if exact_rewards else 0.0

    print("\n" + "=" * 60)
    print("📊 PREFIX COMPLETION TEST RESULTS")
    print("=" * 60)
    print(f"Test Loss: {avg_loss:.4f}")
    print(f"Exact-Match Accuracy / Reward: {test_acc:.2f}% ({test_correct}/{test_total})")
    print(f"Average Reward: {avg_reward:.4f}")

    print("\n🔍 Sample Predictions:")
    shown = 0
    with torch.no_grad():
        for images, target_idx, prefix_text, full_word in test_loader:
            logits, _ = model(images.to(cfg.device))
            preds = logits.argmax(dim=1).cpu()

            for i in range(min(4, len(full_word))):
                pred_word = idx2word.get(preds[i].item(), "?")
                true_word = full_word[i]
                prefix = prefix_text[i]
                status = "✅" if pred_word == true_word else "❌"
                print(f"  {status} '{prefix}_' -> True: '{true_word}' | Pred: '{pred_word}'")
                shown += 1
                if shown >= 8:
                    break
            if shown >= 8:
                break

    print("=" * 60)
    return test_acc


def main():
    print("🧠 Prefix Word Completion (Exact-Match Reward)")
    print("=" * 60)

    print("\n📂 Step 1: Loading markdown files...")
    words = load_markdown_files(cfg.book_folder)

    print("\n📚 Step 2: Building vocabulary...")
    word2idx, idx2word, vocab_words = build_vocabulary(words, max_vocab_size=2000)
    cfg.num_classes = len(vocab_words)

    print(f"\n🖼️ Step 3: Creating prefix dataset (max_prefix_len={cfg.max_prefix_len})...")
    dataset = PrefixDataset(words, word2idx, max_samples=40000)

    print("\n📊 Step 4: Splitting data...")
    total = len(dataset)
    train_size = int(0.7 * total)
    val_size = int(0.15 * total)
    test_size = total - train_size - val_size

    train_dataset, val_dataset, test_dataset = torch.utils.data.random_split(
        dataset, [train_size, val_size, test_size]
    )

    train_loader = DataLoader(
        train_dataset, batch_size=cfg.batch_size, shuffle=True, collate_fn=collate_prefix_batch
    )
    val_loader = DataLoader(
        val_dataset, batch_size=cfg.batch_size, shuffle=False, collate_fn=collate_prefix_batch
    )
    test_loader = DataLoader(
        test_dataset, batch_size=cfg.batch_size, shuffle=False, collate_fn=collate_prefix_batch
    )

    print(f"  Train: {len(train_dataset)} | Val: {len(val_dataset)} | Test: {len(test_dataset)}")

    print("\n🏗️ Step 5: Building model...")
    model = PrefixWordMLP(num_classes=cfg.num_classes, embed_dim=cfg.embed_dim)
    print(model)

    print("\n🚀 Step 6: Training prefix completion model...")
    model = train_model(model, train_loader, val_loader, epochs=cfg.epochs, lr=cfg.lr)

    print("\n🧪 Step 7: Testing...")
    test_acc = test_model(model, test_loader, idx2word)

    torch.save(
        {
            "model_state_dict": model.state_dict(),
            "word2idx": word2idx,
            "idx2word": idx2word,
            "config": {
                "embed_dim": cfg.embed_dim,
                "num_classes": cfg.num_classes,
                "min_prefix_len": cfg.min_prefix_len,
                "max_prefix_len": cfg.max_prefix_len,
            },
        },
        "prefix_reward_completion.pt",
    )

    print("\n💾 Model saved to 'prefix_reward_completion.pt'")
    print(f"Final Test Accuracy: {test_acc:.2f}%")


if __name__ == "__main__":
    main()
