"""
Structured-Factorial Model: Work-Based Training, Simple Inference
==================================================================
Training:   High λ_sec → constrained path discovery (hard work)
Inference:  Low λ_sec  → path verification (easy check)
"""

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, Subset
from torchvision import datasets, transforms
import numpy as np
from dataclasses import dataclass, field
from typing import List, Tuple, Dict
import math


# ============================================================================
# STRUCTURED-FACTORIAL SERVER
# ============================================================================

class StructuredFactorialServer:
    """
    Server computes the Structured-Factorial F_S(n) and provides:
    - Training: High λ_sec (work required to discover valid paths)
    - Inference: Low λ_sec (verification of pre-computed paths)
    """
    
    def __init__(self, n_layers: int, hidden_dims: List[int], gamma: float = 1.5):
        self.n_layers = n_layers
        self.hidden_dims = hidden_dims
        self.gamma = gamma
        
        # F_S(n) = 1 for chain topology (only one valid path)
        self.F_S = 1
        self.n_factorial = math.factorial(n_layers)
        self.confinement_ratio = self.F_S / self.n_factorial
        
        # Training work parameters
        self.lambda_base_train = 5.0   # High damping for training (work)
        self.lambda_base_infer = 0.1   # Low damping for inference (verify)
        
        self.kappa = 2.0
        self.eta = 0.5
        
        print(f"Server: F_S({n_layers})={self.F_S}, R={self.confinement_ratio:.6f}")
    
    def compute_lambda(self, H_AS: float = 0.0, mode: str = 'train') -> float:
        """
        Security Threshold Equation:
        - Training: λ is HIGH (must do work to find paths)
        - Inference: λ is LOW (just verify pre-computed paths)
        """
        lambda_base = (self.lambda_base_train if mode == 'train' 
                       else self.lambda_base_infer)
        
        # Confinement factor (n! / F_S)^γ
        confinement_factor = (self.n_factorial / self.F_S) ** self.gamma
        
        lambda_sec = lambda_base + self.kappa * H_AS * confinement_factor
        return lambda_sec
    
    def compute_work(self) -> float:
        """
        Work = -log(R) = log(n!) - log(F_S)
        Energy required to find valid path through constraint topology.
        """
        return math.log(self.n_factorial) - math.log(self.F_S)


# ============================================================================
# CLIENT MODEL WITH WORK-BASED TRAINING
# ============================================================================

@dataclass
class SFConfig:
    input_size: int = 784
    num_classes: int = 10
    hidden_dims: List[int] = field(default_factory=lambda: [100,100])
    n_layers: int = 2
    gamma: float = 1.5
    batch_size: int = 100
    epochs: int = 5
    learning_rate: float = 0.001
    test_split: float = 0.2
    work_coefficient: float = 0.97  # How much λ affects training


class WorkConstrainedModel(nn.Module):
    """
    Model that:
    - TRAINING: Uses high-λ constrained forward pass (work-based)
    - INFERENCE: Uses low-λ simple verification
    """
    
    def __init__(self, config: SFConfig, server: StructuredFactorialServer):
        super().__init__()
        self.config = config
        self.server = server
        
        # Build layers
        layers = []
        in_dim = config.input_size
        for dim in config.hidden_dims:
            layers.append(nn.Linear(in_dim, dim))
            in_dim = dim
        
        self.hidden_layers = nn.ModuleList(layers)
        self.classifier = nn.Linear(in_dim, config.num_classes)
        
        # Pre-computed valid paths (from server)
        self.valid_paths = self._compute_valid_paths()
        
    def _compute_valid_paths(self) -> List[List[int]]:
        """Compute valid path orderings (F_S(n) = 1 for chain)"""
        return [list(range(self.config.n_layers))]  # Only one valid path
    
    def forward_constrained(self, x: torch.Tensor, mode: str = 'train') -> Tuple:
        """
        Constrained forward pass.
        
        Training mode:   Apply high λ_sec → constrained gradient flow (work)
        Inference mode:  Apply low λ_sec  → free gradient flow (verify)
        """
        lambda_sec = self.server.compute_lambda(mode=mode)
        work = self.server.compute_work()
        
        h = x
        for i, layer in enumerate(self.hidden_layers):
            h = layer(h)
            h = torch.relu(h)
            
            # Apply damping based on λ_sec (topological friction)
            damping = 1.0 / (1 + lambda_sec * (i + 1) / self.config.n_layers)
            h = h * damping
        
        logits = self.classifier(h)
        
        return logits, lambda_sec, work
    
    def forward(self, x: torch.Tensor, mode: str = 'train') -> torch.Tensor:
        """Standard interface"""
        logits, _, _ = self.forward_constrained(x, mode)
        return logits


# ============================================================================
# TRAINING PIPELINE WITH WORK TERM
# ============================================================================

class WorkBasedPipeline:
    """
    Training requires work (high λ_sec) to discover valid paths.
    Inference is simple verification (low λ_sec).
    """
    
    def __init__(self, config: SFConfig):
        self.config = config
        self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
        
        # Server
        self.server = StructuredFactorialServer(
            n_layers=config.n_layers,
            hidden_dims=config.hidden_dims,
            gamma=config.gamma
        )
        
        # Model
        self.model = WorkConstrainedModel(config, self.server).to(self.device)
        
        # Optimizer
        self.optimizer = optim.Adam(self.model.parameters(), lr=config.learning_rate)
        self.scheduler = optim.lr_scheduler.CosineAnnealingLR(
            self.optimizer, T_max=config.epochs, eta_min=1e-5
        )
        
        # Work coefficient (β in thermodynamics)
        self.beta = config.work_coefficient
        
        # Data
        self.train_loader, self.test_loader = self._prepare_data()
        
        self.history = {
            'train_loss': [], 'train_acc': [],
            'test_loss': [], 'test_acc': [],
            'lambda_train': [], 'work': []
        }
    
    def _prepare_data(self):
        transform = transforms.Compose([
            transforms.ToTensor(),
            transforms.Normalize((0.1307,), (0.3081,))
        ])
        
        dataset = datasets.MNIST(root='../../data', train=True, download=True, transform=transform)
        n = len(dataset)
        
        perm = np.random.permutation(n)
        test_size = int(n * self.config.test_split)
        
        train_loader = DataLoader(
            Subset(dataset, perm[:-test_size]), 
            batch_size=self.config.batch_size, shuffle=True
        )
        test_loader = DataLoader(
            Subset(dataset, perm[-test_size:]), 
            batch_size=self.config.batch_size, shuffle=False
        )
        
        print(f"Train: {n - test_size}, Test: {test_size}")
        return train_loader, test_loader
    
    def train_epoch(self, epoch: int) -> Tuple[float, float, float, float]:
        """
        Training with work penalty.
        High λ_sec during training = constrained optimization (work required).
        """
        self.model.train()
        total_loss, correct, total = 0, 0, 0
        
        # High λ for training (work mode)
        lambda_train = self.server.compute_lambda(mode='train')
        
        for data, target in self.train_loader:
            data, target = data.to(self.device), target.to(self.device)
            data = data.view(-1, self.config.input_size)
            
            self.optimizer.zero_grad()
            
            # Constrained forward pass (work-based)
            output, lambda_sec, work = self.model.forward_constrained(data, mode='train')
            
            # Loss with work penalty: P ∝ exp(-βW)
            base_loss = nn.CrossEntropyLoss()(output, target)
            work_penalty = self.beta * work * lambda_sec  # Work term
            loss = base_loss + work_penalty
            
            loss.backward()
            
            # Gradient clipping (topological constraint)
            torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0)
            
            self.optimizer.step()
            
            total_loss += loss.item()
            correct += (output.argmax(1) == target).sum().item()
            total += target.size(0)
        
        self.scheduler.step()
        
        avg_loss = total_loss / len(self.train_loader)
        acc = 100. * correct / total
        return avg_loss, acc, lambda_train, work
    
    @torch.no_grad()
    def evaluate(self, mode: str = 'infer') -> Tuple[float, float]:
        """
        Inference is simple verification (low λ_sec).
        """
        self.model.eval()
        total_loss, correct, total = 0, 0, 0
        
        for data, target in self.test_loader:
            data, target = data.to(self.device), target.to(self.device)
            data = data.view(-1, self.config.input_size)
            
            # Simple forward (verification)
            output = self.model(data, mode=mode)
            loss = nn.CrossEntropyLoss()(output, target)
            
            total_loss += loss.item()
            correct += (output.argmax(1) == target).sum().item()
            total += target.size(0)
        
        return total_loss / len(self.test_loader), 100. * correct / total
    
    def train(self) -> Dict:
        """Full training loop"""
        print("\n" + "="*65)
        print("WORK-BASED TRAINING (High λ_sec = Constrained Discovery)")
        print("="*65)
        print(f"Work per step: W = -log(R) = {self.server.compute_work():.4f}")
        print(f"Training λ_base: {self.server.lambda_base_train}")
        print(f"Inference λ_base: {self.server.lambda_base_infer}")
        print("="*65 + "\n")
        
        for epoch in range(1, self.config.epochs + 1):
            train_loss, train_acc, lambda_train, work = self.train_epoch(epoch)
            test_loss, test_acc = self.evaluate(mode='infer')
            
            # Store
            self.history['train_loss'].append(train_loss)
            self.history['train_acc'].append(train_acc)
            self.history['test_loss'].append(test_loss)
            self.history['test_acc'].append(test_acc)
            self.history['lambda_train'].append(lambda_train)
            self.history['work'].append(work)
            
            print(f"Epoch {epoch:2d}/{self.config.epochs} | "
                  f"Train: {train_loss:.4f}/{train_acc:.1f}% | "
                  f"Test: {test_loss:.4f}/{test_acc:.1f}% | "
                  f"λ_train: {lambda_train:.2f}")
        
        print("\n" + "="*65)
        print("TRAINING COMPLETE")
        print("="*65)
        
        return self.history
    
    def summary(self):
        """Final summary"""
        print("\n" + "="*65)
        print("SUMMARY: Training vs Inference")
        print("="*65)
        print(f"  Training: High work (λ={self.server.lambda_base_train})")
        print(f"           Must discover valid paths through topology")
        print(f"  Inference: Low work (λ={self.server.lambda_base_infer})")
        print(f"             Simple verification of pre-computed paths")
        print(f"\n  Final Train Acc: {self.history['train_acc'][-1]:.2f}%")
        print(f"  Final Test Acc:  {self.history['test_acc'][-1]:.2f}%")
        print(f"  Work done: {sum(self.history['work']):.2f}")


# ============================================================================
# MAIN
# ============================================================================

def main():
    print("\n" + "="*70)
    print("STRUCTURED-FACTORIAL: Work-Based Training, Simple Inference")
    print("Theory: Training discovers paths (high λ), Inference verifies (low λ)")
    print("="*70)
    
    config = SFConfig(
        hidden_dims=[512, 256, 128],
        n_layers=3,
        epochs=15,
        batch_size=128,
        work_coefficient=0.005
    )
    
    pipeline = WorkBasedPipeline(config)
    history = pipeline.train()
    pipeline.summary()
    
    return pipeline, history


if __name__ == "__main__":
    pipeline, history = main()
