
    iq                        d dl Zd dlmZ d dlmZmZmZ d dlm	Z	 d dl
Z
d dlZ e
j                  e
j                  d        e
j                  e      Ze	 G d d             Ze	 G d	 d
             Z G d d      Z G d d      Z G d d      Z G d d      Zedk(  rej/                  d        ed      d   j1                  dd      Z ed      d   dz  j5                  e      Z ed      d   j1                  dd      Z ed      d   dz  j5                  e      Zej1                  dddd      Zej1                  dddd      Z ej/                  d        edd d!d"#      Z!e!jE                  eed$d%d!&       e!jG                  e e      Z$ej/                  d'e$d(       ej/                  d)       e dd* Z%e!jL                  jO                  e%      Z(ej/                  d+e%jR                          ej/                  d,e(jR                          ej/                  d-e(jU                         d.d/e(jW                         d.d0       yy)1    N)read)TupleOptionalCallable)	dataclassz)%(asctime)s - %(levelname)s - %(message)s)levelformatc                       e Zd ZU dZeeeef   ed<   dZeed<   dZeeef   ed<   dZ	eed<   d	Z
eed
<   dZeed<   dZeed<   dZeed<   y)TransformConfigz6Configuration for the Adversarial Reference Transform.input_shape@   reference_channels   r   reference_grid_size   contemplation_depthMbP?learning_rateg?momentumg-C6?regularization   progress_log_intervalN)__name__
__module____qualname____doc__r   int__annotations__r   r   r   r   floatr   r   r        ref01.pyr   r      sh    @sC}%%  +1sCx1   M5 He NE !$3$r"   r   c                   l    e Zd ZU dZej
                  ed<   ej
                  ed<   ej
                  ed<   y)ReferenceBankz6Stores learned reference points and their activations.	positionsfeaturesattention_weightsN)r   r   r   r   npndarrayr   r!   r"   r#   r%   r%      s#    @zzjjzz!r"   r%   c                      e Zd ZdZ	 	 	 ddeeeef   dedeeef   defdZd Zdej                  d	ej                  d
ej                  fdZ
dej                  d
ej                  fdZ	 	 d dej                  dej                  dedee   dee   d
ej                  fdZdej                  d
ej                  fdZdej                  dej                  fdZdej                  dej                  fdZ	 	 	 d!dej                  dedee   dee   d
ej                  f
dZd"dej                  dedefdZdej                  d
ej                  fdZy)#AdversarialReferenceTransformaH  
    Non-convolutional reference transform inspired by Conv2D.
    
    Instead of sliding kernels, this transform:
    1. Maintains a bank of reference points (learned positions)
    2. Uses adversarial contemplation to select optimal references
    3. Transforms input through reference-based interpolation
    
    Key Insight: Conv2D's power comes from local receptive fields.
    We achieve similar behavior by:
    - Learning what positions matter (generator)
    - Judging if transform is "conv-like" (discriminator)
    - Iterative refinement through adversarial training
    r   r   r   r   c                 L    t        ||||      | _        | j                          y )Nr   r   r   r   )r   config_initialize_components)selfr   r   r   r   s        r#   __init__z&AdversarialReferenceTransform.__init__6   s*     &#1 3 3	
 	##%r"   c           	         | j                   j                  \  }}}t        j                  t        j                  dd| j                   j
                  d         t        j                  dd| j                   j
                  d         d      \  }}t        j                  |j                         |j                         gd      }| j                   j                  t        |      z
  }|dkD  r7t        j                  j                  |d      }t        j                  ||g      }t        |j                  t        j                        t        j                  j!                  t        |      | j                   j                        dz  t        j"                  t        |      t        j                        	      | _        t'        |d
z  d|t        j(                  g d            | _        t-        |dz        | _        g | _        y)z$Initialize all learnable components.r      ij)indexingaxis   {Gz?dtype)r&   r'   r(   r      r:   r:   r:   r:   
input_sizehidden_sizeoutput_sizer   )input_channelsN)r/   r   r)   meshgridlinspacer   stackravelr   lenrandomrandvstackr%   astypefloat32randnonesreference_bankMLPClassifierarraytransform_mlpDiscriminatordiscriminatorcontemplation_history)	r1   HWCgrid_ygrid_xinitial_positionsn_extraextra_positionss	            r#   r0   z4AdversarialReferenceTransform._initialize_componentsF   s   ++))1a KK1dkk==a@AKK1dkk==a@A

 HHfllnflln%EAN ++0037H3IIQ; iinnWa8O "		+<o*N O+'..rzz:YY__S):%;T[[=[=[\_cc ggc*;&<BJJO
 +1u((#;<	
 +!a%@ &("r"   Xr&   returnc                 B   |j                   \  }}}}|j                   d   }|dddf   |dz
  z  }|dddf   |dz
  z  }	t        j                  |      j                  t              }
t        j                  |	      j                  t              }t        j
                  |
dz   d|dz
        }t        j
                  |dz   d|dz
        }||
z
  }|	|z
  }t        j                  |||ft        j                        }t        |      D ]  }t        |      D ]  }||||   |
|   f   }||||   |
|   f   }||||   ||   f   }||||   ||   f   }|d||   z
  z  d||   z
  z  |d||   z
  z  ||   z  z   |||   z  d||   z
  z  z   |||   z  ||   z  z   |||f<     |S )a  
        Extract features at reference positions via bilinear interpolation.
        
        This replaces conv2d's local receptive field with reference-based sampling.
        
        Args:
            X: Input tensor (B, H, W, C)
            positions: Reference positions (N_ref, 2) in [0, 1]
            
        Returns:
            Features at references (B, N_ref, C)
        r   Nr4   r;   )	shaper)   floorrL   r   clipzerosrM   range)r1   r_   r&   BrW   rX   rY   N_refpxpyx0y0x1y1wxwyr'   bnv00v01v10v11s                          r#   _compute_reference_featuresz9AdversarialReferenceTransform._compute_reference_featuresn   s    WW
1a" q!t_A&q!t_A& XXb\  %XXb\  %WWR!VQA&WWR!VQA&"W"W88QqM<q 	A5\ 2a5"Q%(2a5"Q%(2a5"Q%(2a5"Q%( 1r!u9%RU31r!u9%1-."Q%K1r!u9-. "Q%K"Q%'( A	 r"   c                 "   |j                   \  }}}}t        j                  |      }|ddddddf   |ddddddf   z
  |ddddddf<   |dddddf   |dddddf   z
  |dddddf<   |dddddf   |dddddf   z
  |dddddf<   t        j                  |      }|ddddf   |ddddf   z
  |ddddddf<   |dddf   |dddf   z
  |dddddf<   |dddf   |dddf   z
  |dddddf<   t        j                  |      }|ddddddf   |ddddddf   z   |ddddddf   z   |ddddddf   z   d|ddddddf   z  z
  |ddddddf<   t        j                  ||||gd      S )	z
        Compute spatial context features (gradient, laplacian approximations).
        
        Mimics Conv2D's edge detection capability without convolutions.
        Nr9   r4   r      r7   )rb   r)   
zeros_likeconcatenate)	r1   r_   rg   rW   rX   rY   grad_xgrad_y	laplacians	            r#   _compute_spatial_contextz6AdversarialReferenceTransform._compute_spatial_context   s    WW
1a q!q!QRx[1Q3B3Y<7q!QrTzAq!G*qAqz1q!QwQ2X;1a84q!Rx q!q!"uX!SbS&	1q!B$zAqD'AadG+q!QwQU8a2h.q"ax MM!$	aQrTkNQq#2#qt|_,a2qrkNq!B$|_-!QrT1R4-  ! 	!QrT1R4-  ~~q&&)<2FFr"   Nreference_featuresepochbatch_indextotal_batchesc                    |j                   \  }}}}	| j                  |      }
t        | j                  j                        D ]  }t        j                         }|j                  |||z  |	      }t        |j                   d   d      }t        j                  j                  |j                   d   |d      }|dd|ddf   }t        j                  |||z  |	ft        j                        }t        ||z        D ]  }||z  }||z  }t        j                  ||z  ||z  ggt        j                        }t        j                  j!                  | j"                  j$                  ddddf   |z
  d      }t        j&                  |      d| }|dd||dz   ddf   }|dd|ddf   j)                  dd	
      }|
dd|||dz   d|	f   }t        j*                  |||gd      }t        |      D ];  }||df   }| j,                  j/                  |j                  dd            }||||f<   = |dz   | j                  j0                  z  dk(  s|dz   ||z  k(  sdt        j                         |z
  }|dz   } | t3        |d      z  }!|||dz    d| }"nd}"t4        j7                  d|dz   |"|dz   | j                  j                  | ||z  |!        | j9                  |      }#| j:                  j=                  |#|d      \  }$}%|%dkD  r| j?                  ||#       | jA                  ||       t4        jC                  d| d|$dd|%d        j                  ||||	      S )a  
        Adversarial contemplation: Generator refines references, Discriminator judges.
        
        This iterative process forces the reference transform to capture local
        receptive field behavior without explicit convolutions.
        
        Args:
            X: Input (B, H, W, C)
            reference_features: Features at current references (B, N_ref, C)
            epoch: Current training epoch
            
        Returns:
            Refined reference features (B, H, W, C)
        r4      F)replaceNr;   r9   r7   Tr8   keepdimsrz   r   &.>/?zFART epoch %d batch %s contemplation %d/%d: %d/%d positions, %.1f pos/sr   )realfaker   gffffff?zContemplation z	: D_loss=.4fz, D_acc=.3f)"rb   r   rf   r/   r   timeperf_counterreshapeminr)   rI   choicere   rM   rR   linalgnormrP   r&   argsortmeanr}   rS   forwardr   maxloggerinfo_mimic_conv2drU   
train_step_improve_generator_update_reference_positionsdebug)&r1   r_   r   r   r   r   rg   rW   rX   rY   spatial_contextcontemplation_step
step_startX_flatn_refref_idxsampled_refsoutputpos_idxpos_ypos_xpos_norm	distances
local_refspixelref_featurescontextcombinedrq   
combined_btransformedelapsedpositions_donepositions_per_secondbatch_labelconv_like_featuresd_loss
d_accuracys&                                         r#   _adversarial_contemplationz8AdversarialReferenceTransform._adversarial_contemplation   s   , WW
1a77:"'(G(G"H S	**,J YYq!a%+F *003R8Eii&&'9'?'?'BESX&YG-a!m<L XXq!a%m2::>F Q< .1! 88eai%;$<BJJOIINN''11!RaR%88C + 	  ZZ	26E:
 q''!)"3Q671!Z2BCHHaZ^H_)!UE%'M2A2*EF >>5,*HrR q 5A!)!Q$J"&"4"4"<"<Z=O=OPQSU=V"WK)4F1g:&5 q[DKK$E$EEJ{a!e+"//1J>G%,q[N+9C<N+N(".=3L)4q(9=/&J&)KK`	#*Q.77&A,	K.d "&!3!3A!6 "&!3!3!>!>'# "? "FJ C''0BC ,,Q7LL !3 4IfS\R\]`QabcS	j ~~aAq))r"   c           	      L   |j                   \  }}}}d}|dz  }t        j                  |d||f||fdfd      }t        j                  ||||ft        j                        }	t        |      D ],  }
t        |      D ]  }|	|dd|
|
|z   |||z   ddf   z  }	 . |	||z  z  }	|	S )z
        Create a pseudo-conv2d output for discriminator comparison.
        
        Uses reference-based aggregation to mimic convolution behavior.
        r   r9   )r   r   reflect)moder;   N)rb   r)   padre   rM   rf   )r1   r_   rg   rW   rX   rY   kernel_sizer   X_paddedr   dydxs               r#   r   z+AdversarialReferenceTransform._mimic_conv2d,  s     WW
1a Q 66!fsCj3*fEIV 1aA,bjj9$ 	;BK( ;(1bAgr"Q$w#9::;	; 	+++r"   r   c                 8   |j                   \  }}}}|j                  dk(  r|j                  ||||      }n'|j                  dk7  rt        d|j                          t	        j
                  t        | j                  j                              }t        | j                  j                        D ]e  \  }}	t        |	d   |dz
  z        t        |	d   |dz
  z        }}
t	        j                  t	        j                  |dd||
ddf               }|||<   g ||j                         z
  }|dz  }| j                  xj                  |ddt        j                  f   t	        j                  ddgg      z  z  c_        t	        j                  | j                  j                  dd      | j                  _        | j                  xj                   d|dz  z   z  c_        | j                  xj                   | j                  j                   j#                         z  c_        y)	z
        Update reference positions based on reconstruction quality.
        
        References that help reconstruct the transformed output move closer to
        important positions; others move away (adversarial positioning).
        r   r{   z:Expected transformed to have 3 or 4 dimensions, got shape r   r4   Nr:   皙?)rb   ndimr   
ValueErrorr)   re   rH   rP   r&   	enumerater   r   absnewaxisrR   rd   r(   sum)r1   r_   r   rg   rW   rX   rY   errorsrr   posri   rj   contributiongradientposition_updatess                  r#   r   z9AdversarialReferenceTransform._update_reference_positionsD  s    WW
1aq %--aAq9K"L[M^M^L_` 
 #d11;;<= 3 3 = => 	%FAsQ1q5)*CA!a%0@,AB 77266+aRl*C#DEL$F1I	% FKKM)#d?%%)9!RZZ-)H288VWYZU[T\K])]]%(*0C0C0M0MqRS(T% 	--!fsl2BC---1D1D1V1V1Z1Z1\\-r"   r   r   c           	      T   ||z
  }t        j                  |t        t        |j                  dz
              d      }t        j
                  |t         j                        j                  dd      }t        j                  |dd      }||j                  dd      z  }t        d	      D ]|  }t         j                  j                  d| j                  j                  j                  d
         j                  t         j                        }| j                  j!                  ||       ~ y)zE
        Improve generator when discriminator is too strong.
        r4   Fr   r;   rz   gư>NTr   r   )r)   r   tuplerf   r   asarrayrM   r   rd   r   rI   rJ   rS   W1rb   rL   update)r1   r   r   diff	mean_difftarget_X_samples           r#   r   z0AdversarialReferenceTransform._improve_generatorj  s    
 d{GGDuU499q=-A'BUS	IRZZ8@@BGt,&**!d*33 q 	8Ayy~~a););)>)>)D)DQ)GHOOPRPZPZ[H%%h7	8r"   c                    |j                   \  }}}}| j                  || j                  j                        }	| j	                  ||	|||      }
| j                  j
                  }|j                  dddt        |            j                  |d      ddddddd|
j                   d   f   }|
t        j                  |      z  }|S )z
        Forward pass through the adversarial reference transform.
        
        Args:
            X: Input tensor (B, H, W, C)
            
        Returns:
            Transformed output (B, H, W, C)
        r   r   r   r4   rz   r7   N)rb   rw   rP   r&   r   r(   r   rH   repeatr)   r   )r1   r_   r   r   r   rg   rW   rX   rY   r   r   	attentionattention_reshapedr   s                 r#   r   z%AdversarialReferenceTransform.forwardz  s      WW
1a "==))
 55#' 6 
 ''99	&..q!S^

&&
Q1&<{'8'8'<&<<>
 rwwy11r"   X_trainepochs
batch_sizec                 H   t         j                  d| d       |j                  d   }t        ||z   dz
  |z  d      }t	        |      D ]V  }t
        j                  j                  |      }d}d}	t        j                         }
t	        d||      D ]  }||z  }||||z    }||   }t        j                         }| j                  ||||      }t        j                  ||z
  dz        }||z  }|	dz  }	t        j                         |z
  }t         j                  d|dz   ||dz   |||t        |      t        |d      z          |t        |	d      z  }t         j                  d	|dz   ||t        j                         |
z
         | j                  j                  |       Y y
)z
        Fit the transform using adversarial contemplation.
        
        Args:
            X_train: Training data (N, H, W, C)
            epochs: Number of training epochs
            batch_size: Batch size for contemplation
        z0Starting adversarial contemplation training for z epochsr   r4   r   r9   zHART epoch %d/%d batch %d/%d: loss=%.4f, batch_time=%.2fs, samples/s=%.2fr   z9ART epoch %d/%d complete: avg_loss=%.4f, epoch_time=%.2fsN)r   r   rb   r   rf   r)   rI   permutationr   r   r   r   rH   rV   append)r1   r   r   r   Nr   r   indices
epoch_lossbatch_countepoch_startbatch_startr   	batch_idxX_batchbatch_start_timer   
recon_lossbatch_elapsedavg_losss                       r#   fitz!AdversarialReferenceTransform.fit  s    	FvhgVWMM!Q^a/J>B6] .	8Eii++A.GJK++-K$Q:6 )Z7#Kj0HI	!),#'#4#4#6   +"/	 &   WWfw&61%<=
j(
q  $ 1 1 36F F^AI!O!!L3}d#;;	'< "CQ$77HKKK	!!#k1 &&--h7].	8r"   c                 $    | j                  |      S )z,Transform new data using learned references.)r   )r1   r_   s     r#   	transformz'AdversarialReferenceTransform.transform  s    ||Ar"   )r   r   r   )NN)r   NN)
       )r   r   r   r   r   r   r2   r0   r)   r*   rw   r   r   r   r   r   r   r   r   r   r!   r"   r#   r,   r,   &   s   $ #%/5#$&3S=)&  & #38_	&
 !& &(P.RZZ .BJJ .SUS]S] .`G"** G GF &*'+n*::n* JJn* 	n*
 c]n*  }n* 
n*`rzz bjj 0$]RZZ $]bjj $]L8rzz 8 8& %)'+*::* * c]	*
  }* 
*X<82:: <8s <8S <8|2:: "** r"   r,   c                       e Zd ZdZddedefdZdej                  dej                  dee	e	f   fdZ
	 dd	ej                  d
ej                  de	dee	e	f   fdZy)rT   z
    Discriminator network that judges if a transform is "conv-like".
    
    The discriminator compares reference-based transforms against
    ground-truth convolution outputs, providing adversarial feedback.
    rC   rA   c                 2   t         j                  j                  |dz  |      dz  | _        t        j                  d|f      | _        t         j                  j                  |d      dz  | _        t        j                  d      | _        d| _        d| _	        y )Nr9   r:   r4   )r4   r4         ?)
r)   rI   rN   r   re   b1W2b2real_probabilityfake_probability)r1   rC   rA   s      r#   r2   zDiscriminator.__init__  st    ))//.1"4kBTI((A{+,))//+q1D8((6" !$ #r"   real_featuresfake_featuresr`   c                 r   |j                  |j                  d   d      }|j                  |j                  d   d      }t        j                  ||gd      }|j                  d   }t        j                  |ddd|f         }t        j                  |dd|df         }t        |      t        |      fS )z
        Classify features as real (conv-like) or fake (reference-based).
        
        Returns:
            Tuple of (discriminator score for real, score for fake)
        r   rz   r4   r7   N)r   rb   r)   r}   r   r    )	r1   r  r  	real_flat	fake_flatcombined_real
real_width
score_real
score_fakes	            r#   r   zDiscriminator.forward  s     "))-*=*=a*@"E	!))-*=*=a*@"E	 (
  __Q'
WW]1kzk>:;
WW]1jk>:;
Z %
"333r"   r   r   r   c           	      j   | j                  ||      \  }}d}||z
  }| xj                  ||z  dz  z  c_        | xj                  ||z  dz  z  c_        t        dt	        dd|dz  z               | _        t        dt	        dd|dz  z               | _        t        |      }ddt        |      z   z  }||fS )z}
        Single training step for the discriminator.
        
        Returns:
            Tuple of (loss, accuracy)
        r   r:   r   gGz?r         ?)r   r   r   r   r   r  r  r   )	r1   r   r   r   r
  r  accuracyerrorlosss	            r#   r   zDiscriminator.train_step  s     "&dD!9
J  Z'=5(4//=5(500 !$D#dC*s:J4J*K L #D#dC*s:J4J*K L5z#E
*+X~r"   N)r   )r   )r   r   r   r   r   r2   r)   r*   r   r    r   r   r!   r"   r#   rT   rT     s    $s $ $4RZZ 4

 4uUZ\aUaOb 46  %	jj jj 	
 
ue|	r"   rT   c            	          e Zd ZdZ	 ddedededej                  fdZdej                  d	ej                  fd
Zdej                  d	ej                  fdZ	dej                  d	ej                  fdZ
dej                  d	ej                  fdZdej                  dej                  dej                  fdZdej                  dej                  fdZy)rQ   z<MLP Classifier with SPDER activation (from previous update).Nr@   rA   rB   r   c                 `   t         j                  j                  ||      dz  | _        t        j                  d|f      | _        t         j                  j                  ||      dz  | _        t        j                  d|f      | _        ||| _	        y t        j                  g d      | _	        y )Nr:   r4   r>   )
r)   rI   rN   r   re   r   r   r   rR   lr)r1   r@   rA   rB   r   s        r#   r2   zMLPClassifier.__init__>  s~    ))//*k:TA((A{+,))//+{;dB((A{+,#0#<-"((KcBdr"   xr`   c                     t        j                  |      t        j                  t        j                  |      dz         z  S )z2SPDER: sin(x) * sqrt(|x|) - Periodic with damping.:0yE>)r)   sinsqrtr   )r1   r  s     r#   spderzMLPClassifier.spderF  s+    vvay277266!9t#3444r"   c                 (   t        j                  |      dz   }t        j                  |      }t        j                  |      }t        j                  |dk(  d|      }|t        j
                  |      z  |d|z  z  t        j                  |      z  z   S )zDerivative of SPDER.r  r   r4   r9   )r)   r   r  signwherecosr  )r1   r  abs_x
sqrt_abs_xsign_xs        r#   spder_derivativezMLPClassifier.spder_derivativeJ  sr    q	D WWU^
&A+q&1BFF1I%1z>)BbffQi(OOOr"   c                     t        j                  |t        j                  |dd      z
        }|t        j                  |dd      z  S )Nr4   Tr   )r)   expr   r   )r1   r  exp_xs      r#   softmaxzMLPClassifier.softmaxR  s:    q266!!d;;<rvve!d;;;r"   r_   c                 d   |dz  dz
  }t        j                  || j                        | j                  z   | _        | j                  | j                        | _        t        j                  | j                  | j                        | j                  z   | _	        | j                  | j                        S )N     _@r  )r)   dotr   r   z1r  a1r   r   z2r%  )r1   r_   X_norms      r#   r   zMLPClassifier.forwardV  sx    US&&)DGG3**TWW%&&$''*TWW4||DGG$$r"   y_truey_predc                    |j                   d   }|dz  dz
  }||z
  }t        j                  | j                  j                  |      |z  }t        j
                  |dd      |z  }t        j                  || j                  j                        }	|	| j                  | j                        z  }
t        j                  |j                  |
      |z  }t        j
                  |
dd      |z  }||||fS )Nr   r'  r  Tr   )	rb   r)   r(  r*  Tr   r   r!  r)  )r1   r_   r-  r.  mr,  dz2dW2db2da1dz1dW1db1s                r#   backwardzMLPClassifier.backward]  s    LLOUSvoffTWWYY$q(ffSq4014ffS$''))$D))$''22ffVXXs#a'ffSq4014Cc!!r"   c                 ~   | j                  |      }| j                  |||      \  }}}}| xj                  | j                  d   |z  z  c_        | xj                  | j                  d   |z  z  c_        | xj
                  | j                  d   |z  z  c_        | xj                  | j                  d   |z  z  c_        y )Nr   r4   r9   r   )r   r9  r   r  r   r   r   )r1   r_   r-  r.  r7  r8  r3  r4  s           r#   r   zMLPClassifier.updatel  s    a!]]1ff=S#s4771:##4771:##4771:##4771:##r"   )N)r   r   r   r   r   r)   r*   r2   r  r!  r%  r   r9  r   r!   r"   r#   rQ   rQ   ;  s    F .2e3 eS es e "

e5rzz 5bjj 5P"** P P< <

 <% %

 %""** "bjj ""** "$

 $BJJ $r"   rQ   c                   2   e Zd ZdZ	 	 	 	 ddeeeef   dededefdZdej                  dej                  fd	Z		 dd
ej                  dej                  dededef
dZ
dej                  dej                  fdZdej                  dej                  defdZy)ART_MLP_Classifierz
    Combined Adversarial Reference Transform + MLP Classifier.
    
    The ART layer preprocesses the input, then feeds to MLP for classification.
    r   r   rA   rB   c           	          t        ||dd      | _        t        t        j                  |      ||t        j
                  g d            | _        y )Nr   r   r.   r>   r?   )r,   artrQ   r)   prodrR   mlp)r1   r   r   rA   rB   s        r#   r2   zART_MLP_Classifier.__init__  sI     1#1 & !	
 !ww{+##((#;<	
r"   r_   r`   c                 $   t        |j                        dk(  r?t        t        j                  |j                  d               }|j                  d||d      }| j                  j                  |      }|j                  |j                  d   d      S )z!Apply ART preprocessing to input.r9   r4   rz   r   )rH   rb   r   r)   r  r   r>  r   )r1   r_   sizer   s       r#   
preprocesszART_MLP_Classifier.preprocess  sw    qww<1rwwqwwqz*+D		"dD!,A hh&&q) "";#4#4Q#7<<r"   r   y_train
art_epochs
mlp_epochsr   c                    t         j                  d       | j                  j                  |||       t         j                  d       | j	                  |      }t        |      D ]  }t        j                  j                  dt        |      |      }||   }	||   }
| j                  j                  |	t        j                  d      |
          |dz  dk(  sr| j                  ||      }t         j                  d| d|d	        y
)zTrain the combined pipeline.z4Phase 1: Training Adversarial Reference Transform...)r   r   z#Phase 2: Training MLP Classifier...r   r   r4   z
MLP Epoch z: accuracy=r   N)r   r   r>  r   rC  rf   r)   rI   randintrH   r@  r   eyescore)r1   r   rD  rE  rF  r   X_preprocessediidxr   y_batchaccs               r#   r   zART_MLP_Classifier.fit  s     	JKWZJG9:1z" 	BA))##As7|Z@C$S)GclGHHOOGRVVBZ%891uzjj9j;s3i@A	Br"   c                     | j                  |      }t        j                  | j                  j	                  |      d      S )Nr4   r7   )rC  r)   argmaxr@  r   )r1   r_   X_procs      r#   predictzART_MLP_Classifier.predict  s0    #yy))&1::r"   r-  c                     |j                   dk(  r1t        j                  | j                  j	                  |      d      }n| j                  |      }t        j                  ||k(        S )Nr9   r4   r7   )r   r)   rQ  r@  r   rS  r   )r1   r_   r-  r.  s       r#   rJ  zART_MLP_Classifier.score  sM    66Q;YYtxx//2;F\\!_Fwwv'((r"   N)   rV  r4   r   d   r   )r   rW  rW  )r   r   r   r   r   r   r2   r)   r*   rC  r   rS  r    rJ  r!   r"   r#   r<  r<  y  s     -8"$
3S=)
  
 	

 
.=BJJ =2:: = KNB2:: B

 BB-0BDGB&; ;

 ;)rzz )2:: )% )r"   r<  __main__zLoading data...z../X_train.wavr4   rz   i  z../y_train.wav	   z../X_test.wavz../y_test.wavrV  z$Initializing ART + MLP Classifier...rU  r   rW  r   )r   r   rA   rB   r   r9   )rE  rF  r   zFinal Test Accuracy: r   z)Demonstrating standalone ART transform...r   zInput shape: zART output shape: zOutput range: [r   z, ]),numpyr)   scipy.io.wavfiler   typingr   r   r   dataclassesr   loggingr   basicConfigINFO	getLoggerr   r   r   r%   r,   rT   rQ   r<  r   r   r   rL   r   rD  X_testy_testX_train_img
X_test_img
classifierr   rJ  r  
demo_inputr>  r   
art_outputrb   r   r   r!   r"   r#   <module>rj     sb    ! , , !     ',,/Z [			8	$ 	% 	% 	% " " "@ @NG G\7$ 7$|G) G)\ z
KK!"#$Q'//C8G$%a(1,44S9G/"1%--b#6F?#A&*2237F //"b"a0KBA.J KK67#	J NN;A!PSNT 
F3H
KK'~67 KK;<BQJ''
3J
KK-
 0 0123
KK$Z%5%5$678
KK/*.."23!7r*..:J39OqQRE r"   