import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
import numpy as np
import re

# ============================================================
# 1. DATASET: Chunked "books" → (prompt, target, fragment_pool)
# ============================================================
class BookRecombinationDataset(Dataset):
    def __init__(self, raw_texts, chunk_size=32, max_fragments=128, vocab_size=1024):
        self.vocab_size = vocab_size
        self.max_fragments = max_fragments
        
        # Simulate tokenization for demo (replace with real tokenizer)
        self.fragments = []
        for text in raw_texts:
            # Split into overlapping chunks
            chunks = [text[i:i+chunk_size] for i in range(0, max(len(text)-chunk_size+1, 1), chunk_size//2)]
            self.fragments.extend(chunks)
        self.fragments = self.fragments[:max_fragments]  # Fixed pool for demo
        
        # Generate synthetic (prompt, target) pairs by recombining fragments
        self.samples = []
        for _ in range(200):
            # Randomly sample 2-4 fragments as "prompt"
            k = np.random.randint(2, 5)
            prompt_frags = np.random.choice(self.fragments, size=k, replace=False)
            prompt = " [SEP] ".join(prompt_frags)
            
            # Target: reorder + slightly mutate (simulates human/AI recombination)
            target_frags = prompt_frags.copy()
            np.random.shuffle(target_frags)
            # Add noise/overlap to simulate "fuzzy" recombination
            target = " ".join([f[:int(len(f)*np.random.uniform(0.7, 1.0))] for f in target_frags])
            
            self.samples.append((prompt, target))
            
        # Simple integer tokenization
        self.char2idx = {c: i+2 for i, c in enumerate(sorted(set("".join(self.fragments))))}
        self.char2idx["[PAD]"] = 0
        self.char2idx["[SEP]"] = 1
        self.idx2char = {v: k for k, v in self.char2idx.items()}
        
    def encode(self, text, max_len):
        ids = [self.char2idx.get(c, 1) for c in text]
        ids = ids[:max_len]
        ids += [0] * (max_len - len(ids))
        return torch.tensor(ids, dtype=torch.long)
        
    def __len__(self): return len(self.samples)
    def __getitem__(self, idx):
        prompt, target = self.samples[idx]
        frag_ids = torch.stack([self.encode(f, 32) for f in self.fragments])
        return (
            self.encode(prompt, 128),
            self.encode(target, 128),
            frag_ids,
            torch.tensor([len(self.fragments)]),  # actual pool size
            torch.tensor(len(target))             # target length for loss masking
        )

# ============================================================
# 2. MODEL: Retriever + Fuzzy Recombiner
# ============================================================
class FuzzyRecombiner(nn.Module):
    def __init__(self, vocab_size, embed_dim=128, num_heads=4, num_layers=2, dropout=0.1):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embed_dim)
        self.drop = nn.Dropout(dropout)

        # Retrieval: OWN embedding layer (not shared with generation)
        self.retrieve_embedding = nn.Embedding(vocab_size, embed_dim)
        self.retrieve_attn = nn.Linear(embed_dim, embed_dim)
        self.retrieve_key = nn.Linear(embed_dim, embed_dim)
        self.pool_query = nn.Parameter(torch.randn(embed_dim))
        self.retrieve_gate = nn.Parameter(torch.tensor(0.0))  # starts favoring lexical (sigmoid(0)=0.5, but neural is unscaled)

        # Recombination decoder
        decoder_layer = nn.TransformerDecoderLayer(
            d_model=embed_dim, nhead=num_heads, dim_feedforward=embed_dim*4, dropout=dropout, batch_first=True
        )
        self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=num_layers)
        self.out_proj = nn.Linear(embed_dim, vocab_size)

    def _weighted_pool(self, embeddings, mask=None):
        """Attention-weighted pooling: [B, L, D] -> [B, D]"""
        attn = torch.matmul(embeddings, self.pool_query)  # [B, L]
        if mask is not None:
            attn = attn.masked_fill(mask == 0, float('-inf'))
        weights = F.softmax(attn, dim=-1).unsqueeze(-1)  # [B, L, 1]
        return (embeddings * weights).sum(dim=1)  # [B, D]

    def retrieve(self, prompt_ids, fragment_ids):
        # Use retrieval-specific embeddings
        prompt_emb = self.retrieve_embedding(prompt_ids)  # [B, L_p, D]
        q = self.retrieve_attn(self._weighted_pool(prompt_emb, mask=(prompt_ids > 0)))  # [B, D]
        
        if fragment_ids.dim() == 3:
            B, N, L_f = fragment_ids.shape
            frag_flat = fragment_ids.view(B * N, L_f)
            frag_emb_flat = self.retrieve_embedding(frag_flat)  # [B*N, L_f, D]
            frag_mask_flat = (frag_flat > 0)
            
            k_flat = self.retrieve_key(self._weighted_pool(frag_emb_flat, mask=frag_mask_flat))  # [B*N, D]
            k = k_flat.view(B, N, -1)  # [B, N, D]
            
            # Neural scores: scaled dot product [B, N]
            neural_scores = torch.matmul(q.unsqueeze(1), k.transpose(1, 2)).squeeze(1) / (k.size(-1) ** 0.5)
            
            # Lexical overlap: [B, N]
            prompt_vocab = F.one_hot(prompt_ids.clamp(min=0), num_classes=self.retrieve_embedding.num_embeddings).float()
            prompt_vocab[:, :, 0] = 0
            prompt_vocab = prompt_vocab.sum(dim=1).clamp(min=0, max=1)  # [B, V]
            
            frag_onehot = F.one_hot(frag_flat.clamp(min=0), num_classes=self.retrieve_embedding.num_embeddings).float()
            frag_onehot[:, :, 0] = 0
            frag_vocab = frag_onehot.sum(dim=1).view(B, N, -1).clamp(min=0, max=1)  # [B, N, V]
            
            overlap = (frag_vocab * prompt_vocab.unsqueeze(1)).sum(dim=-1) / frag_vocab.sum(dim=-1).clamp(min=1)
            
            # Gated combination: learnable weight for neural vs lexical
            alpha = torch.sigmoid(self.retrieve_gate)
            scores = alpha * neural_scores + (1 - alpha) * overlap * 5.0
            f_emb = k
        else:
            N = fragment_ids.size(0)
            frag_emb = self.retrieve_embedding(fragment_ids)
            k = self.retrieve_key(self._weighted_pool(frag_emb, mask=(fragment_ids > 0)))  # [N, D]
            neural_scores = torch.matmul(q, k.transpose(0, 1)) / (k.size(-1) ** 0.5)
            
            prompt_vocab = F.one_hot(prompt_ids.clamp(min=0), num_classes=self.retrieve_embedding.num_embeddings).float()
            prompt_vocab[:, :, 0] = 0
            prompt_vocab = prompt_vocab.sum(dim=1).clamp(min=0, max=1)
            
            frag_onehot = F.one_hot(fragment_ids.clamp(min=0), num_classes=self.retrieve_embedding.num_embeddings).float()
            frag_onehot[:, :, 0] = 0
            frag_vocab = frag_onehot.sum(dim=1).clamp(min=0, max=1)
            
            overlap = (frag_vocab * prompt_vocab.unsqueeze(0)).sum(dim=-1) / frag_vocab.sum(dim=-1).clamp(min=1)
            
            alpha = torch.sigmoid(self.retrieve_gate)
            scores = alpha * neural_scores + (1 - alpha) * overlap * 5.0
            f_emb = k
            
        return scores, f_emb
        
    def forward(self, prompt_ids, fragment_ids, target_ids):
        scores, f_emb = self.retrieve(prompt_ids, fragment_ids)

        # Soft-retrieve top-K fragments (differentiable)
        weights = F.softmax(scores / 0.5, dim=-1)  # [B, N]

        # Weighted fragment context
        # f_emb: [B, N, D] or [N, D], weights: [B, N] -> ctx: [B, D]
        if f_emb.dim() == 3:
            # Batched: [B, N, D]
            ctx = (weights.unsqueeze(-1) * f_emb).sum(dim=1)  # [B, D]
        else:
            # Unbatched: [N, D]
            ctx = (weights.unsqueeze(-1) * f_emb.unsqueeze(0)).sum(dim=1)  # [B, D]
        ctx = ctx.unsqueeze(1)  # [B, 1, D]

        # Decoder memory: concatenate prompt context + retrieved fragments
        prompt_ctx = self.embedding(prompt_ids)  # [B, L_p, D]
        memory = torch.cat([prompt_ctx, ctx], dim=1)  # [B, L_p + 1, D]

        # Teacher-forced decoding
        tgt_emb = self.embedding(target_ids[:, :-1])  # shift for autoregressive
        logits = self.decoder(tgt_emb, memory)  # [B, L_t-1, D]
        logits = self.out_proj(logits)  # [B, L_t-1, V]

        return logits, scores, weights

# ============================================================
# 3. FUZZY SUPERVISED LOSS
# ============================================================
class FuzzySupervisedLoss(nn.Module):
    def __init__(self, model, vocab_size, label_smoothing=0.1, alpha_ret=1.0, alpha_align=0.3):
        super().__init__()
        self.model = model
        self.vocab_size = vocab_size
        self.ce = nn.CrossEntropyLoss(label_smoothing=label_smoothing, reduction='none')
        self.alpha_ret = alpha_ret
        self.alpha_align = alpha_align

    def forward(self, logits, scores, weights, target_ids, target_lens, prompt_ids, fragment_ids):
        B, L, V = logits.shape
        mask = torch.arange(L, device=target_ids.device).unsqueeze(0) < target_lens[:, None] - 1
        loss_gen = (self.ce(logits.reshape(-1, V), target_ids[:, 1:].reshape(-1)) * mask.reshape(-1)).sum() / mask.clamp(min=1).sum()

        # --- Retrieval loss: cross-entropy with soft overlap labels ---
        # Compute target-fragment overlap as soft labels
        if fragment_ids.dim() == 3:
            B, N, L_f = fragment_ids.shape
            frag_flat = fragment_ids.view(B * N, L_f)
            tgt_expanded = target_ids.unsqueeze(1).expand(-1, N, -1).reshape(B * N, -1)
            
            frag_present = F.one_hot(frag_flat.clamp(min=0), num_classes=self.vocab_size).sum(dim=1).float()
            tgt_present = F.one_hot(tgt_expanded.clamp(min=0), num_classes=self.vocab_size).sum(dim=1).float()
            frag_present[:, 0] = 0
            tgt_present[:, 0] = 0
            
            intersection = (frag_present * tgt_present).sum(dim=-1).view(B, N)
            frag_size = frag_present.sum(dim=-1).view(B, N).clamp(min=1)
            overlap_ratio = intersection / frag_size  # [B, N]
            
            frag_emb = self.model.retrieve_key(self.model.retrieve_embedding(fragment_ids.view(B * N, L_f)).mean(dim=1)).view(B, N, -1)
        else:
            N = fragment_ids.size(0)
            frag_present = F.one_hot(fragment_ids.clamp(min=0), num_classes=self.vocab_size).sum(dim=1).float()
            frag_present[:, 0] = 0
            overlap_ratio = torch.zeros(B, N, device=target_ids.device)
            for b in range(B):
                tgt_p = F.one_hot(target_ids[b].clamp(min=0), num_classes=self.vocab_size).sum(dim=0).float()
                tgt_p[0] = 0
                intersection = (frag_present * tgt_p).sum(dim=-1)
                frag_size = frag_present.sum(dim=-1).clamp(min=1)
                overlap_ratio[b] = intersection / frag_size
            frag_emb = self.model.retrieve_key(self.model.retrieve_embedding(fragment_ids).mean(dim=1))

        # --- Retrieval loss: cross-entropy with target-fragment overlap ---
        # Soft target: overlap ratio normalized to distribution
        target_dist = overlap_ratio / overlap_ratio.sum(dim=-1, keepdim=True).clamp(min=1e-8)
        target_dist = target_dist + 1e-8  # prevent log(0)
        target_dist = target_dist / target_dist.sum(dim=-1, keepdim=True)
        
        log_probs = F.log_softmax(scores, dim=-1)
        ret_loss = -(target_dist * log_probs).sum(dim=-1).mean()

        # --- Alignment loss ---
        retrieved_ctx = (weights.unsqueeze(-1) * frag_emb).sum(dim=1)  # [B, D]
        tgt_emb = self.model.retrieve_key(self.model.retrieve_embedding(target_ids).mean(dim=1))  # [B, D]
        align_loss = 1.0 - F.cosine_similarity(retrieved_ctx, tgt_emb, dim=-1).mean()

        return loss_gen + self.alpha_ret * ret_loss + self.alpha_align * align_loss, {
            'gen': loss_gen.item(), 'ret': ret_loss.item(), 'align': align_loss.item()
        }

# ============================================================
# 4. TRAINING LOOP
# ============================================================
def train(model, dataloader, loss_fn, optimizer, epochs=3, device='cpu'):
    model.train()
    for epoch in range(epochs):
        total_loss = 0
        for prompts, targets, fragments, pool_sizes, tgt_lens in dataloader:
            prompts, targets, fragments = prompts.to(device), targets.to(device), fragments.to(device)
            tgt_lens = tgt_lens.to(device)

            optimizer.zero_grad()
            logits, scores, weights = model(prompts, fragments, targets)
            loss, components = loss_fn(logits, scores, weights, targets, tgt_lens, prompts, fragments)
            loss.backward()
            optimizer.step()
            total_loss += loss.item()
        print(f"Epoch {epoch+1} | Loss: {total_loss/len(dataloader):.4f} | Gen: {components['gen']:.3f} | Ret: {components['ret']:.3f} | Align: {components['align']:.3f}")

# ============================================================
# 5. INFERENCE
# ============================================================
@torch.no_grad()
def recombine_text(model, prompt_text, dataset, max_len=64, top_k=32, device='cpu'):
    model.eval()
    prompt_ids = dataset.encode(prompt_text, 128).unsqueeze(0).to(device)
    fragments = dataset.fragments
    frag_ids = torch.stack([dataset.encode(f, 32) for f in fragments]).to(device)
    
    scores, f_emb = model.retrieve(prompt_ids, frag_ids)
    actual_k = min(top_k, scores.size(-1))
    topk_idx = torch.topk(scores, actual_k, dim=-1).indices.squeeze()
    
    retrieved = [fragments[i] for i in topk_idx.cpu().numpy()]
    retrieved_ids = torch.stack([dataset.encode(r, 32) for r in retrieved]).unsqueeze(0).to(device)
    
    # Autoregressive decoding
    prompt_ctx = model.embedding(prompt_ids)
    scores, f_emb = model.retrieve(prompt_ids, frag_ids)
    ctx = (F.softmax(scores, dim=-1).unsqueeze(-1) * f_emb.unsqueeze(0)).sum(dim=1)
    ctx = ctx.unsqueeze(1)
    memory = torch.cat([prompt_ctx, ctx], dim=1)
    
    generated = torch.zeros(1, 1, dtype=torch.long, device=device)
    for _ in range(max_len):
        tgt_emb = model.embedding(generated)
        out = model.decoder(tgt_emb, memory)
        logits = model.out_proj(out[:, -1, :])
        next_token = torch.argmax(logits, dim=-1, keepdim=True)  # [1, 1]
        generated = torch.cat([generated, next_token], dim=1)
        
    out_text = "".join([dataset.idx2char[t.item()] for t in generated.squeeze() if t.item() != 0])
    return out_text, retrieved

# ============================================================
# 6. RUN EXAMPLE
# ============================================================
if __name__ == "__main__":
    # Simulated "stack of books"
    raw_books = [
        "The quick brown fox jumps over the lazy dog near the river bank.",
        "Machine learning models optimize parameters through gradient descent.",
        "In ancient forests, moss covered stones whispered secrets to the wind.",
        "Neural networks approximate functions by stacking nonlinear transformations.",
        "The detective followed clues through crowded streets and dimly lit alleys."
    ]
    
    dataset = BookRecombinationDataset(raw_books, vocab_size=128)
    loader = DataLoader(dataset, batch_size=4, shuffle=True)
    
    model = FuzzyRecombiner(vocab_size=dataset.vocab_size, embed_dim=128).to('cpu')
    loss_fn = FuzzySupervisedLoss(model, vocab_size=dataset.vocab_size)
    opt = torch.optim.Adam(model.parameters(), lr=1e-3)
    
    train(model, loader, loss_fn, opt, epochs=30, device='cpu')
    
    prompt = "The quick model followed clues near the forest bank."
    out, frags = recombine_text(model, prompt, dataset, device='cpu')
    print("\n[Recombined Output]")
    print(out)
    print("\n[Retrieved Fragments]")
    for i, f in enumerate(frags): print(f"{i+1}. {f}")
