
    mil                        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
mZ ddlmZ ddlmZ ddlZddlmZmZmZmZ ddlZddlmZ ddlZddlmZ  G d d	ej:                        Z G d
 dej:                        Z G d dej:                        Z  G d dej:                        Z!dde"de#de!fdZ$ G d d      Z%	 	 	 dde&de"de"deeef   fdZ'd Z(e)dk(  r e(       Z*yy)z
Flux ResNet-18 for CIFAR-10
=============================
ResNet-18 architecture with Flux Algebra tracking.
Trains on CIFAR-10 with uncertainty quantification and entropy collapse.

Based on: The Algebra of Flux (Conditional Collapse Theory)
    N)datasets
transforms)
DataLoader)DictListTupleOptional)tqdm)
FluxTensorc                        e Zd ZdZddedef fdZd Z fdZde	j                  de	j                  fd	Zd
e	j                  dedefdZdefdZ xZS )FluxBatchNorm2dz_
    Batch Normalization with Flux tracking.
    Tracks uncertainty in running statistics.
    num_featuresmomentumc                 v    t         |           t        j                  ||      | _        d | _        d | _        y )N)r   )super__init__nnBatchNorm2dbn
gamma_flux	beta_flux)selfr   r   	__class__s      7/home/per/Documents/flux algebra/flux_resnet_cifar10.pyr   zFluxBatchNorm2d.__init__"   s0    ..A     c                    | j                   It        | j                  j                  j                  j                         t        j                  | j                  j                  j                        dz  t        j                  | j                  j                  j                              | _         t        | j                  j                  j                  j                         t        j                  | j                  j                  j                        dz  t        j                  | j                  j                  j                              | _
        yy)z5Initialize flux tensors after parameters are created.N皙?st)r   r   r   weightdataclonetorch	ones_like
zeros_likebiasr   r   s    r   initialize_fluxzFluxBatchNorm2d.initialize_flux*   s    ??"(##))+//$''.."5"56<""477>>#6#67DO
 (!!'')//$'',,"3"34s:""477<<#4#45DN #r   c                 N   t        |   |i | | j                  >| j                  j                  | j                  j                  j
                        | _        | j                  >| j                  j                  | j                  j                  j
                        | _        | S z4Override to() to move flux tensors to target device.)r   tor   r   r!   devicer   r   argskwargsr   s      r   r,   zFluxBatchNorm2d.to8   su    
D#F#??&"oo001F1FGDO>>%!^^..tww~~/D/DEDNr   xreturnc                    | j                          | j                  j                  j                  |j                  k7  rT| j                  j	                  |j                        | _        | j
                  j	                  |j                        | _        | j                  j                  | j                  j                  _        | j
                  j                  | j                  j                  _        | j                  |      S N)
r)   r   vr-   r,   r   r   r!   r"   r'   r   r1   s     r   forwardzFluxBatchNorm2d.forwardA   s    ??##qxx/"oo00:DO!^^..qxx8DN"oo// NN,,wwqzr   losslrworkc                 &   | j                          t        j                         5  | j                  j                  j
                  | | j                  j                  j
                  z  }t        | j                  j                  |z   | j                  j                  t        j                  |      z   | j                  j                  |z         j                  |      | _        | j                  j                  | j                  j                  _        | j                  j                  j
                  | | j                  j                  j
                  z  }t        | j                  j                  |z   | j                  j                  t        j                  |      z   | j                  j                  |z         j                  |      | _        | j                  j                  | j                  j                  _        d d d        y # 1 sw Y   y xY wr4   )r)   r$   no_gradr   r!   gradr   r   r   absr    collapser5   r"   r'   r   r   r8   r9   r:   deltas        r   flux_updatezFluxBatchNorm2d.flux_updateL   sp   ]]_ 	5ww~~"".dggnn111",GGNNU*OO%%		%(88OO%%-# (4.	 
 '+oo&7&7# ww||  ,dggll///!+GGLL5(NN$$uyy'77NN$$u," (4.	 
 %)NN$4$4!'	5 	5 	5s   GHHc                     | j                          | j                  j                  j                         j	                         | j
                  j                  j                         j	                         dS )N)gamma_entropybeta_entropy)r)   r   r   meanitemr   r(   s    r   entropy_statszFluxBatchNorm2d.entropy_statsc   sV    !__..335::< NN,,11388:
 	
r   r   )__name__
__module____qualname____doc__intfloatr   r)   r,   r$   Tensorr7   rB   dictrH   __classcell__r   s   @r   r   r      sj    
S E 	 	%,, 	5 5% 5u 5.
t 
r   r   c                        e Zd ZdZ	 	 	 ddedededededef fdZdd	efd
Z fdZ	d Z
dej                  dej                  fdZdej                  dedefdZdefdZdefdZ xZS )
FluxConv2dz1
    Convolutional layer with Flux tracking.
    in_channelsout_channelskernel_sizepaddingstrider'   c                 ~    t         |           t        j                  ||||||      | _        d | _        d | _        y )N)rY   rZ   r'   )r   r   r   Conv2dconvweight_flux	bias_flux)r   rV   rW   rX   rY   rZ   r'   r   s          r   r   zFluxConv2d.__init__r   s?     	II{F
	  r   init_entropyc                    | j                   `t        | j                  j                  j                  j                         t        j                  | j                  j                  j                        |z  t        j                  | j                  j                  j                              | _         | j                  j                  t        | j                  j                  j                  j                         t        j                  | j                  j                  j                        |z  t        j                  | j                  j                  j                              | _
        yyy)zInitialize flux tensors.Nr   )r^   r   r]   r!   r"   r#   r$   r%   r&   r'   r_   )r   r`   s     r   r)   zFluxConv2d.initialize_flux   s    #)		  %%++-//$))"2"2"7"78<G""499#3#3#8#89 D
 yy~~)!+IINN''--/oodiinn&9&9:\I&&tyy~~':':;" * $r   c                 N   t        |   |i | | j                  >| j                  j                  | j                  j                  j
                        | _        | j                  >| j                  j                  | j                  j                  j
                        | _        | S r+   )r   r,   r^   r]   r!   r-   r_   r.   s      r   r,   zFluxConv2d.to   s~    
D#F#'#//224993C3C3J3JKD>>%!^^..tyy/?/?/F/FGDNr   c                     | j                   /| j                   j                  | j                  j                  _        | j
                  0| j
                  j                  | j                  j                  _        yy)z/Copy flux values back to underlying parameters.N)r^   r5   r]   r!   r"   r_   r'   r(   s    r   _sync_flux_to_paramszFluxConv2d._sync_flux_to_params   sU    '$($4$4$6$6DII!>>%"&.."2"2DIINN &r   r1   r2   c                 ~   | j                          | j                  j                  j                  |j                  k7  r`| j                  j	                  |j                        | _        | j
                  *| j
                  j	                  |j                        | _        | j                          | j                  |      S r4   )r)   r^   r5   r-   r,   r_   rd   r]   r6   s     r   r7   zFluxConv2d.forward   s    $$0#//22188<D~~)!%!2!2188!<!!#yy|r   r8   r9   r:   c                 R   | j                          t        j                         5  | j                  j                  j
                  | | j                  j                  j
                  z  }t        | j                  j                  |z   | j                  j                  t        j                  |      z   | j                  j                  |z         j                  |      | _        | j                  j                  | j                  j                  _        | j                  j                  | j                  j                  j
                  | | j                  j                  j
                  z  }t        | j                  j                  |z   | j                  j                  t        j                  |      z   | j                  j                  |z         j                  |      | _        | j                  j                  | j                  j                  _        d d d        y # 1 sw Y   y xY wr4   )r)   r$   r<   r]   r!   r=   r   r^   r   r>   r    r?   r5   r"   r'   r_   r@   s        r   rB   zFluxConv2d.flux_update   s   ]]_ 	7yy$$0dii..333#-II$$u,$$&&5)99$$&&.$ (4.	  
 )-(8(8(:(:		  % yy~~)diinn.A.A.Mdiinn111!+IINNU*NN$$uyy'77NN$$u," (4.	 
 '+nn&6&6		#'	7 	7 	7s   G/HH&c                     | j                          | j                  j                  |      | _        | j                  j                  !| j
                  j                  |      | _        y y r4   )r)   r^   r?   r]   r'   r_   r   r:   s     r   r?   zFluxConv2d.collapse   sP    ++44T:99>>%!^^44T:DN &r   c                 $   | j                          d| j                  j                  j                         j	                         i}| j
                  j                  5| j                  j                  j                         j	                         |d<   |S )Nweight_entropybias_entropy)r)   r^   r   rF   rG   r]   r'   r_   )r   statss     r   rH   zFluxConv2d.entropy_stats   so    !4#3#3#5#5#:#:#<#A#A#CD99>>%$(NN$4$4$9$9$;$@$@$BE.!r   )r      FrI   )rJ   rK   rL   rM   rN   boolr   rO   r)   r,   rd   r$   rP   r7   rB   r?   rQ   rH   rR   rS   s   @r   rU   rU   m   s       	
   "E 3 %,, 7 7% 7u 7.;U ;t r   rU   c            
            e Zd ZdZdZ	 	 ddedededeej                     f fdZ	de
j                  d	e
j                  fd
Zde
j                  dedefdZdefdZd	ee   fdZ xZS )FluxBasicBlocku   
    ResNet Basic Block (for ResNet-18/34) with Flux tracking.
    
    Architecture:
        x → Conv → BN → ReLU → Conv → BN → +x → ReLU
    rm   rV   rW   rZ   
downsamplec                    t         |           || _        t        ||dd|      | _        t        |      | _        t        ||dd      | _        t        |      | _        t        j                  d      | _        || _        y )N   rm   rY   rZ   rY   Tinplace)r   r   rq   rU   conv1r   bn1conv2bn2r   ReLUrelurZ   )r   rV   rW   rZ   rq   r   s        r   r   zFluxBasicBlock.__init__   so     	$  \1aPVW
"<0lAqI
"<0GGD)	r   r1   r2   c                    |}| j                   | j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }||z  }| j                  |      }|S r4   )rq   rx   ry   r}   rz   r{   )r   r1   identityouts       r   r7   zFluxBasicBlock.forward   s{    ??&q)Hjjmhhsmiinjjohhsmxiin
r   r8   r9   r:   c                 F   | j                   j                  |||       | j                  j                  |||       | j                  j                  |||       | j                  j                  |||       | j
                  t        | j
                  t        j                        r|t        | j
                  d   t              r | j
                  d   j                  |||       t        | j
                  d   t              r!| j
                  d   j                  |||       yyyy)z%Update all flux states in this block.Nr   rm   )rx   rB   ry   rz   r{   rq   
isinstancer   
SequentialrU   r   )r   r8   r9   r:   s       r   rB   zFluxBasicBlock.flux_update  s    

tR.T2t,

tR.T2t,??&:door}}+U$//!,j9"..tR>$//!,o>"..tR> ? ,V&r   c                 L   | j                   j                  |       | j                  j                  |       | j                  bt	        | j                  t
        j                        r=t	        | j                  d   t              r| j                  d   j                  |       y y y y )Nr   )rx   r?   rz   rq   r   r   r   rU   rh   s     r   r?   zFluxBasicBlock.collapse  sy    

D!

D!??&:door}}+U$//!,j9"++D1 : ,V&r   c                     g }| j                   | j                  | j                  | j                  fD ]?  }|j	                         }|j                  |j                         D cg c]  }| c}       A |S c c}w )z)Collect all entropy values in this block.)rx   ry   rz   r{   rH   extendvalues)r   entropy_listlayerrl   r5   s        r   get_all_entropyzFluxBasicBlock.get_all_entropy  si    jj$((DJJA 	=E'')EELLN ;q ;<	=  !<s   !	A5
)rm   N)rJ   rK   rL   rM   	expansionrN   r	   r   Moduler   r$   rP   r7   rO   rB   r?   r   r   rR   rS   s   @r   rp   rp      s     I *.  	
 RYY'& %,, "? ?% ?u ?2U 2e r   rp   c                   0    e Zd ZdZ	 	 ddej
                  dee   dedef fdZ	d Z
	 ddej
                  ded	ed
edej                  f
dZ fdZdej                  dej                  fdZdej                  dedefdZdefdZdefdZdefdZ xZS )
FluxResNetzi
    ResNet with Flux Algebra tracking.
    
    Supports ResNet-18, 34, 50, 101, 152 architectures.
    blocklayersnum_classesr`   c                 f   t         |           d| _        || _        t	        ddddd      | _        t        d      | _        t        j                  d      | _
        t        j                         | _        | j                  |d|d         | _        | j                  |d|d   d	
      | _        | j                  |d|d	   d	
      | _        | j                  |d|d   d	
      | _        t        j$                  d      | _        t	        d|j(                  z  |d      | _        | j-                          y )N@   rs   rm   rt   Trv   r         rZ      i   rm   rm   )r   r   rV   r`   rU   rx   r   ry   r   r|   r}   Identitymaxpool_make_layerlayer1layer2layer3layer4AdaptiveAvgPool2davgpoolr   fc_initialize_all_flux)r   r   r   r   r`   r   s        r   r   zFluxResNet.__init__'  s    	(  2q!A>
"2&GGD)	{{} &&ub&)<&&uc6!9Q&G&&uc6!9Q&G&&uc6!9Q&G ++F3S5??2KC 	!!#r   c                 $   | j                   j                  | j                         | j                  j                          t	        | j                  j
                  j                  j                  j                         t        j                  | j                  j
                  j                  j                        | j                  z  t        j                  | j                  j
                  j                  j                              | j                  _        t	        | j                  j
                  j                  j                  j                         t        j                  | j                  j
                  j                  j                        | j                  z  t        j                  | j                  j
                  j                  j                              | j                  _        | j                  j                  | j                         y)z'Initialize flux tensors for all layers.r   N)rx   r)   r`   ry   r   r   r!   r"   r#   r$   r%   r&   r   r'   r   r   r(   s    r   r   zFluxResNet._initialize_all_fluxE  sB   

""4#4#45  "(HHKK##))+oodhhkk005569J9JJtxx{{11667

 (HHKK!!'')oodhhkk..334t7H7HHtxx{{//445

 	 1 12r   channelsblocksrZ   r2   c                    d }|dk7  s| j                   ||j                  z  k7  rPt        j                  t	        | j                   ||j                  z  d|      t        ||j                  z              }g }|j                   || j                   |||             ||j                  z  | _         t        d|      D ]$  }|j                   || j                   |             & t        j                  | S )Nrm   r   )rV   r   r   r   rU   r   appendrange)r   r   r   r   rZ   rq   r   _s           r   r   zFluxResNet._make_layerU  s     
Q;$**h.HH4++X-GSYZ5?? :;J
 eD,,h
KL#eoo5q&! 	=AMM% 0 0(;<	= }}f%%r   c                    t        |   |i |  | j                  j                  |i |  | j                  j                  |i |  | j                  j                  |i | | j
                  | j                  | j                  | j                  fD ]  }|D ]  } |j                  j                  |i |  |j                  j                  |i |  |j                  j                  |i |  |j                  j                  |i | |j                  |j                  D ]<  }t        |d      st        |t        j                        s+ |j                  |i | >   | S )z8Override to() to move all flux tensors to target device.r,   )r   r,   rx   ry   r   r   r   r   r   rz   r{   rq   hasattrr   r   r   )r   r/   r0   layer_groupr   moduler   s         r   r,   zFluxResNet.tol  s>   
D#F#

t&v&T$V$

D#F# KKdkk4;;O 		7K$ 7//		d-f-//		d-f-##/"'"2"2 7"640Z		5R%FIIt6v677		7 r   r1   c                    | j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }| j                  |      }| j                  |      }| j                  |      }| j                  |      }|j                  |j                  d      d      }|S )Nr   )rx   ry   r}   r   r   r   r   r   r   r   viewsizer6   s     r   r7   zFluxResNet.forward~  s    JJqMHHQKIIaLLLOKKNKKNKKNKKNLLOGGAJFF166!9b!r   r8   r9   r:   c                 L   | j                   j                  |||       | j                  j                  |||       | j                  | j                  | j
                  | j                  fD ]  }|D ]  }|j                  |||         | j                  j                  |||       y)z"Update flux states for all layers.N)rx   rB   ry   r   r   r   r   r   )r   r8   r9   r:   r   r   s         r   flux_update_allzFluxResNet.flux_update_all  s    

tR.T2t, KKdkk4;;O 	2K$ 2!!$D12	2 	D"d+r   c                    | j                   j                  |       | j                  j                  |       | j                  | j                  | j
                  | j                  fD ]  }|D ]  }|j                  |         y)z"Apply collapse to all flux states.N)rx   r?   r   r   r   r   r   )r   r:   r   r   s       r   collapse_allzFluxResNet.collapse_all  si    

D! KKdkk4;;O 	%K$ %t$%	%r   c                     | j                   j                         | j                  j                         | j                  j                         d}t	        | j
                  | j                  | j                  | j                  g      D ]u  \  }}g }|D ]!  }|j                  |j                                # t        j                  |      t        j                  |      t        j                  |      d|d|dz    <   w |S )z!Get comprehensive entropy report.)rx   ry   r   )rF   maxminr   rm   )rx   rH   ry   r   	enumerater   r   r   r   r   r   nprF   r   r   )r   reportir   layer_entropiesr   s         r   get_entropy_reportzFluxResNet.get_entropy_report  s     ZZ--/88))+'''')
 (dkk4;;PTP[P[(\] 	NA{ O$ @&&u'<'<'>?@ 0vvo.vvo.%FU1Q3%=!		 r   c                    d}| j                   j                         | j                  j                         fD ]  }|t        |j	                               z  }  | j
                  j                         j	                         D ]  }||z  }	 | j                  | j                  | j                  | j                  fD ]#  }|D ]  }|j                         D ]  }||z  }	  % |S )z.Calculate total entropy across all parameters.        )rx   rH   ry   sumr   r   r   r   r   r   r   )r   totalrl   r   r   ents         r   total_entropyzFluxResNet.total_entropy  s     jj..0$((2H2H2JK 	)ES((E	) WW**,335 	EUNE	 !KKdkk4;;O 	!K$ ! 002 !CSLE!!	!
 r   
   r   )rm   )rJ   rK   rL   rM   r   r   r   rN   rO   r   r   r   r   r,   r$   rP   r7   r   r   rQ   r   r   rR   rS   s   @r   r   r      s     !$yy$ S	$ 	$
 $<3* &yy& & 	&
 & 
&.$ %,, "	,ELL 	,e 	,5 	,% %D (u r   r   r   r`   r2   c                 *    t        t        g d| |      S )z$Create ResNet-18 with Flux tracking.)r   r   r   r   )r   rp   r   r`   s     r   flux_resnet18r     s    nlKNNr   c                      e Zd ZdZ	 	 	 	 	 	 ddedej                  dedededee	   de
d	efd
Zde
de
defdZ	 ddede
de
dedef
dZ ej"                         ddededefd       Z	 	 ddedede
dedef
dZddee	   fdZy)FluxTrainerz
    Training loop with Flux Algebra integration.
    
    Tracks entropy, applies collapse, and monitors learning dynamics.
    modelr-   r9   r   weight_decaywork_schedulecollapse_every_ncollapse_workc	                    || _         || _        t        j                  |j	                         D 	cg c]  }	|	j
                  s|	 c}	|||      | _        d | _        || _        || _	        || _
        g g g g g g g d| _        d| _        y c c}	w )N)r9   r   r   
train_loss	train_accval_lossval_accentropylearning_rateswork_budgetsr   )r   r-   optimSGD
parametersrequires_grad	optimizer	schedulerr   r   r   historyepoch)
r   r   r-   r9   r   r   r   r   r   ps
             r   r   zFluxTrainer.__init__  s     
 ((*>1aooQ>%	
  + 0*  
 
3 ?s   B	B	r   total_epochsr2   c                    | j                   dk(  ry| j                   dk(  rdd||z  z  z   S | j                   dk(  rdd|z  z  S | j                   dk(  r5dddt        j                  t        j                  |z  |z        z   z  d	z  z   S y)
z&Compute work budget based on schedule.constantg{Gz?linearg{Gz?exponentialg?cosinerm   r   )r   r   cospi)r   r   r   s      r   compute_work_budgetzFluxTrainer.compute_work_budget  s    +8+$%,"6777=03%<((8+$!bffRUUU]\-I&J"JKaOOOr   train_loaderverbosec                 
   | j                   j                          | j                  ||      }t        j                         }d}d}d}	t        |d|dz    d| d|       }
t        |
      D ]  \  }\  }}|j                  | j                        |j                  | j                        }}| j                  |      } |||      }| j                  j                          |j                          | j                  j                  d   d   }| j                  j                          | j                   j                  |j                         |d	z  |       |dz   | j                   z  dk(  r(| j                   j#                  | j$                  d	z         ||j'                         |j)                  d      z  z  }|j+                  d      \  }}|	|j)                  d      z  }	||j-                  |      j/                         j'                         z  }|
j1                  |j'                         d
d|z  |	z  dd| j                   j3                         dd        ||	z  }d|z  |	z  }||| j                   j3                         |dS )z'Train for one epoch with flux tracking.r   r   Epoch rm   /z [Train]descdisabler9   r   .4f      Y@.2f%)r8   accr   )r8   accuracyr   work_budget)r   trainr   r   CrossEntropyLossr
   r   r,   r-   r   	zero_gradbackwardparam_groupsstepr   detachr   r   r   rG   r   r   eqr   set_postfixr   )r   r   r   r   r   r   	criterionrunning_losscorrectr   pbar	batch_idxinputstargetsoutputsr8   r9   r   	predictedavg_lossr   s                        r   train_epochzFluxTrainer.train_epoch  sZ    	

..ulC'')	Lay,x'P 'K) -6dO "	(I($ii4gjj6MGF jj(GWg.D NN$$&MMO ,,Q/5B NN! JJ&&t{{}b3hL A!6!66!;

''(:(:S(@A DIIK&++a.88L";;q>LAyW\\!_$Ey||G,0027799G99;s+w,u,S13"jj668= ="	H  %''>E)  zz//1&	
 	
r   
val_loaderc                    | j                   j                          t        j                         }d}d}d}t	        |d|       }|D ]  \  }}	|j                  | j                        |	j                  | j                        }	}| j                  |      }
 ||
|	      }||j                         |j                  d      z  z  }|
j                  d      \  }}||	j                  d      z  }||j                  |	      j                         j                         z  }|j                  |j                         dd|z  |z  dd	d
        ||z  }d|z  |z  }||dS )z!Evaluate model on validation set.r   r   
Evaluatingr   rm   r   r   r   r   )r8   r   )r8   r   )r   evalr   r   r
   r,   r-   rG   r   r   r  r   r  )r   r  r   r  r  r  r   r  r
  r  r  r8   r   r  r  r   s                   r   evaluatezFluxTrainer.evaluateV  sW    	

'')	J\w;G# 	OFG$ii4gjj6MGFjj(GWg.DDIIK&++a.88L";;q>LAyW\\!_$Ey||G,0027799G99;s+w,u,S13 	   %''>E)  
 	
r   
num_epochsc                 "   t        dd        t        d       t        d        t        d| j                          t        d| j                  j                         d       t        d| j                          t        d d       t
        j                  j                  | j                  t        |dz        t        |d	z        gd
      | _
        d}d}t        |      D ]  }|| _        | j                  ||||      }| j                  ||      }	| j                  j                          | j                  j                   d   d   }
| j"                  d   j%                  |d          | j"                  d   j%                  |d          | j"                  d   j%                  |	d          | j"                  d   j%                  |	d          | j"                  d   j%                  |d          | j"                  d   j%                  |
       | j"                  d   j%                  |d          |s|dz   dz  dk(  r+t        d|dz   dd|d   dd|	d   dd |d   dd!|
d"
       |	d   |kD  s|	d   }|| j                  j'                         | j                  j'                         ||d   |d   d#} t        dd        t        d$       t        d%|dd&       t        d'| j"                  d   d(   d       t        d d       |S ))zFull training loop.
<============================================================zFlux ResNet TrainingzDevice: Initial entropy: r   zWork schedule:       ?g      ?r   )
milestonesgammar   Nr   r9   r   r8   r   r   r   r   r   r   r   r   rm   r   r   3dz
 | Train: z	% | Val: z% | Entropy: z | LR: r   )r   model_stateoptimizer_stater   r   r   zTraining Complete!Best Validation Accuracy: r   zFinal Entropy: r   )printr-   r   r   r   r   lr_schedulerMultiStepLRr   rN   r   r   r   r  r  r  r   r   r   
state_dict)r   r   r  r  r   best_val_accbest_model_stater   train_stats	val_stats
current_lrs              r   r   zFluxTrainer.trainz  s,    	6(m$&&'!$**":":"<S!ABC 2 2345m ++77NNJ,-s:3D/EF 8 
 :& (	EDJ **<
GTK j':I NN!44Q7=J LL&--k&.ABLL%,,[-DELL$++If,=>LL#**9Z+@ALL#**;y+ABLL)*11*=LL(//M0JK 519*a/uQwrl +  +J 7< ='
3C8 9""-i"8!= >',	. / $|3(4"#'::#8#8#:'+~~'@'@'B+!,Z!8*95$ C(	T 	6(m"$*<*<A>?Y 7 ;C@ABmr   N	save_pathc                    t        j                  ddd      \  }}|d   j                  | j                  d   ddd	       |d   j                  | j                  d
   ddd	       |d   j	                  d       |d   j                  d       |d   j                  d       |d   j                          |d   j                  dd       |d   j                  | j                  d   ddd	       |d   j                  | j                  d   ddd	       |d   j	                  d       |d   j                  d       |d   j                  d       |d   j                          |d   j                  dd       |d   j                  | j                  d   dddd       |d   j	                  d       |d   j                  d       |d   j                  d       |d   j                          |d   j                  dd       |d   j                  | j                  d    d!d"dd       |d   j	                  d       |d   j                  d!       |d   j                  d#       |d   j                          |d   j                  dd       |d   j                  d$       t        j                          |r&t        j                  |d%d&'       t        d(|        t        j                          y))*zPlot training history.r   )   r   )figsize)r   r   r   Trainors   )labelmarker
markersizer   
Validationr   EpochLosszLoss over TimeTg333333?)alpha)r   rm   r   r   zAccuracy (%)zAccuracy over Time)rm   r   r   zTotal Entropyred)r0  colorr1  r2  EntropyzFlux Entropy over Timer   r   zLearning RategreenzLearning Rate Schedulelog   tight)dpibbox_incheszPlots saved to N)pltsubplotsplotr   
set_xlabel
set_ylabel	set_titlelegendgrid
set_yscaletight_layoutsavefigr!  show)r   r*  figaxess       r   plot_historyzFluxTrainer.plot_history  s   LLAx8	T 	T
\2'#Z[\T
Z0S]^_T
g&T
f%T
-.T
T
C( 	T
[1YZ[T
Y/|C\]^T
g&T
n-T
12T
T
C( 	T
Y/e\_lmnT
g&T
i(T
56T
T
C( 	T
%56oU\ehuvwT
g&T
o.T
56T
T
C(T
e$KK	s@OI;/0
r   )r   ?Mb@?r      r  )T)d   Tr4   )rJ   rK   rL   rM   r   r$   r-   rO   r	   strrN   r   r   r   rn   rQ   r  r<   r  r   rN   r   r   r   r     sX    "'/ !")) ) 	)
 ) )  }) ) )V C E ( @
 @
 @
 	@

 @
 
@
D U]]_!
: !
 !
 !
 !
N J  J  J  	J 
 J  
J X-hsm -r   r   data_dir
batch_sizenum_workersc           	         t        j                  t        j                  dd      t        j                         t        j                         t        j
                  dd      g      }t        j                  t        j                         t        j
                  dd      g      }t        j                  | dd|      }t        j                  | dd|      }t        ||d|d	      }t        ||d|d	      }||fS )
z+Create CIFAR-10 train and test DataLoaders.       ru   )gHPs?gec]?g~jt?)gV-?g^I+?g(?T)rootr   download	transformF)rV  shufflerW  
pin_memory)	r   Compose
RandomCropRandomHorizontalFlipToTensor	Normalizer   CIFAR10r   )	rU  rV  rW  train_transformtest_transformtrain_datasettest_datasetr   test_loaders	            r   get_cifar10_loadersrk    s    !((b!,'')57OP	* O  ''57OP) N $$!	M ## 	L L K $$r   c                     t        d       t        d       t        d       t        d       ddddd	d
ddddt        j                  j                         rdndd} t        d       | j	                         D ]  \  }}t        d| d|         t                t        j
                  | d         }t        d       t        | d   | d         \  }}t        dt        |j                                t        dt        |j                         d       t        d       t        d| d         j                  |      }t        d |j                         D              }t        d  |j                         D              }t        d!|d"       t        d#|d"       t        d$|j                         d%d       t        ||| d&   | d'   | d(   | d)   | d*   | d+   ,      }	t        d-       t        j                         }
|	j!                  ||| d.   d/0      }t        j                         |
z
  }t        d1|d%d2|d3z  d%d4       t        d5       |j#                  |d6          t        d       t        d7       t        d8       |j%                          t'        j(                         }d9}d:}d:}d:gdz  }d:gdz  }g d;}t        j*                         5  t-        |d<=      D ]"  \  }}|j                  |      |j                  |      }} ||      } |||      }||j/                         |j1                  d:      z  z  }|j3                  d>      \  }}||j1                  d:      z  }||j5                  |      j                         j/                         z  }t7        |j1                  d:            D ]O  }||   }||xx   ||   j5                  |      j                         j/                         z  cc<   ||xx   d>z  cc<   Q % 	 d?d?d?       ||z  }d@|z  |z  }t        dA|dB       t        dC|d%dD       t        dE| dF|        t        dG       t        dH       t9        |      D ]0  \  }}||   d:kD  sd@||   z  ||   z  }t        d|dId|dJdD       2 t        d       t        dK       t        d8       |j;                         } t        dL|j                         d%       t        dM|	j<                  dN   dO   d%       t        dP       | j	                         D ]?  \  }!}"t?        |"t@              sdQ|"v st        d|!dRdS|"dQ   dBdT|"dU   dBdV|"dW   dB       A t        d       t        dX       t        d8       | |||dY   |dZ   |	j<                  dN   dO   |d[|	j<                  dZ   |	j<                  dY   |	j<                  dN   d\d]}#t        d^|dY   d%dD       t        d_|d%dD       t        d`|d%da       t        db| d   dcz  d%dd|	j<                  dN   dO   d%       t        jB                  de      }$df|$ dg}%t        jD                  |d6   |dh   | ||	j<                  dN   dO   di|%       t        dj|%        dk|$ dl}&tG        |&dm      5 }'tI        jJ                  |	j<                  dn   |	j<                  dZ   |	j<                  do   |	j<                  dY   |	j<                  dN   |	j<                  dp   |	j<                  dq   dr|'ds       d?d?d?       t        dt|&        	 du|$ dv}(|	jM                  |(w       t        d       t        d{       t        d       |#S # 1 sw Y   kxY w# 1 sw Y   `xY w# tN        $ rG})t        dx|)        t        dy       	 |	jM                          n#  t        dz       Y nxY wY d?})~)d?})~)ww xY w)|z#Main training and testing pipeline.z=
============================================================zFlux ResNet-18 on CIFAR-10z;Based on: The Algebra of Flux (Conditional Collapse Theory)z=============================================================
r   2   r   rO  rP  r   r   r  r   cudacpu)rV  r  learning_rater   r   r   r   r   r`   rW  r-   zConfiguration:z  z: r-   zLoading CIFAR-10...rV  rW  )rV  rW  zTrain samples: zTest samples: r  zCreating Flux ResNet-18...r`   r   c              3   <   K   | ]  }|j                           y wr4   )numel.0r   s     r   	<genexpr>zmain.<locals>.<genexpr>`  s     =Qqwwy=s   c              3   V   K   | ]!  }|j                   s|j                          # y wr4   )r   rr  rs  s     r   ru  zmain.<locals>.<genexpr>a  s     TAOO1779Ts   ))zTotal parameters: ,zTrainable parameters: r  r   rp  r   r   r   r   r   )r   r-   r9   r   r   r   r   r   zStarting training...
r  T)r   r  r  r   zTraining time: zs (<   zmin)
z'Loading best model for final testing...r  zFINAL EVALUATION ON TEST SETr  r   r   )
planecarbirdcatdeerdogfroghorseshiptruckTesting)r   rm   Nr   z
Test Loss: r   zTest Accuracy: r   z	Correct: r   z
Per-class Accuracy:z(----------------------------------------8sz6.2fzENTROPY REPORTz
Total Entropy: zFinal Entropy (from training): r   r   z
Layer-wise Entropy:rF   12sz: mean=z, max=r   z, min=r   zTRAINING SUMMARYr   r   )test_accuracy	test_lossr%  best_train_accfinal_entropytraining_time)r   r   r   )configresultsr   r   zFinal Test Accuracy:      zTraining Time:            r   zEntropy Reduction:        iا u    → z%Y%m%d_%H%M%Sflux_resnet18_cifar10_z.pthr  )r  r  r  r  r   z
Model saved to: flux_resnet18_history_z.jsonwr   r   r   r   r   )indentzHistory saved to: flux_resnet18_training_z.png)r*  z
Could not save plot: zTrying to show plot instead...z#Plotting not available (no display)z!Flux ResNet-18 CIFAR-10 Complete!)(r!  r$   rn  is_availableitemsr-   rk  lendatasetr   r,   r   r   r   r   timer   load_state_dictr  r   r   r<   r
   rG   r   r   r  r   r   r   r   r   rQ   strftimesaveopenjsondumprN  	Exception)*r  keyvaluer-   r   rj  r   total_paramstrainable_paramstrainer
start_time
best_stater  r  r  r  r   class_correctclass_totalcifar_classesr
  r  r  r8   r   r  r   r0  avg_test_lossr  
class_namer   entropy_report
layer_namerl   summary	timestampmodel_save_pathhistory_save_pathfplot_save_pathes*                                             r   mainr  2  st   	-	
&'	
GH	- !!JJ335&5F 

lln #
U3%r%!"#	G\\&*+F 

  3,'=)!L+ 
OC 4 456
78	N3{2234B
78 

&'N+ 	bj 
 =%*:*:*<==LTe.>.>.@TT	|A.
/0	"#3A"6
78	e113C8
;< /"
#N+_- 23_-	G 

"#J!,'	  J IIK*,M	OM#.c-2B31Gv
NO 

34	*]34 
-	
()	&M	JJL##%IIGE C"HM#(K>M 
 (#Ki@ 	(OFG$ii/F1CGFFmGWg.Dv{{1~55I";;q>LAyW\\!_$Ey||G,0027799G 7<<?+ (
e$	!(>(B(B(D(I(I(KK$E"a'"(	((& %M7NU*M	M-,
-.	OM#.a
01	IgYaw
'( 

!"	(O"=1 5:q>Aq))KN:CBz"oRDz345 
-	
	&M--/N	e113C8
9:	+GOOI,Fr,J3+O
PQ	
!"+113 E
EeT"vBz#&geFmC-@ Auc*&uc0BD EE 
-	
	&M *&&y1(5$__Y7;*
 !5y1y1
G$ 
&z)'<S&A
CD	&}S&9
;<	&}S&9
;<	&N#h.s35__Y'+C02 3
 o.I /yk>O	JJ!-0%&78&??9-b1  
/
01 15A		% 			!//,7 5
3y1y1%oo.>?#OON;
 Q		 
01
23	929+TB~6 
-	
-.	-Nw( (~	 	   9's+,./	9  "	9789sJ   D4]1#A<]>6^
 1];>^
	__-^>=_>_____main__r   )z./datar   r   )+rM   r$   torch.nnr   torch.nn.functional
functionalFtorch.optimr   torchvisionr   r   torch.utils.datar   matplotlib.pyplotpyplotr@  numpyr   typingr   r   r   r	   r  r
   r  flux_tensorr   r   r   rU   rp   r   rN   rO   r   r   rS  rk  r  rJ   r  rT  r   r   <module>r     s        , '   . .    "
L
bii L
ba aLHRYY HZi iXOs Ou Oz O` `L	 5%5%5% 5% :z!"	5%t_D zfG r   