import numpy as np
from scipy.io.wavfile import read
import sys

# Import original classifier
sys.path.insert(0, '/home/per/Documents/python/state to state')
from c00 import MLPClassifier
from self_learning_wrapper import SelfLearningWrapper

# ============================================================
# Self-Learning Training Loop
# ============================================================
# Philosophy: "What doesn't break is an emergent truth"
# The classifier discovers patterns in random noise that are stable,
# then uses these emergent truths to guide its own learning.
# ============================================================

def main():
    # Load data
    print("Loading data...")
    X_train = read('../X_train.wav')[1].reshape(-1, 784)
    y_train = (read('../y_train.wav')[1] * 9).astype(int)
    X_test = read('../X_test.wav')[1].reshape(-1, 784)
    y_test = (read('../y_test.wav')[1] * 9).astype(int)
    
    # Initialize classifier
    learning_rate = np.random.rand(4)#np.array([0.01, 0.01, 0.01, 0.01])  # Fixed for stability
    f = MLPClassifier(input_size=784, hidden_size=100, output_size=10, 
                      learning_rate=learning_rate)
    
    # Wrap with self-learning
    wrapper = SelfLearningWrapper(
        classifier=f,
        discovery_interval=20,      # Discover truths every 20 iterations
        self_train_interval=10,     # Self-train every 10 iterations
        stability_threshold=0.85    # High threshold for reliable truths
    )
    
    print("\n" + "="*60)
    print("SELF-LEARNING TRAINING")
    print("="*60)
    print(f"Training samples: {len(X_train)}")
    print(f"Test samples: {len(X_test)}")
    print(f"Discovery interval: Every {wrapper.discovery_interval} iterations")
    print(f"Self-training interval: Every {wrapper.self_train_interval} iterations")
    print("="*60 + "\n")
    
    # Phase 1: Pure supervised learning (warm-up)
    print("Phase 1: Supervised warm-up (500 iterations)...")
    for i in range(500):
        idx = np.random.randint(0, 60000, 100)
        X = X_train[idx]
        yt = y_train[idx]
        f.update(X, np.eye(10)[yt])
        
        if i % 100 == 0:
            acc = f.score(X_test[:1000], y_test[:1000])
            print(f"  Iter {i:4d} | Test Accuracy: {acc:.4f}")
    
    print(f"\nWarm-up complete. Test accuracy: {f.score(X_test[:1000], y_test[:1000]):.4f}\n")
    
    # Phase 2: Combined supervised + self-learning
    print("Phase 2: Combined learning (2000 iterations)...")
    print("-" * 60)
    
    best_test_acc = 0
    patience_counter = 0
    
    for i in range(20000):
        idx = np.random.randint(0, 60000, 100)
        X = X_train[idx]
        yt = y_train[idx]
        
        results = wrapper.train_step(X, np.eye(10)[yt])
        
        # Monitor progress
        if i % 200 == 0:
            test_acc = f.score(X_test[:1000], y_test[:1000])
            stats = wrapper.learner.get_truth_statistics()
            
            print(f"Iter {i:5d} | "
                  f"Test Acc: {test_acc:.4f} | "
                  f"Sup Loss: {results['supervised_loss']:.4f} | "
                  f"Stable Found: {results['n_stable']:3d} | "
                  f"Buffer: {stats['buffer_usage']:.1%} | "
                  f"Mean Stability: {stats['mean_stability']:.3f}")
            
            # Early stopping check
            if test_acc > best_test_acc:
                best_test_acc = test_acc
                patience_counter = 0
                # Save best weights
                best_W1 = f.W1.copy()
                best_b1 = f.b1.copy()
                best_W2 = f.W2.copy()
                best_b2 = f.b2.copy()
            else:
                patience_counter += 1
            
            if patience_counter > 10 and i > 1000:
                print(f"\nEarly stopping at iteration {i}")
                break
    
    # Restore best weights
    f.W1 = best_W1
    f.b1 = best_b1
    f.W2 = best_W2
    f.b2 = best_b2
    
    print("\n" + "="*60)
    print("FINAL RESULTS")
    print("="*60)
    
    final_acc = f.score(X_test[:1000], y_test[:1000])
    print(f"Final Test Accuracy: {final_acc:.4f}")
    print(f"Best Test Accuracy: {best_test_acc:.4f}")
    
    # Show emergent truths discovered
    truth_stats = wrapper.learner.get_truth_statistics()
    print(f"\nEmergent Truths Discovered:")
    print(f"  Total stable samples: {truth_stats['total_stable_found']}")
    print(f"  Buffer usage: {truth_stats['buffer_usage']:.1%}")
    print(f"  Mean stability: {truth_stats['mean_stability']:.3f}")
    print(f"\nClass distribution in truth buffer:")
    for cls, count in truth_stats['class_distribution'].items():
        print(f"  Class {cls}: {count:4d}")
    
    # Test prediction with confidence
    print("\n" + "="*60)
    print("PREDICTION WITH CONFIDENCE")
    print("="*60)
    
    X_sample = X_test[:10]
    predictions, confidence, stability = wrapper.predict_with_confidence(X_sample)
    
    print(f"{'Sample':<8} {'Prediction':<12} {'Stability':<10} {'True Label':<12}")
    print("-" * 42)
    for i in range(10):
        print(f"{i:<8} {predictions[i]:<12} {stability[i]:.3f}      {y_test[i]:<12}")
    
    print("\nDone!")


if __name__ == "__main__":
    main()
