
    i"                        d Z ddlZddlmZ ddlmc mZ ddlmZ ddl	m
Z
 ddlmZmZ ddlZddlmZ  ej&                  d       ej(                  j+                  d        ed        ed        ed        G d d	ej.                        Z G d
 dej.                        Z G d dej.                        ZddZ ej8                         d        ZddZedk(  r eddd      Z yy)ze
Resonant Filter MLP - MNIST
Each layer applies a resonant filter with learnable resonance parameter
    N)
DataLoader)datasets
transforms*   <============================================================zRESONANT FILTER MLP - MNISTc                   .     e Zd ZdZ fdZd Zd Z xZS )ResonantFiltera  
    A simple resonant filter layer.
    Output = activation(W @ x + r * y_prev)
    where r is a learnable resonance coefficient (feedback).
    
    This creates frequency-selective behavior - resonance amplifies
    certain signal components based on the feedback strength.
    c                    t         |           t        j                  t	        j
                  ||      dz        | _        t        j                  t	        j                  |            | _        t        j                  t	        j                  d            | _
        | j                  dt	        j                  d|      d       y )Ng?g?state   F)
persistent)super__init__nn	Parametertorchrandnweightzerosbiastensor	resonanceregister_buffer)selfin_featuresout_features	__class__s      ex03.pyr   zResonantFilter.__init__!   s    ll5;;|[#ID#PQLL\!:;	 ell3&78 	Wekk!\&BuU    c                     t        j                  || j                  | j                        }|| j                  | j
                  z  z   }|j                         j                  dd      | _        |S )Nr   T)dimkeepdim)Flinearr   r   r   r   detachmean)r   xouts      r   forwardzResonantFilter.forward,   sY    hhq$++tyy1 DNNTZZ// ZZ\&&1d&;

r   c                 8    | j                   j                          y N)r   zero_)r   s    r   resetzResonantFilter.reset9   s    

r   )__name__
__module____qualname____doc__r   r)   r-   __classcell__r   s   @r   r	   r	      s    	Vr   r	   c                   :     e Zd ZdZdddgdf fd	Zd Zd Z xZS )	ResonantFilterMLPzMLP with resonant filter layers        
   c                     t         |           |g|z   |gz   }t        j                  t	        t        |      dz
        D cg c]  }t        ||   ||dz             c}      | _        y c c}w Nr   )r   r   r   
ModuleListrangelenr	   filtersr   
input_sizehidden_sizesoutput_sizesizesir   s         r   r   zResonantFilterMLP.__init__@   sl    |+{m;}}3u:>*&
 58U1Q3Z0&
  &
s   A+c                     |j                  |j                  d      d      }t        | j                        D ]>  \  }} ||      }|t	        | j                        dz
  k  s*t        j                  |      }@ |S Nr   r   )viewsize	enumerater?   r>   r#   relu)r   r'   rE   filters       r   r)   zResonantFilterMLP.forwardI   si    FF166!9b!"4<<0 	IAvq	A3t||$q((FF1I	
 r   c                 F    | j                   D ]  }|j                           y r+   )r?   r-   )r   rM   s     r   reset_stateszResonantFilterMLP.reset_statesS   s    ll 	FLLN	r   )r.   r/   r0   r1   r   r)   rO   r2   r3   s   @r   r5   r5   =   s     )"%S#JB r   r5   c                   4     e Zd ZdZdddgdf fd	Zd Z xZS )StandardMLPz(Baseline: standard MLP without resonancer6   r7   r8   r9   c                     t         |           |g|z   |gz   }t        j                  t	        t        |      dz
        D cg c]!  }t        j                  ||   ||dz            # c}      | _        y c c}w r;   )r   r   r   r<   r=   r>   Linearlayersr@   s         r   r   zStandardMLP.__init__[   sp    |+{m;mm3u:>*%
 IIeAhac
+%
  %
s   &A5c                     |j                  |j                  d      d      }t        | j                        D ]>  \  }} ||      }|t	        | j                        dz
  k  s*t        j                  |      }@ |S rG   )rI   rJ   rK   rT   r>   r#   rL   )r   r'   rE   layers       r   r)   zStandardMLP.forwardc   si    FF166!9b!!$++. 	HAuaA3t{{#a''FF1I	 r   )r.   r/   r0   r1   r   r)   r2   r3   s   @r   rQ   rQ   X   s    2"%S#JB r   rQ   c                    | j                          d}d}d}|r| j                          |D ]  \  }	}
|	j                  |      |
j                  |      }
}	|j                           | |	      } |||
      }|j	                          |j                          ||j                         |	j                  d      z  z  }|j                  d      }|||
k(  j                         j                         z  }||	j                  d      z  } ||z  ||z  fS )Nr   r   r!   )
trainrO   to	zero_gradbackwardstepitemrJ   argmaxsum)modelloader	criterion	optimizerdevicerO   
total_losscorrecttotalr'   youtputlosspreds                 r   train_epochrm   l   s    	KKMJGE 1ttF|QTT&\1q#diikAFF1I--
}}}#DAI??$))++ w..r   c                 h   | j                          d}d}t        | d      r| j                          |D ]y  \  }}|j                  |      |j                  |      }} | |      }|j	                  d      }|||k(  j                         j                         z  }||j                  d      z  }{ ||z  S )Nr   rO   r   rX   )evalhasattrrO   rZ   r_   r`   r^   rJ   )	ra   rb   re   rg   rh   r'   ri   rj   rl   s	            r   evaluaterq      s    	JJLGE un% 1ttF|QTT&\1q}}}#DAI??$))++ U?r   r9   r7   c           
         t        j                  t         j                  j                         rdnd      }t	        d|        t        j                  t        j                         t        j                  dd      g      }t        t        j                  ddd|      |d	      }t        t        j                  dd
|      |      }t        j                         }t	        d       t	        d       t               j                  |      }t!        j"                  |j%                         |      }	g }
g }t'        |       D ]P  }t)        ||||	|      \  }}|
j+                  |       |j+                  |       t	        d|dz   dd|dd|d       R t-        |||      }t	        d|d       t	        d       t/        |j0                        D ]/  \  }}t	        d| d|j2                  j5                         d       1 t	        d       t	        d       t7               j                  |      }t!        j"                  |j%                         |      }g }g }t'        |       D ]R  }t)        |||||d
      \  }}|j+                  |       |j+                  |       t	        d|dz   dd|dd|d       T t-        |||      }t	        d|d       t9        j:                  ddd      \  }}|d    j=                  |
d!d"d#$       |d    j=                  |d%d&d#$       |d    j?                  d'       |d    jA                  d(       |d    jC                  d)       |d    jE                          |d    jG                  dd*+       |d   j=                  |d!d"d#$       |d   j=                  |d%d&d#$       |d   j?                  d'       |d   jA                  d,       |d   jC                  d-       |d   jE                          |d   jG                  dd*+       |d#   jI                  d.d&g||gd/d0gd1d23       |d#   jA                  d4       |d#   jC                  d5       |d#   jK                  d6d7g       |d#   jG                  dd*d89       t'        tM        |j0                              D cg c]  }d:| 	 }}|j0                  D cg c]  }|j2                  j5                          }}|d;   jI                  ||d<d1d23       |d;   jO                  d d2d=d>?       |d;   jA                  d@       |d;   jC                  dA       |d;   jG                  dd*d89       t9        jP                          t9        jR                  dBdCD       t9        jT                          t	        dE       t	        dF       t	        dG       t	        dH|d       t	        dI|d       ||z
  dJz  }t	        dK|dLdM       |
|dN||dNdOS c c}w c c}w )PNcudacpuz	
Device: )g_)Ǻ?)gGr?dataT)rY   download	transform)
batch_sizeshuffleF)rY   rw   )rx   u"   
🧠 Training Resonant Filter MLPz2--------------------------------------------------)lrzEpoch r   2dz: Loss=z.4fz, Train Acc=u$   📊 Resonant Filter Test Accuracy: u%   
🔮 Learned Resonance Coefficients:z  Layer z: resonance = u%   
⚡ Training Standard MLP (baseline))rO   u   📊 Standard Test Accuracy:    )   r|   )figsizer   zb-ozResonant Filter   )label	linewidthzr--sStandardEpochLosszTraining Lossg333333?)alphaAccuracyzTraining AccuracyzResonant
Filterblueredgffffff?black)colorr   	edgecolorzTest AccuracyzTest Accuracy Comparisong?g      ?ri   )r   axisL   purple-g      ?)ri   r   	linestyler   zResonance ValuezLearned Resonance Coefficientszresonant_filter_mlp.png   )dpiz=
============================================================SUMMARYr   z#Resonant Filter MLP Test Accuracy: z#Standard MLP Test Accuracy:        d   zDifference: z+.2f%)
train_losstest_acc)resonantstandard)+r   re   rs   is_availableprintr   ComposeToTensor	Normalizer   r   MNISTr   CrossEntropyLossr5   rZ   optimAdam
parametersr=   rm   appendrq   rK   r?   r   r^   rQ   pltsubplotsplot
set_xlabel
set_ylabel	set_titlelegendgridbarset_ylimr>   axhlinetight_layoutsavefigshow)n_epochsrx   rz   re   rw   train_loadertest_loaderrc   r   opt_resresonant_train_lossesresonant_train_accsepochrk   accresonant_test_accrE   fr   opt_stdstandard_lossesstandard_accsstandard_test_accfigaxeslayer_names
res_valuesdiffs                               r   run_experimentr      s   \\EJJ$;$;$=&5IF	Jvh
  ""Y	2$ I
 vTDINtL vUi@K
 ##%I 

/0	(O "%%f-Hjj,,.26Gx K,	7FS	c$$T*""3'uQwrl'$s<CyIJK !;?	01B30G
HI 

23(**+ D1>!++*:*:*<S)ABCD 

23	(O}'Hjj,,.26GOMx K,	7Fafg	ct$S!uQwrl'$s<CyIJK !;?	)*;C)@
AB Q73IC 	GLL&5FRSLTGLL&
aLHGwGvGo&GNNGLLSL! 	GLL$e3DPQLRGLLjALFGwGz"G)*GNNGLLSL! 	GKK#Z0"$56uoSG  E 	G'G01Gc3Z GLLSsL+ %*#h.>.>*?$@AqQqc7AKA.6.>.>?!++""$?J?GKKZxsgKVGOOaw#OEG()G67GLLSsL+KK)s3HHJ 
/	)	(O	/0A#/F
GH	/0A#/F
GH 11S8D	Ld1
%& $9FWX#2@QR + B?s   .W>
!X__main__gMbP?)r   rx   rz   )T)r9   r7   g{Gz?)!r1   r   torch.nnr   torch.nn.functional
functionalr#   torch.optimr   torch.utils.datar   torchvisionr   r   numpynpmatplotlib.pyplotpyplotr   manual_seedrandomseedr   Moduler	   r5   rQ   rm   no_gradrq   r   r.   results r   r   <module>r      s   
      ' ,     "  		r  h # $ h#RYY #L		 6")) (/4  &yx zbSTBG r   