import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader, Dataset
import torchvision
import torchvision.transforms as transforms
from dataclasses import dataclass, field
from typing import Dict, List, Tuple, Optional, Callable
from enum import Enum
import numpy as np
import logging
from collections import defaultdict

logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)


# =============================================================================
# DEVICE CONFIGURATION
# =============================================================================

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
logger.info(f"Using device: {device}")


# =============================================================================
# ERROR TAXONOMY & LOSS REGISTRY
# =============================================================================

class LossCategory(Enum):
    SPATIAL = "spatial"
    SPECTRAL = "spectral"
    STATISTICAL = "statistical"
    COMPONENT = "component"
    PATTERN = "pattern"
    ADVERSARIAL = "adversarial"
    REGULARIZATION = "regularization"


@dataclass
class LossConfig:
    name: str
    category: LossCategory
    weight: float = 1.0
    enabled: bool = True


@dataclass
class LossResult:
    name: str
    value: torch.Tensor
    gradient: torch.Tensor
    diagnostics: Dict
    category: LossCategory
    weight: float
    normalized_value: float


@dataclass
class MultiLossState:
    losses: Dict[str, LossResult]
    total_loss: torch.Tensor
    weighted_contributions: Dict[str, float]
    dominant_loss: str
    gradient_norms: Dict[str, float]
    optimization_advice: List[str]


# =============================================================================
# BASE LOSS FUNCTIONS (PyTorch)
# =============================================================================

class BaseLossFunction(nn.Module):
    """Base class for all diagnostic loss functions."""
    
    def __init__(self, config: LossConfig):
        super().__init__()
        self.config = config
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        raise NotImplementedError


# =============================================================================
# SPATIAL LOSSES
# =============================================================================

class SpatialConcentrationLoss(BaseLossFunction):
    """Penalizes spatially clustered errors."""
    
    def __init__(self, config: LossConfig):
        super().__init__(config)
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        # Handle batch dimension
        if error.dim() == 4:  # (B, 1, H, W)
            error_2d = error.mean(dim=1)  # (B, H, W)
            B = error_2d.shape[0]
            
            # Compute concentration per sample
            concentrations = []
            for b in range(B):
                e = error_2d[b]
                local_windows = []
                window_size = 5
                H, W = e.shape
                
                for y in range(0, H - window_size, window_size):
                    for x in range(0, W - window_size, window_size):
                        window = e[y:y+window_size, x:x+window_size]
                        local_windows.append(torch.var(window))
                
                local_var = torch.mean(torch.stack(local_windows)) if local_windows else torch.tensor(0.0)
                global_var = torch.var(e) + 1e-8
                concentrations.append(local_var / global_var)
            
            concentration = torch.mean(torch.stack(concentrations))
            
            # Gradient: based on variance differences
            grad_y = torch.gradient(error_2d, dim=1)[0]
            grad_x = torch.gradient(error_2d, dim=2)[0]
            gradient_magnitude = (grad_y**2 + grad_x**2).sqrt().mean(dim=0)
            
        else:  # Single sample (H, W)
            e = error
            local_windows = []
            window_size = 5
            H, W = e.shape
            
            for y in range(0, H - window_size, window_size):
                for x in range(0, W - window_size, window_size):
                    window = e[y:y+window_size, x:x+window_size]
                    local_windows.append(torch.var(window))
            
            local_var = torch.mean(torch.stack(local_windows)) if local_windows else torch.tensor(0.0)
            global_var = torch.var(e) + 1e-8
            concentration = local_var / global_var
            
            grad_y, grad_x = torch.gradient(e, dim=0), torch.gradient(e, dim=1)
            gradient_magnitude = (grad_y**2 + grad_x**2).sqrt()
        
        loss_value = torch.relu(concentration - 1.0)
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient_magnitude,
            diagnostics={'concentration': concentration.item()},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(concentration.item(), 2.0) / 2.0
        )


class SpatialEntropyLoss(BaseLossFunction):
    """Encourages uniform error distribution (high entropy)."""
    
    def __init__(self, config: LossConfig):
        super().__init__(config)
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        if error.dim() == 4:
            error_mag = error.abs().mean(dim=1).flatten(1)  # (B, H*W)
            B = error_mag.shape[0]
            
            entropies = []
            grads = []
            for b in range(B):
                em = error_mag[b]
                em_sum = em.sum() + 1e-8
                p = em / em_sum
                
                entropy = -(p * torch.log(p + 1e-8)).sum()
                max_entropy = torch.log(torch.tensor(em.numel(), dtype=torch.float32))
                normalized_entropy = entropy / (max_entropy + 1e-8)
                
                entropies.append(normalized_entropy)
                grads.append(error.abs()[b] * (1 - normalized_entropy))
            
            loss_value = 1.0 - torch.mean(torch.stack(entropies))
            gradient = torch.stack(grads).mean(dim=0)
            
        else:
            error_mag = error.abs().flatten()
            p = error_mag / (error_mag.sum() + 1e-8)
            entropy = -(p * torch.log(p + 1e-8)).sum()
            max_entropy = torch.log(torch.tensor(error_mag.numel(), dtype=torch.float32))
            normalized_entropy = entropy / (max_entropy + 1e-8)
            
            loss_value = 1.0 - normalized_entropy
            gradient = error.abs() * (1 - normalized_entropy)
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'entropy': entropy.item() if error.dim() == 4 else entropy},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=loss_value.item() if isinstance(loss_value, torch.Tensor) else loss_value
        )


class HotspotPenaltyLoss(BaseLossFunction):
    """Penalizes top-k highest error regions."""
    
    def __init__(self, config: LossConfig, n_hotspots: int = 5, penalty_scale: float = 2.0):
        super().__init__(config)
        self.n_hotspots = n_hotspots
        self.penalty_scale = penalty_scale
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        if error.dim() == 4:
            error_mag = error.abs().mean(dim=1)  # (B, H, W)
            B = error_mag.shape[0]
            
            loss_values = []
            gradients = []
            
            for b in range(B):
                em = error_mag[b].flatten()
                total_size = em.numel()
                
                # Find threshold for top-k
                k = min(self.n_hotspots, total_size // 10)
                threshold = torch.kthvalue(em, max(1, total_size - k))[0]
                
                hotspot_mask = em >= threshold
                hotspot_errors = em[hotspot_mask]
                
                if hotspot_errors.numel() > 0:
                    loss_val = (hotspot_errors ** self.penalty_scale).mean()
                    grad = torch.zeros_like(em)
                    grad[hotspot_mask] = hotspot_errors * self.penalty_scale
                    loss_values.append(loss_val)
                    gradients.append(grad.reshape_as(error_mag[b]))
                else:
                    loss_values.append(torch.tensor(0.0))
                    gradients.append(torch.zeros_like(error_mag[b]))
            
            loss_value = torch.mean(torch.stack(loss_values))
            gradient = torch.stack(gradients).mean(dim=0)
            
        else:
            error_mag = error.abs().flatten()
            total_size = error_mag.numel()
            k = min(self.n_hotspots, total_size // 10)
            threshold = torch.kthvalue(error_mag, max(1, total_size - k))[0]
            
            hotspot_mask = error_mag >= threshold
            hotspot_errors = error_mag[hotspot_mask]
            
            loss_value = (hotspot_errors ** self.penalty_scale).mean() if hotspot_errors.numel() > 0 else torch.tensor(0.0)
            
            gradient = torch.zeros_like(error)
            if hotspot_errors.numel() > 0:
                gradient.flatten()[hotspot_mask] = hotspot_errors * self.penalty_scale
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'n_hotspots': k},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(loss_value.item() if isinstance(loss_value, torch.Tensor) else loss_value, 1.0)
        )


# =============================================================================
# SPECTRAL LOSSES
# =============================================================================

class SpectralBandLoss(BaseLossFunction):
    """Loss based on error energy in different frequency bands."""
    
    def __init__(self, config: LossConfig, band: str = 'high', target_ratio: float = 0.2):
        super().__init__(config)
        self.band = band
        self.target_ratio = target_ratio
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        if error.dim() == 4:
            error_2d = error.mean(dim=1)  # (B, H, W)
            B = error_2d.shape[0]
            
            band_energies = []
            gradients = []
            
            for b in range(B):
                e = error_2d[b]
                H, W = e.shape
                
                # FFT
                fft = torch.fft.fft2(e)
                fft_shift = torch.fft.fftshift(fft)
                magnitude = torch.abs(fft_shift)
                
                # Frequency bands
                center_y, center_x = H // 2, W // 2
                y_coords = torch.arange(H, dtype=torch.float32, device=e.device).unsqueeze(1).expand(H, W)
                x_coords = torch.arange(W, dtype=torch.float32, device=e.device).unsqueeze(0).expand(H, W)
                
                distance = ((y_coords - center_y)**2 + (x_coords - center_x)**2).sqrt()
                max_dist = torch.sqrt(torch.tensor(center_y**2 + center_x**2, dtype=torch.float32, device=e.device))
                
                if self.band == 'low':
                    mask = distance < max_dist * 0.25
                elif self.band == 'mid':
                    mask = (distance >= max_dist * 0.25) & (distance < max_dist * 0.75)
                else:  # high
                    mask = distance >= max_dist * 0.75
                
                band_energy = (magnitude[mask] ** 2).sum()
                total_energy = (magnitude ** 2).sum() + 1e-8
                ratio = band_energy / total_energy
                
                band_energies.append(ratio)
                gradients.append(magnitude * ratio)  # Simplified gradient
            
            loss_value = torch.mean(torch.stack([
                torch.abs(r - torch.tensor(self.target_ratio, device=r.device)) 
                for r in band_energies
            ]))
            gradient = torch.stack(gradients).mean(dim=0)
            
        else:
            e = error
            H, W = e.shape
            
            fft = torch.fft.fft2(e)
            fft_shift = torch.fft.fftshift(fft)
            magnitude = torch.abs(fft_shift)
            
            center_y, center_x = H // 2, W // 2
            y_coords = torch.arange(H, dtype=torch.float32, device=e.device).unsqueeze(1).expand(H, W)
            x_coords = torch.arange(W, dtype=torch.float32, device=e.device).unsqueeze(0).expand(H, W)
            distance = ((y_coords - center_y)**2 + (x_coords - center_x)**2).sqrt()
            max_dist = torch.sqrt(torch.tensor(center_y**2 + center_x**2, dtype=torch.float32, device=e.device))
            
            if self.band == 'low':
                mask = distance < max_dist * 0.25
            elif self.band == 'mid':
                mask = (distance >= max_dist * 0.25) & (distance < max_dist * 0.75)
            else:
                mask = distance >= max_dist * 0.75
            
            band_energy = (magnitude[mask] ** 2).sum()
            total_energy = (magnitude ** 2).sum() + 1e-8
            ratio = band_energy / total_energy
            
            loss_value = torch.abs(ratio - self.target_ratio)
            gradient = magnitude * loss_value
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'band': self.band, 'ratio': ratio.item() if isinstance(ratio, torch.Tensor) else ratio},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(loss_value.item() if isinstance(loss_value, torch.Tensor) else loss_value, 1.0)
        )


class SpectralSkewnessLoss(BaseLossFunction):
    """Penalizes asymmetric frequency distributions."""
    
    def __init__(self, config: LossConfig, target_skewness: float = 0.0):
        super().__init__(config)
        self.target_skewness = target_skewness
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        if error.dim() == 4:
            error_2d = error.mean(dim=1)
            B = error_2d.shape[0]
            
            skewnesses = []
            for b in range(B):
                e = error_2d[b]
                fft = torch.fft.fft2(e)
                fft_shift = torch.fft.fftshift(fft)
                magnitude = torch.abs(fft_shift)
                
                h_mean = magnitude.mean(dim=0)
                h_skew = ((h_mean - h_mean.mean()) / (h_mean.std() + 1e-8) ** 3).mean()
                v_mean = magnitude.mean(dim=1)
                v_skew = ((v_mean - v_mean.mean()) / (v_mean.std() + 1e-8) ** 3).mean()
                
                skewnesses.append((abs(h_skew) + abs(v_skew)) / 2)
            
            skewness = torch.mean(torch.stack(skewnesses))
        else:
            fft = torch.fft.fft2(error)
            fft_shift = torch.fft.fftshift(fft)
            magnitude = torch.abs(fft_shift)
            
            h_mean = magnitude.mean(dim=0)
            h_skew = ((h_mean - h_mean.mean()) / (h_mean.std() + 1e-8) ** 3).mean()
            v_mean = magnitude.mean(dim=1)
            v_skew = ((v_mean - v_mean.mean()) / (v_mean.std() + 1e-8) ** 3).mean()
            
            skewness = (abs(h_skew) + abs(v_skew)) / 2
        
        loss_value = torch.abs(skewness - self.target_skewness)
        gradient = error * torch.sign(skewness - self.target_skewness)
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'skewness': skewness.item() if isinstance(skewness, torch.Tensor) else skewness},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(skewness.item() if isinstance(skewness, torch.Tensor) else skewness, 2.0) / 2.0
        )


# =============================================================================
# STATISTICAL LOSSES
# =============================================================================

class MeanErrorLoss(BaseLossFunction):
    """Penalizes non-zero mean error (bias correction)."""
    
    def __init__(self, config: LossConfig):
        super().__init__(config)
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        mean_error = error.mean()
        loss_value = torch.abs(mean_error)
        gradient = torch.sign(mean_error) * torch.ones_like(error)
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'mean': mean_error.item()},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(abs(mean_error.item()) * 10, 1.0)
        )


class SkewnessLoss(BaseLossFunction):
    """Encourages symmetric error distribution."""
    
    def __init__(self, config: LossConfig, target_skew: float = 0.0):
        super().__init__(config)
        self.target_skew = target_skew
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        error_flat = error.flatten()
        mean = error_flat.mean()
        std = error_flat.std() + 1e-8
        
        skewness = ((error_flat - mean) / std ** 3).pow(3).mean()
        loss_value = torch.abs(skewness - self.target_skew)
        
        standardized = (error_flat - mean) / std
        gradient = standardized.pow(2) * torch.sign(skewness - self.target_skew)
        gradient = gradient.reshape_as(error)
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'skewness': skewness.item()},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(abs(skewness.item()), 3.0) / 3.0
        )


class KurtosisLoss(BaseLossFunction):
    """Encourages normal-like kurtosis (≈3)."""
    
    def __init__(self, config: LossConfig, target_kurtosis: float = 3.0):
        super().__init__(config)
        self.target_kurtosis = target_kurtosis
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        error_flat = error.flatten()
        mean = error_flat.mean()
        std = error_flat.std() + 1e-8
        
        kurtosis = ((error_flat - mean) / std ** 4).pow(4).mean()
        loss_value = torch.abs(kurtosis - self.target_kurtosis)
        
        standardized = (error_flat - mean) / std
        gradient = standardized.pow(3) * torch.sign(kurtosis - self.target_kurtosis)
        gradient = gradient.reshape_as(error)
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'kurtosis': kurtosis.item()},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(abs(kurtosis.item() - 3) / 5, 1.0)
        )


class BimodalityLoss(BaseLossFunction):
    """Detects and penalizes bimodal error distributions."""
    
    def __init__(self, config: LossConfig, n_bins: int = 20):
        super().__init__(config)
        self.n_bins = n_bins
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        error_flat = error.flatten()
        
        # Compute histogram
        min_val, max_val = error_flat.min(), error_flat.max()
        bin_edges = torch.linspace(min_val, max_val, self.n_bins + 1, device=error.device)
        
        hist = torch.histc(error_flat, bins=self.n_bins, min=min_val.item(), max=max_val.item())
        hist = hist / (hist.sum() + 1e-8)
        
        # Find peaks
        peaks = []
        for i in range(1, len(hist) - 1):
            if hist[i] > hist[i-1] and hist[i] > hist[i+1]:
                peaks.append((i, hist[i].item()))
        
        if len(peaks) >= 2:
            peaks.sort(key=lambda x: x[1], reverse=True)
            peak1, peak2 = peaks[0], peaks[1]
            
            valley_start = min(peak1[0], peak2[0])
            valley_end = max(peak1[0], peak2[0])
            valley_height = min(hist[valley_start:valley_end+1].min().item(), hist.min().item())
            
            peak_heights = (peak1[1] + peak2[1]) / 2
            bimodality_score = max(0, 1.0 - valley_height / (peak_heights + 1e-8))
        else:
            bimodality_score = 0.0
        
        loss_value = torch.tensor(max(0, bimodality_score - 0.3), device=error.device)
        
        # Gradient: push valley up
        gradient = torch.zeros_like(error)
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'bimodality_score': bimodality_score, 'n_peaks': len(peaks)},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=bimodality_score
        )


# =============================================================================
# COMPONENT LOSSES
# =============================================================================

class ChannelErrorLoss(BaseLossFunction):
    """Penalizes unequal error across channels."""
    
    def __init__(self, config: LossConfig, target_equal: bool = True):
        super().__init__(config)
        self.target_equal = target_equal
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        if error.dim() == 4:
            # Per-channel error
            channel_errors = error.abs().mean(dim=(2, 3))  # (B, C)
            B, C = channel_errors.shape
            
            mean_errors = channel_errors.mean(dim=0)  # (C,)
        else:
            n_channels = error.size(0) // 784
            channel_errors = error.reshape(n_channels, 784).abs().mean(dim=1)
            mean_errors = channel_errors
        
        if self.target_equal:
            target = mean_errors.mean()
            loss_value = ((mean_errors - target) ** 2).mean()
            gradient = (mean_errors - target).unsqueeze(0).unsqueeze(-1).unsqueeze(-1).expand_as(error) * 2
        else:
            loss_value = mean_errors.mean()
            gradient = torch.sign(error)
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'mean_errors': mean_errors.mean().item()},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(loss_value.item() if isinstance(loss_value, torch.Tensor) else loss_value, 1.0)
        )


class CorrelationLoss(BaseLossFunction):
    """Penalizes correlated errors across dimensions."""
    
    def __init__(self, config: LossConfig, target_correlation: float = 0.0):
        super().__init__(config)
        self.target_correlation = target_correlation
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        # Compute correlation between spatial blocks
        error_2d = error if error.dim() == 2 else error.reshape(error.size(0), -1)
        
        if error_2d.size(0) > 1:
            corr_matrix = torch.corrcoef(error_2d + 1e-8)
            
            n = corr_matrix.shape[0]
            off_diagonal = []
            for i in range(n):
                for j in range(i+1, n):
                    off_diagonal.append(corr_matrix[i, j].abs())
            
            mean_corr = torch.mean(torch.stack(off_diagonal)) if off_diagonal else torch.tensor(0.0)
        else:
            mean_corr = torch.tensor(0.0)
        
        loss_value = torch.abs(mean_corr - self.target_correlation)
        gradient = error * mean_corr
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'mean_correlation': mean_corr.item()},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(mean_corr.item(), 1.0)
        )


# =============================================================================
# PATTERN LOSSES
# =============================================================================

class EdgeErrorLoss(BaseLossFunction):
    """Separately tracks error in edge vs smooth regions."""
    
    def __init__(self, config: LossConfig, edge_weight: float = 1.5, smooth_weight: float = 1.0):
        super().__init__(config)
        self.edge_weight = edge_weight
        self.smooth_weight = smooth_weight
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        if error.dim() == 4:
            error_2d = error.mean(dim=1)  # (B, H, W)
            B = error_2d.shape[0]
            
            if reference is not None:
                ref = reference.mean(dim=1) if reference.dim() == 4 else reference
            else:
                ref = error_2d
        else:
            error_2d = error
            B = 1
            ref = reference if reference is not None else error
        
        # Compute edges
        grad_y = torch.gradient(ref, dim=1)[0] if ref.dim() == 2 else torch.zeros_like(error_2d)
        grad_x = torch.gradient(ref, dim=2)[0] if ref.dim() == 3 else torch.zeros_like(error_2d)
        edge_magnitude = (grad_y.abs() + grad_x.abs())
        
        threshold = torch.kthvalue(edge_magnitude.flatten(), int(edge_magnitude.numel() * 0.75))[0]
        edge_mask = edge_magnitude > threshold
        smooth_mask = ~edge_mask
        
        if error_2d.dim() == 3:
            edge_error = error_2d.abs()[edge_mask].mean() if edge_mask.any() else torch.tensor(0.0)
            smooth_error = error_2d.abs()[smooth_mask].mean() if smooth_mask.any() else torch.tensor(0.0)
        else:
            edge_error = error_2d.abs()[edge_mask].mean() if edge_mask.any() else torch.tensor(0.0)
            smooth_error = error_2d.abs()[smooth_mask].mean() if smooth_mask.any() else torch.tensor(0.0)
        
        loss_value = self.edge_weight * edge_error + self.smooth_weight * smooth_error
        
        gradient = torch.zeros_like(error_2d)
        if error_2d.dim() == 3:
            gradient[edge_mask] = self.edge_weight * torch.sign(error_2d[edge_mask])
            gradient[smooth_mask] = self.smooth_weight * torch.sign(error_2d[smooth_mask])
        else:
            gradient[edge_mask] = self.edge_weight * torch.sign(error_2d[edge_mask])
            gradient[smooth_mask] = self.smooth_weight * torch.sign(error_2d[smooth_mask])
        
        if error.dim() == 4:
            gradient = gradient.unsqueeze(1).expand_as(error).mean(dim=1)
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'edge_error': edge_error.item(), 'smooth_error': smooth_error.item()},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(loss_value.item() if isinstance(loss_value, torch.Tensor) else loss_value, 1.0)
        )


# =============================================================================
# ADVERSARIAL LOSSES
# =============================================================================

class ErrorDiscriminator(nn.Module):
    """Discriminator that judges error quality."""
    
    def __init__(self, input_dim: int = 784, hidden_dim: int = 64):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(input_dim, hidden_dim),
            nn.LeakyReLU(0.2),
            nn.Linear(hidden_dim, hidden_dim),
            nn.LeakyReLU(0.2),
            nn.Linear(hidden_dim, 1),
            nn.Sigmoid()
        )
        
    def forward(self, error: torch.Tensor) -> torch.Tensor:
        error_flat = error.flatten(1)[:, :784]  # Handle batch
        return self.net(error_flat)


class AdversarialLoss(BaseLossFunction):
    """Adversarial loss using discriminator to judge error quality."""
    
    def __init__(self, config: LossConfig, discriminator: nn.Module = None):
        super().__init__(config)
        self.discriminator = discriminator
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        if self.discriminator is None:
            loss_value = (error ** 2).mean()
            gradient = 2 * error
        else:
            quality_score = self.discriminator(error)
            loss_value = 1.0 - quality_score.mean()
            
            # Gradient from discriminator feedback
            score_scale = 1 - quality_score
            while score_scale.dim() < error.dim():
                score_scale = score_scale.unsqueeze(-1)
            gradient = error * score_scale
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'quality_score': quality_score.mean().item() if self.discriminator else 0.5},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=loss_value.item() if isinstance(loss_value, torch.Tensor) else loss_value
        )


# =============================================================================
# REGULARIZATION LOSSES
# =============================================================================

class GradientPenaltyLoss(BaseLossFunction):
    """Penalizes large gradients (unstable predictions)."""
    
    def __init__(self, config: LossConfig, penalty: float = 1.0):
        super().__init__(config)
        self.penalty = penalty
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        if error.dim() == 4:
            error_2d = error.mean(dim=1)  # (B, H, W)
            
            grad_y = torch.gradient(error_2d, dim=1)[0]
            grad_x = torch.gradient(error_2d, dim=2)[0]
            grad_mag = (grad_y**2 + grad_x**2).sqrt().mean()
            
            # Laplacian
            laplacian = torch.gradient(grad_y, dim=1)[0] + torch.gradient(grad_x, dim=2)[0]
        else:
            grad_y, grad_x = torch.gradient(error, dim=0), torch.gradient(error, dim=1)
            grad_mag = (grad_y**2 + grad_x**2).sqrt().mean()
            laplacian = torch.gradient(grad_y, dim=0)[0] + torch.gradient(grad_x, dim=1)[0]
        
        loss_value = grad_mag ** 2
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=laplacian,
            diagnostics={'mean_gradient': grad_mag.item()},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(loss_value.item() if isinstance(loss_value, torch.Tensor) else loss_value, 1.0)
        )


class LipschitzPenaltyLoss(BaseLossFunction):
    """Ensures smooth function behavior (K-Lipschitz)."""
    
    def __init__(self, config: LossConfig, K: float = 1.0):
        super().__init__(config)
        self.K = K
        
    def forward(self, error: torch.Tensor, X: torch.Tensor = None, 
                reference: torch.Tensor = None) -> LossResult:
        if error.dim() == 4:
            error_2d = error.mean(dim=1)
        else:
            error_2d = error
        
        H, W = error_2d.shape[-2:]
        
        h_diff = error_2d[..., :, 1:] - error_2d[..., :, :-1]
        v_diff = error_2d[..., 1:, :] - error_2d[..., :-1, :]
        
        max_diff = max(
            torch.abs(h_diff).max().item() if h_diff.numel() > 0 else 0,
            torch.abs(v_diff).max().item() if v_diff.numel() > 0 else 0
        )
        
        loss_value = torch.relu(torch.tensor(max_diff - self.K, device=error.device))
        gradient = torch.zeros_like(error)
        
        return LossResult(
            name=self.config.name,
            value=loss_value,
            gradient=gradient,
            diagnostics={'max_diff': max_diff, 'K': self.K},
            category=self.config.category,
            weight=self.config.weight,
            normalized_value=min(loss_value.item() if isinstance(loss_value, torch.Tensor) else loss_value, 1.0)
        )


# =============================================================================
# MULTI-LOSS OPTIMIZER
# =============================================================================

class MultiLossOptimizer:
    """Optimizer that combines multiple diagnostic losses."""
    
    def __init__(self, shape: Tuple[int, ...] = (1, 28, 28)):
        self.shape = shape
        self.loss_functions = nn.ModuleDict()
        self.discriminator = ErrorDiscriminator(input_dim=784)
        self.optimizer_disc = optim.Adam(self.discriminator.parameters(), lr=0.001)
        
        self._initialize_loss_functions()
        
    def _initialize_loss_functions(self):
        """Initialize all loss functions."""
        
        # Spatial losses
        self.loss_functions['spatial_concentration'] = SpatialConcentrationLoss(
            LossConfig('spatial_concentration', LossCategory.SPATIAL, weight=0.5)
        )
        self.loss_functions['spatial_entropy'] = SpatialEntropyLoss(
            LossConfig('spatial_entropy', LossCategory.SPATIAL, weight=0.3)
        )
        self.loss_functions['hotspot_penalty'] = HotspotPenaltyLoss(
            LossConfig('hotspot_penalty', LossCategory.SPATIAL, weight=0.5),
            n_hotspots=5, penalty_scale=2.0
        )
        
        # Spectral losses
        self.loss_functions['spectral_high'] = SpectralBandLoss(
            LossConfig('spectral_high', LossCategory.SPECTRAL, weight=0.3),
            band='high', target_ratio=0.2
        )
        self.loss_functions['spectral_mid'] = SpectralBandLoss(
            LossConfig('spectral_mid', LossCategory.SPECTRAL, weight=0.3),
            band='mid', target_ratio=0.5
        )
        self.loss_functions['spectral_low'] = SpectralBandLoss(
            LossConfig('spectral_low', LossCategory.SPECTRAL, weight=0.3),
            band='low', target_ratio=0.3
        )
        self.loss_functions['spectral_skewness'] = SpectralSkewnessLoss(
            LossConfig('spectral_skewness', LossCategory.SPECTRAL, weight=0.3)
        )
        
        # Statistical losses
        self.loss_functions['mean_error'] = MeanErrorLoss(
            LossConfig('mean_error', LossCategory.STATISTICAL, weight=0.5)
        )
        self.loss_functions['skewness'] = SkewnessLoss(
            LossConfig('skewness', LossCategory.STATISTICAL, weight=0.3)
        )
        self.loss_functions['kurtosis'] = KurtosisLoss(
            LossConfig('kurtosis', LossCategory.STATISTICAL, weight=0.3),
            target_kurtosis=3.0
        )
        self.loss_functions['bimodality'] = BimodalityLoss(
            LossConfig('bimodality', LossCategory.STATISTICAL, weight=0.4)
        )
        
        # Component losses
        self.loss_functions['channel_error'] = ChannelErrorLoss(
            LossConfig('channel_error', LossCategory.COMPONENT, weight=0.3)
        )
        self.loss_functions['correlation'] = CorrelationLoss(
            LossConfig('correlation', LossCategory.COMPONENT, weight=0.3)
        )
        
        # Pattern losses
        self.loss_functions['edge_error'] = EdgeErrorLoss(
            LossConfig('edge_error', LossCategory.PATTERN, weight=0.4),
            edge_weight=1.5, smooth_weight=1.0
        )
        
        # Adversarial loss
        self.loss_functions['adversarial'] = AdversarialLoss(
            LossConfig('adversarial', LossCategory.ADVERSARIAL, weight=0.5),
            discriminator=self.discriminator
        )
        
        # Regularization losses
        self.loss_functions['gradient_penalty'] = GradientPenaltyLoss(
            LossConfig('gradient_penalty', LossCategory.REGULARIZATION, weight=0.2)
        )
        self.loss_functions['lipschitz_penalty'] = LipschitzPenaltyLoss(
            LossConfig('lipschitz_penalty', LossCategory.REGULARIZATION, weight=0.2)
        )
    
    def compute_all_losses(
        self,
        error: torch.Tensor,
        X: torch.Tensor = None,
        reference: torch.Tensor = None
    ) -> Dict[str, LossResult]:
        """Compute all enabled losses."""
        results = {}
        
        for name, loss_fn in self.loss_functions.items():
            try:
                with torch.no_grad():
                    result = loss_fn.forward(error, X, reference)
                results[name] = result
            except Exception as e:
                logger.warning(f"Loss {name} failed: {e}")
        
        return results
    
    def aggregate_losses(
        self,
        losses: Dict[str, LossResult],
        dynamic_weighting: bool = True
    ) -> Tuple[torch.Tensor, torch.Tensor, MultiLossState]:
        """Aggregate all losses into single loss and gradient."""
        total_loss = torch.tensor(0.0, device=device)
        combined_gradient = torch.zeros_like(list(losses.values())[0].gradient)
        weighted_contributions = {}
        gradient_norms = {}
        
        # Dynamic weight adjustment
        if dynamic_weighting:
            loss_values = [l.normalized_value for l in losses.values() if l.weight > 0]
            if loss_values:
                max_loss = max(loss_values)
                min_loss = min(loss_values)
                
                for name, loss in losses.items():
                    if loss.weight > 0 and max_loss > min_loss:
                        normalized = (loss.normalized_value - min_loss) / (max_loss - min_loss + 1e-8)
                        loss.weight = loss.weight * (1 + normalized)
        
        # Aggregate
        for name, loss in losses.items():
            if loss.weight > 0:
                weighted_loss = loss.value * loss.weight
                total_loss = total_loss + weighted_loss
                weighted_contributions[name] = weighted_loss.item()
                
                grad_norm = loss.gradient.norm() + 1e-8
                normalized_gradient = loss.gradient / grad_norm
                combined_gradient = combined_gradient + normalized_gradient * loss.weight
                gradient_norms[name] = grad_norm.item()
        
        # Find dominant loss
        dominant_loss = max(weighted_contributions, key=weighted_contributions.get) if weighted_contributions else 'mse'
        
        state = MultiLossState(
            losses=losses,
            total_loss=total_loss,
            weighted_contributions=weighted_contributions,
            dominant_loss=dominant_loss,
            gradient_norms=gradient_norms,
            optimization_advice=self._generate_advice(losses, dominant_loss)
        )
        
        return total_loss, combined_gradient, state
    
    def _generate_advice(self, losses: Dict[str, LossResult], dominant_loss: str) -> List[str]:
        """Generate optimization advice."""
        advice = []
        
        spatial_losses = [l for l in losses.values() if l.category == LossCategory.SPATIAL]
        if spatial_losses:
            avg_spatial = np.mean([l.normalized_value for l in spatial_losses])
            if avg_spatial > 0.5:
                advice.append("Focus on reducing spatial concentration - errors are clustered")
        
        if dominant_loss == 'hotspot_penalty':
            advice.append("Focus training on identified high-error regions")
        elif dominant_loss == 'edge_error':
            advice.append("Improve edge rendering accuracy")
        
        return advice


# =============================================================================
# MLP CLASSIFIER WITH SPDER ACTICATION
# =============================================================================

class SPDERActivation(nn.Module):
    """SPDER: sin(x) * sqrt(|x|) - Periodic with damping."""
    
    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return torch.sin(x) * torch.sqrt(torch.abs(x) + 1e-8)


class SPDERDerivative(nn.Module):
    """Derivative of SPDER for backpropagation."""
    
    def forward(self, x: torch.Tensor) -> torch.Tensor:
        abs_x = torch.abs(x) + 1e-8
        sqrt_abs_x = torch.sqrt(abs_x)
        sign_x = torch.sign(x)
        sign_x = torch.where(sign_x == 0, torch.ones_like(sign_x), sign_x)
        return sqrt_abs_x * torch.cos(x) + (sign_x / (2 * sqrt_abs_x)) * torch.sin(x)


class MultiLossMLP(nn.Module):
    """MLP Classifier with Multi-Loss Diagnostic Learning."""
    
    def __init__(self, input_size: int = 784, hidden_size: int = 256, 
                 output_size: int = 10, use_spder: bool = True):
        super().__init__()
        
        self.input_size = input_size
        self.hidden_size = hidden_size
        self.output_size = output_size
        self.use_spder = use_spder
        
        # Layers
        self.fc1 = nn.Linear(input_size, hidden_size)
        self.fc2 = nn.Linear(hidden_size, output_size)
        
        # Activations
        if use_spder:
            self.activation = SPDERActivation()
            self.activation_deriv = SPDERDerivative()
        else:
            self.activation = nn.ReLU()
            self.activation_deriv = None
        
        # Multi-loss optimizer
        self.multi_loss_optimizer = MultiLossOptimizer(shape=(1, 28, 28))
        
        # Gradient tracking
        self.register_hook()
        
    def register_hook(self):
        """Register hooks for gradient analysis."""
        pass
        
    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = x.view(x.size(0), -1)
        self.z1 = self.fc1(x)
        self.a1 = self.activation(self.z1)
        self.z2 = self.fc2(self.a1)
        
        return self.z2
    
    def compute_multi_loss_gradient(self, X: torch.Tensor, y_pred: torch.Tensor) -> torch.Tensor:
        """Compute gradient from multi-loss diagnostic system."""
        # Use input statistics for diagnostic gradient
        error_for_analysis = X.view(X.size(0), 1, 28, 28)  # Reshape for analysis
        
        losses = self.multi_loss_optimizer.compute_all_losses(
            error_for_analysis,
            X=X,
            reference=None
        )
        
        _, combined_gradient, state = self.multi_loss_optimizer.aggregate_losses(losses)
        
        return combined_gradient, state
    
    def get_spder_gradient(self, upstream_grad: torch.Tensor) -> torch.Tensor:
        """Get gradient through SPDER activation."""
        if self.use_spder and self.activation_deriv is not None:
            spder_deriv = self.activation_deriv(self.z1)
            return upstream_grad * spder_deriv
        return upstream_grad
    
    def predict(self, x: torch.Tensor) -> torch.Tensor:
        with torch.no_grad():
            return torch.argmax(self.forward(x), dim=1)
    
    def accuracy(self, x: torch.Tensor, y: torch.Tensor) -> float:
        return (self.predict(x) == y).float().mean().item()


# =============================================================================
# TRAINING WITH MULTI-LOSS DIAGNOSTICS
# =============================================================================

def train_with_multi_loss(
    model: MultiLossMLP,
    train_loader: DataLoader,
    test_loader: DataLoader,
    epochs: int = 20,
    lr: float = 0.001,
    multi_loss_weight: float = 0.1
):
    """Train model with multi-loss diagnostic learning."""
    
    optimizer = optim.Adam(model.parameters(), lr=lr)
    
    # Loss for classification
    ce_loss = nn.CrossEntropyLoss()
    
    loss_history = defaultdict(list)
    
    for epoch in range(epochs):
        model.train()
        epoch_losses = defaultdict(float)
        epoch_correct = 0
        epoch_total = 0
        
        for batch_idx, (data, target) in enumerate(train_loader):
            data, target = data.to(device), target.to(device)
            
            optimizer.zero_grad()
            
            # Forward pass
            logits = model(data)
            
            # Classification loss
            loss_ce = ce_loss(logits, target)
            
            # Multi-loss diagnostic state for monitoring.
            # These losses are currently computed from the input tensor under
            # no_grad, so they do not provide a trainable signal to the MLP.
            _, loss_state = model.compute_multi_loss_gradient(data, logits)
            
            # Optimize only the trainable classification objective.
            total_loss = loss_ce
            
            # Backward pass
            total_loss.backward()
            
            # Apply gradient clipping
            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
            
            optimizer.step()

            print(total_loss.item())
            
            # Track statistics
            epoch_losses['ce'] += loss_ce.item()
            epoch_losses['total'] += total_loss.item()
            epoch_losses['multi'] += loss_state.total_loss.item()
            
            # Track accuracy
            pred = logits.argmax(dim=1)
            epoch_correct += pred.eq(target).sum().item()
            epoch_total += target.size(0)
            
            # Log batch progress
            if batch_idx % 100 == 0:
                logger.info(
                    f"Epoch {epoch} [{batch_idx}/{len(train_loader)}] "
                    f"CE: {loss_ce.item():.4f} Multi(diag): {loss_state.total_loss.item():.4f}"
                )
        
        # Epoch summary
        acc = epoch_correct / epoch_total
        logger.info(f"\n=== Epoch {epoch} Summary ===")
        logger.info(f"Training Accuracy: {acc:.4f}")
        logger.info(
            f"Losses - CE: {epoch_losses['ce']/len(train_loader):.4f}, "
            f"Multi(diag): {epoch_losses['multi']/len(train_loader):.4f}"
        )
        
        # Log top losses
        sorted_losses = sorted(
            loss_state.weighted_contributions.items(),
            key=lambda x: x[1],
            reverse=True
        )[:3]
        logger.info(f"Top contributors: {sorted_losses}")
        
        if loss_state.optimization_advice:
            logger.info(f"Advice: {loss_state.optimization_advice[0]}")
        
        # Test evaluation
        test_acc = evaluate(model, test_loader)
        logger.info(f"Test Accuracy: {test_acc:.4f}\n")
        
        # Store history
        loss_history['train_acc'].append(acc)
        loss_history['test_acc'].append(test_acc)
        loss_history['ce_loss'].append(epoch_losses['ce'] / len(train_loader))
        loss_history['multi_loss'].append(epoch_losses['multi'] / len(train_loader))
    
    return loss_history


def evaluate(model: MultiLossMLP, data_loader: DataLoader) -> float:
    """Evaluate model on given data loader."""
    model.eval()
    correct = 0
    total = 0
    
    with torch.no_grad():
        for data, target in data_loader:
            data, target = data.to(device), target.to(device)
            output = model(data)
            pred = output.argmax(dim=1)
            correct += pred.eq(target).sum().item()
            total += target.size(0)
    
    return correct / total


# =============================================================================
# ADVERSARIAL TRAINING WITH DISCRIMINATOR
# =============================================================================

def train_adversarial(
    model: MultiLossMLP,
    discriminator: nn.Module,
    train_loader: DataLoader,
    epochs: int = 20
):
    """Adversarial training with discriminator for error quality."""
    
    optimizer_G = optim.Adam(model.parameters(), lr=0.001)
    optimizer_D = optim.Adam(discriminator.parameters(), lr=0.001)
    
    ce_loss = nn.CrossEntropyLoss()
    
    for epoch in range(epochs):
        model.train()
        discriminator.train()
        
        for batch_idx, (data, target) in enumerate(train_loader):
            data, target = data.to(device), target.to(device)
            batch_size = data.size(0)
            
            # ========== Train Discriminator ==========
            optimizer_D.zero_grad()
            
            # Real errors (from classification mistakes)
            output = model(data)
            pred = output.argmax(dim=1)
            errors = (pred.float() - target.float()).abs().reshape(batch_size, 1, 28, 28)
            
            # Fake errors (model output patterns)
            fake_errors = output - 1.0 / model.output_size  # Uniform prediction
            
            # Discriminator scores
            real_score = discriminator(errors)
            fake_score = discriminator(fake_errors.view(batch_size, -1)[:, :784].unsqueeze(1).expand(-1, 1, 28, 28))
            
            # D loss
            d_loss = -(torch.log(real_score + 1e-8) + torch.log(1 - fake_score + 1e-8)).mean()
            d_loss.backward()
            optimizer_D.step()
            
            # ========== Train Generator (MLP) ==========
            optimizer_G.zero_grad()
            
            # Classification loss
            loss_ce = ce_loss(output, target)
            
            # Adversarial loss (fool discriminator)
            g_loss = -torch.log(discriminator(fake_errors.view(batch_size, -1)[:, :784].unsqueeze(1).expand(-1, 1, 28, 28)) + 1e-8).mean()
            
            total_g_loss = loss_ce + 0.1 * g_loss
            total_g_loss.backward()
            optimizer_G.step()
            
            if batch_idx % 200 == 0:
                logger.info(f"Epoch {epoch} [{batch_idx}] D_loss: {d_loss.item():.4f} "
                          f"G_loss: {g_loss.item():.4f}")
        
        # Evaluate
        train_acc = evaluate(model, train_loader)
        logger.info(f"Epoch {epoch}: Train Accuracy = {train_acc:.4f}")


# =============================================================================
# MAIN EXECUTION
# =============================================================================

if __name__ == "__main__":
    logger.info("=== PyTorch Multi-Loss Diagnostic Learning with MNIST ===\n")
    
    # ==========================================================================
    # DATA LOADING (MNIST via torchvision)
    # ==========================================================================
    
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.5,), (0.5,))
    ])
    
    # Load MNIST
    train_dataset = torchvision.datasets.MNIST(
        root='./data',
        train=True,
        download=True,
        transform=transform
    )
    
    test_dataset = torchvision.datasets.MNIST(
        root='./data',
        train=False,
        download=True,
        transform=transform
    )
    
    train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=2)
    test_loader = DataLoader(test_dataset, batch_size=128, shuffle=False, num_workers=2)
    
    logger.info(f"Training samples: {len(train_dataset)}")
    logger.info(f"Test samples: {len(test_dataset)}")
    
    # ==========================================================================
    # MODEL INITIALIZATION
    # ==========================================================================
    
    model = MultiLossMLP(
        input_size=784,
        hidden_size=256,
        output_size=10,
        use_spder=False
    ).to(device)
    
    discriminator = model.multi_loss_optimizer.discriminator.to(device)
    
    logger.info(f"Model parameters: {sum(p.numel() for p in model.parameters()):,}")
    logger.info(f"Discriminator parameters: {sum(p.numel() for p in discriminator.parameters()):,}")
    
    # ==========================================================================
    # TRAINING MODES
    # ==========================================================================
    
    print("\n" + "="*60)
    print("TRAINING MODE 1: Standard Multi-Loss Training")
    print("="*60 + "\n")
    
    loss_history = train_with_multi_loss(
        model=model,
        train_loader=train_loader,
        test_loader=test_loader,
        epochs=10,
        lr=0.001,
        multi_loss_weight=0.1
    )
    
    # ==========================================================================
    # EVALUATION
    # ==========================================================================
    
    print("\n" + "="*60)
    print("FINAL EVALUATION")
    print("="*60 + "\n")
    
    final_test_acc = evaluate(model, test_loader)
    logger.info(f"Final Test Accuracy: {final_test_acc:.4f}")
    
    # Show loss history
    print("\n--- Training History ---")
    for i, (train_acc, test_acc) in enumerate(zip(loss_history['train_acc'], loss_history['test_acc'])):
        print(f"Epoch {i}: Train={train_acc:.4f}, Test={test_acc:.4f}")
    
    # ==========================================================================
    # OPTIONAL: ADVERSARIAL TRAINING
    # ==========================================================================
    
    print("\n" + "="*60)
    print("TRAINING MODE 2: Adversarial Training (Optional)")
    print("="*60 + "\n")
    
    # Uncomment to run adversarial training
    # train_adversarial(model, discriminator, train_loader, epochs=10)
    
    # ==========================================================================
    # DEMONSTRATION: Multi-Loss Analysis
    # ==========================================================================
    
    print("\n" + "="*60)
    print("MULTI-LOSS DIAGNOSTIC ANALYSIS")
    print("="*60 + "\n")
    
    model.eval()
    
    # Analyze a batch
    sample_data, sample_targets = next(iter(test_loader))
    sample_data = sample_data.to(device)
    
    with torch.no_grad():
        output = model(sample_data)
        
        # Reshape for loss analysis
        error_for_analysis = sample_data.view(sample_data.size(0), 1, 28, 28)
        
        # Compute all losses
        losses = model.multi_loss_optimizer.compute_all_losses(error_for_analysis)
        
        print("Loss breakdown for a sample batch:")
        for name, loss in sorted(losses.items(), key=lambda x: x[1].value.item(), reverse=True):
            if loss.value.item() > 0.01:
                print(f"  {name:25s}: {loss.value.item():.4f} (weight: {loss.weight:.2f})")
    
    # ==========================================================================
    # SAVE/LOAD
    # ==========================================================================
    
    torch.save({
        'model_state_dict': model.state_dict(),
        'discriminator_state_dict': discriminator.state_dict(),
        'loss_history': dict(loss_history)
    }, 'multi_loss_mlp_mnist.pth')
    
    logger.info("\nModel saved to 'multi_loss_mlp_mnist.pth'")
    
    # Load
    checkpoint = torch.load('multi_loss_mlp_mnist.pth')
    model.load_state_dict(checkpoint['model_state_dict'])
    
    final_acc = evaluate(model, test_loader)
    logger.info(f"Loaded model test accuracy: {final_acc:.4f}")
