
    ^i	                     8    d dl Zd dlmZ d dlmZ  G d d      Zy)    N)EmergentTruthLearner)StabilityAnalyzerc                   B    e Zd ZdZ	 	 	 d	dZd Zd Zd Z	 	 d
dZd Z	y)SelfLearningWrapperaI  
    Wraps MLPClassifier to add self-learning capabilities.
    
    Usage:
        wrapper = SelfLearningWrapper(f)
        wrapper.train_with_self_learning(X_train, y_train, X_test, y_test)
        
        # Use enhanced prediction
        pred = wrapper.predict(X)
        confidence = wrapper.predict_with_confidence(X)
    c                     || _         t        |      | _        t        |      | _        || _        || _        || _        d| _        g g g g d| _	        y)a.  
        Args:
            classifier: MLPClassifier instance (your f from c00.py)
            discovery_interval: Discover emergent truths every N iterations
            self_train_interval: Self-train every N iterations
            stability_threshold: Minimum stability to consider as truth
        r   )supervised_lossself_learning_lossstability_scoresn_stable_foundN)

classifierr   learnerr   analyzerdiscovery_intervalself_train_intervalstability_threshold	iterationhistory)selfr   r   r   r   s        B/home/per/Documents/python/state to state/self_learning_wrapper.py__init__zSelfLearningWrapper.__init__   sW     %+J7)*5"4#6 #6 !"$ " 	
    c                 8    | j                   j                  |      S )zStandard prediction.)r   predict)r   Xs     r   r   zSelfLearningWrapper.predict.   s    &&q))r   c                 J    | j                   j                  |      \  }}}|||fS )a&  
        Prediction with stability-based confidence.
        
        Returns:
            predictions: Array of predicted labels
            confidence: Array [batch_size, 10] - stability-weighted probabilities
            stability: Array [batch_size] - how stable each prediction is
        )r   measure_stability)r   r   r
   predictionsavg_confidences        r   predict_with_confidencez+SelfLearningWrapper.predict_with_confidence2   s/     9=8W8WXY8Z5+~N,<<<r   c                    | j                   j                  ||       | j                   j                  |      }| j                   j                  ||      }| j                  |dddd}| j                  | j
                  z  dk(  r,| j                  j                  d| j                        }||d<   | j                  | j                  z  dk(  r!| j                  j                  d      }||d<   | j                  d	z  dk(  r`t        j                  j                  dd
      }| j                  j                  |      \  }	}
}
t!        t        j"                  |	            |d<   | j$                  d   j'                  |       | j$                  d   j'                  |d          | j$                  d   j'                  |d          | xj                  dz  c_        |S )z
        Single training step with optional self-learning.
        
        Args:
            X: Training batch
            y_onehot: One-hot encoded labels
            
        Returns:
            dict with losses and stats
        Nr   )r   r   r	   	stabilityn_stabled   )	n_samplesr   r"   )
batch_sizer	   2   i  r!   r   r      )r   updateforwardcompute_lossr   r   r   discover_from_noiser   r   
self_trainnprandomrandnr   r   floatmeanr   append)r   r   y_onehoty_predsup_lossresultsr"   	self_lossX_test_noiser!   _s              r   
train_stepzSelfLearningWrapper.train_step>   s    	q(+((+??//&A '"&
 >>D333q8||77$($<$< 8 H #+GJ>>D4449//3/?I,5G() >>B!#99??34L"mm==lKOIq!#();#<GK &'..x8)*11':N2OP%&--gj.AB!r   c                    t        |      D ]  }t        j                  j                  dt	        |      |      }||   }||   }	| j                  |t        j                  d      |	         }
|sb|dz  dk(  sk| j                  j                         }t        d|dd|
d   dd	|
d
   r|
d
   nddd|
d   dd|d   d
        y)aV  
        Full training loop combining supervised and self-learning.
        
        Args:
            X_train: Training data [n_samples, 784]
            y_train: Training labels [n_samples]
            n_iterations: Number of iterations
            batch_size: Batch size for supervised learning
            verbose: Print progress
        r   
   r#   zIter 6dz | Sup Loss: r   z.4fz | Self Loss: r	   zN/Az>8z | Stable: r"   3dz | Buffer: buffer_usagez.1%N)
ranger-   r.   randintlenr:   eyer   get_truth_statisticsprint)r   X_trainy_trainn_iterationsr%   verboseiidxX_batchy_batchr6   statss               r   
train_loopzSelfLearningWrapper.train_loopo   s     |$ 	>A))##As7|Z@CclGclG oogrvvbz'/BCG1s7a<99;aV $##*+<#=c"B C$ELMaEbG,@$Ahmnp#q r!!(!4R 8 9!!&~!6s ;	= >	>r   c                     | j                   j                         }|| j                  | j                  d   r| j                  d   d   nddd}|S )z7Get comprehensive report on discovered emergent truths.r   N)
iterationsfinal_supervised_loss)emergent_truthstraining_history)r   rD   r   r   )r   rN   reports      r   get_truth_reportz$SelfLearningWrapper.get_truth_report   sV    113  %"nnPTP\P\]nPo6G)H)Luy!
 r   N)r<      g333333?)i  r#   T)
__name__
__module____qualname____doc__r   r   r   r:   rO   rW    r   r   r   r      s9    
 %'%&%)
6*
=/b 9=+/>:r   r   )numpyr-   emergent_truthr   stability_analyzerr   r   r]   r   r   <module>ra      s     / 0S Sr   