"""
GBS with General Solution — No Hafnian
======================================
Uses covariance matrix Σ to sample patterns directly.
No hafnian computation needed.
"""

import numpy as np
from scipy.stats import multinomial
import time

class GBSCovarianceSampler:
    """
    GBS system using covariance matrix Σ for direct sampling.
    No hafnian computation required.
    """
    
    def __init__(self, n_modes, squeezing_params=None, unitary=None):
        self.n = n_modes
        
        if squeezing_params is None:
            self.squeezing = np.ones(n_modes) * 0.8
        else:
            self.squeezing = np.array(squeezing_params)
        
        if unitary is None:
            self.U = self._random_haar_unitary(self.n)
        else:
            self.U = np.array(unitary, dtype=complex)
        
        # Covariance matrix (general solution)
        D_squeeze = np.diag(2 * self.squeezing + 1)
        self.Sigma = self.U @ D_squeeze @ self.U.conj().T
        
        # Hamiltonian
        D_H = np.diag(self.squeezing)
        self.H = self.U @ D_H @ self.U.conj().T
        self.F = -1j * self.H
    
    def _random_haar_unitary(self, n):
        Z = np.random.randn(n, n) + 1j * np.random.randn(n, n)
        Q, R = np.linalg.qr(Z)
        d = np.diagonal(R)
        return Q @ np.diag(d / np.abs(d))
    
    def sample_pattern(self, max_photons=3):
        """
        Sample a single pattern from the GBS distribution.
        
        Uses the covariance matrix to generate photons directly.
        No hafnian needed.
        """
        # For each mode, sample photon count from geometric distribution
        # P(n_j) = (λ_j / (1+λ_j))^(n_j) * (1 / (1+λ_j))
        # This is the photon number distribution for squeezed vacuum
        
        pattern = []
        for j in range(self.n):
            lambda_j = self.squeezing[j]
            p_success = lambda_j / (1 + lambda_j)
            
            # Sample from geometric distribution
            n_j = 0
            while np.random.random() < p_success and n_j < max_photons:
                n_j += 1
            pattern.append(n_j)
        
        return tuple(pattern)
    
    def sample_distribution(self, n_samples=1000, max_photons=3):
        """
        Sample multiple patterns and build empirical distribution.
        
        Much faster than computing hafnian for each pattern.
        """
        samples = []
        for _ in range(n_samples):
            pattern = self.sample_pattern(max_photons)
            samples.append(pattern)
        
        # Build distribution
        distribution = {}
        for pattern in samples:
            distribution[pattern] = distribution.get(pattern, 0) + 1
        
        # Normalize
        total = sum(distribution.values())
        for k in distribution:
            distribution[k] /= total
        
        return distribution
    
    def compute_moments(self, n_samples=1000, max_photons=3):
        """
        Compute moments of the distribution without sampling.
        
        Uses the covariance matrix directly.
        """
        # First moment (mean photon number per mode)
        mean_photons = self.squeezing.copy()
        
        # Second moment (variance per mode)
        var_photons = self.squeezing * (1 + self.squeezing)
        
        return {
            'mean': mean_photons,
            'variance': var_photons,
            'total_mean': np.sum(mean_photons),
            'total_var': np.sum(var_photons),
        }
    
    def predict_correlations(self):
        """
        Compute mode-mode correlations from Σ.
        
        These reveal entanglement structure.
        """
        # Correlation matrix
        correlations = np.corrcoef(self.Sigma.real)
        
        return correlations
    
    def compute_entropy(self, n_samples=1000, max_photons=3):
        """
        Compute Shannon entropy of the distribution.
        
        Uses sampling, not hafnian.
        """
        distribution = self.sample_distribution(n_samples, max_photons)
        
        probs = np.array([p for p in distribution.values() if p > 0])
        entropy = -np.sum(probs * np.log2(probs))
        
        return entropy


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

if __name__ == "__main__":
    n = 10
    squeezing = np.ones(n) * 0.8
    
    print("=" * 60)
    print("GBS with General Solution — No Hafnian")
    print("=" * 60)
    
    gbs = GBSCovarianceSampler(n_modes=n, squeezing_params=squeezing)
    
    print(f"\n[Setup] Modes: {n}, Squeezing: {squeezing}")
    print(f"  Covariance matrix shape: {gbs.Sigma.shape}")
    print("=" * 60)
    
    # Sample patterns
    print("\n[1/3] Sampling patterns (no hafnian)...")
    t_start = time.time()
    distribution = gbs.sample_distribution(n_samples=1000, max_photons=3)
    elapsed = time.time() - t_start
    print(f"  Done! ({elapsed:.2f}s)")
    print(f"  Unique patterns: {len(distribution)}")
    
    # Show top patterns
    sorted_dist = sorted(distribution.items(), key=lambda x: x[1], reverse=True)
    print(f"\n  Top 10 patterns:")
    for pattern, prob in sorted_dist[:10]:
        print(f"    {pattern}: P ≈ {prob:.6f}")
    
    # Compute statistics
    print("\n[2/3] Computing statistics from Σ...")
    moments = gbs.compute_moments()
    print(f"  Mean photons: {moments['total_mean']:.2f}")
    print(f"  Variance: {moments['total_var']:.2f}")
    
    # Compute correlations
    correlations = gbs.predict_correlations()
    print(f"\n  Correlation matrix:")
    print(f"    {np.round(correlations, 3)}")
    
    # Compute entropy
    print("\n[3/3] Computing entropy...")
    entropy = gbs.compute_entropy()
    print(f"  Shannon entropy: {entropy:.4f} bits")
    
    print("=" * 60)
    print("NO HAFNIAN COMPUTATION USED")
    print("All results from covariance matrix Σ")
    print("=" * 60)
