
    iH                     8   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	Z	ddl
mZ ddlmZ ddlZ G d dej                         Z G d dej                         Z G d d	ej                         Zdd
Zd Z G d dej                         Zd Zedk(  r e        yy)uK  
Singularity Memory ODE (SM-ODE) for CIFAR-10 — FIXED VERSION
Based on the theory: Memory stored in shared singularity, not lost between states.

Bugs fixed:
  1. Batch size mismatch in compute_memory_bias (torch.cat crash)
  2. Channel vs memory_dim mismatch in SharedSingularityBlock (blocks 1 & 3 crash)
  3. reset_singularity() never called — singularity grows unboundedly
  4. Singularity updated during evaluation — test data contamination
  5. Registered buffer replaced by plain tensor — breaks state_dict
  6. Baseline trained with misleading singularity-named function
    N)
DataLoaderc                   t     e Zd ZdZd	 fd	Zed        Zej                  d        Z fdZd Z	d Z
d Z xZS )
SingularityMemoryaO  
    Stores historical states in a 'singularity' that influences current predictions.

    Theory: History is not lost between states - it's stored in a singularity field.
    This allows 'variable determinism' where predictions are constrained by laws
    (standard weights) but influenced by accumulated memory (novel approach).
    c                    t         |           || _        || _        || _        d | _        t        j                  ||      | _        t        j                  t        j                  ||z   |      t        j                               | _        t        j                  t        j                  ||      t        j                               | _        ||k7  rt        j                  ||      | _        y t        j                         | _        y N)super__init__feature_dim
memory_dimdecay_singularitynnLinearto_singularity
SequentialSigmoidmemory_gateTanh
kernel_netmemory_to_featureIdentity)selfr
   r   r   	__class__s       ex03.pyr	   zSingularityMemory.__init__%   s    &$

 ! !iiZ@ ==IIj;.
;JJL
 --IIj*-GGI
 $%'YYz;%GD"%'[[]D"    c                     | j                   S r   r   )r   s    r   singularityzSingularityMemory.singularityF   s       r   c                     || _         y r   r   )r   values     r   r   zSingularityMemory.singularityJ   s
    !r   c                 l    t         |   |       | j                   || j                        | _        | S )z;Propagate device / dtype changes to the singularity tensor.)r   _applyr   )r   fnr   s     r   r"   zSingularityMemory._applyN   s3    r( "4#4#4 5Dr   c                     t        | j                               j                  }t        | j                               j                  }t	        j
                  || j                  ||      | _        y)z8Reset singularity for a given batch size.  Device-aware.)devicedtypeN)next
parametersr%   r&   torchzerosr   r   )r   
batch_sizer%   r&   s       r   reset_singularityz#SingularityMemory.reset_singularityU   sO    doo'(//T__&'--!KK
DOO/5UDr   c                    |j                  d      }| j                  | j                  j                  d      |k7  r| j                  |       | j                  j                         }| j	                  |      }|d|z  z   }t        j                  ||gd      }| j                  |      }|t        j                  |      z  dz  }| j                  |      }	||	z   }
|
S )z|
        Compute how past states influence current prediction.
        B(S, y) = integral of K(t-tau) * y(tau) dtau
        r   皙?dimg333333?)
sizer   r,   detachr   r)   catr   tanhr   )r   current_featuresr+   r   kernel_outputmemory_bias
gate_inputgatememory_contribmemory_contrib_projbiased_featuress              r   compute_memory_biasz%SingularityMemory.compute_memory_bias\   s     &**1-
#t'7'7'<'<Q'?:'M"":. &&--/ 4 "C-$77 YY-=>BG

+ 

; 77#= #44^D*-@@r   c                     | j                   syt        j                         5  | j                  |      }| j                  | j
                  z  |j                         z   | _        ddd       y# 1 sw Y   yxY w)a1  
        Update singularity memory with new states.
        S(t) = S(t-1) * decay + new_features

        FIX #4: skip update during evaluation to prevent test-data contamination.
        Keep the update out of autograd so the persistent state never
        captures graph history across batches.
        N)trainingr)   no_gradr   r   r   r3   )r   new_featuresnew_projs      r   update_singularityz$SingularityMemory.update_singularity|   sc     }}]]_ 	S**<8H $ 1 1DJJ >AR RD	S 	S 	Ss   AA,,A5)   gffffff?)__name__
__module____qualname____doc__r	   propertyr   setterr"   r,   r>   rD   __classcell__r   s   @r   r   r      sS    3B ! ! " "D@Sr   r   c                   *     e Zd ZdZd fd	Zd Z xZS )SharedSingularityBlockz
    Multiple processing heads share a common singularity memory.

    Theory: Entangled particles share singularity memory (Phi_AB).
    Here, different feature channels share the same memory field.
    c                 T   t         |           || _        t        ||      | _        t        j                  t        d      D cg c]W  }t        j                  t        j                  ||dd      t        j                  |      t        j                  d            Y c}      | _        t        j                         | _        t        j                  t        j                  |dz  |dz        t        j                         t        j                  |dz  |            | _        y c c}w )N      paddingTinplace      )r   r	   channelsr   r   r   
ModuleListranger   Conv2dBatchNorm2dReLUheadsr   memory_projr   fusion)r   rY   r   _r   s       r   r	   zSharedSingularityBlock.__init__   s      -XzB ]]
 Qx$

 	 MM		(Ha;x(%$
 
 ;;= mmIIhlHqL1GGIIIhlH-
$
s   AD%c           
      0   |j                   \  }}}}| j                  D cg c]
  } ||       }}t        j                  |d      }t	        j
                  |d      j                  d      j                  d      }	| j                  j                  |	      }
| j                  j                  |	       | j                  |
      }|j                  d      j                  d      j                  dd||      }| j                  t        j                  |j                  |d||z        j                  d      |
gd            }|j                  d      j                  d      j                  dd||      }|d|z  z   dt        j                   |      z  z   S c c}w )NrR   r0   r/   g?r.   )shaper_   r)   r4   Fadaptive_avg_pool2dsqueezer   r>   rD   r`   	unsqueezeexpandra   viewmeanr5   )r   xr+   rY   hwheadhead_outputshead_concatpooledr=   memory_alignedmemory_expandedfuseds                 r   forwardzSharedSingularityBlock.forward   sw   %&WW"
Ha -1JJ7DQ77ii!4 &&q!,44R8@@D **>>vF 	++F3 ))/:(2226@@DKKBPRTUWXY EIIZQU388<'
  
 #--b188RAF 3;uzz/'B!BBB3 8s   F)   rF   rG   rH   rI   r	   rv   rL   rM   s   @r   rO   rO      s    
4Cr   rO   c                   0     e Zd ZdZd fd	Zd Zd Z xZS )SMODECIFAR10aF  
    Singularity Memory ODE Network for CIFAR-10

    Key innovations:
    1. Memory singularity stores history, doesn't lose it
    2. Multiple blocks share singularity memory (entanglement-like)
    3. Predictions = f(laws) + alpha * memory_bias
    4. Variable determinism: constrained by weights, influenced by memory
    c                 &   t         |           t        j                  t        j                  dddd      t        j
                  d      t        j                  d      t        j                  dddd      t        j
                  d      t        j                  d            | _        t        d|      | _	        t        d|      | _
        t        d|      | _        t        j                  t        j                  dddd	d
      t        j
                  d      t        j                  d            | _        t        j                  t        j                  dddd	d
      t        j
                  d      t        j                  d            | _        t        d|d	z        | _        t        j                  t        j                   dd      t        j                         t        j"                  d      t        j                   d|            | _        y )NrQ   @   rR   rS   TrU   rw   rE   rX   )striderT   i   g      ?)r   r	   r   r   r\   r]   r^   
input_convrO   shared_singularity1shared_singularity2shared_singularity3transition1transition2r   global_singularityr   Dropout
classifier)r   num_classesr   r   s      r   r	   zSMODECIFAR10.__init__   sx    --IIaQ*NN2GGD!IIb"a+NN2GGD!
 $:"j#I #9#z#J #9#z#J  ==IIb#qA6NN3GGD!
 ==IIc3!Q7NN3GGD!
 #4Ca"H --IIi%GGIJJsOIIc;'	
r   c                 ^   | j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }t        j                  |d      j                  d      j                  d      }| j                  j                  |      }| j                  j                  |       t        j                  |d      j                  d      j                  d      }t        j                  ||gd      }| j                  |      S )NrR   r/   r0   )r~   r   r   r   r   r   re   rf   rg   r   r>   rD   r)   r4   r   )r   rl   rr   memory_featurescombineds        r   rv   zSMODECIFAR10.forward  s   OOA $$Q' Q $$Q' Q $$Q' &&q!,44R8@@D11EEfM226: &&q!,44R8@@D99fo6B?x((r   c                 r    | j                         D ]$  }t        |t              s|j                  |       & y)z6FIX #3: reset every singularity module in the network.N)modules
isinstancer   r,   )r   r+   modules      r   reset_all_singularitiesz$SMODECIFAR10.reset_all_singularities)  s0    lln 	5F&"34((4	5r   )
   rw   )rF   rG   rH   rI   r	   rv   r   rL   rM   s   @r   rz   rz      s    (
T)<5r   rz   c                    t        j                         }t        j                  | j	                         dd      }t        j
                  j                  ||      }d}g }	t        |      D ]  }
| j                          d}d}d}t        | d      r| j                  |j                         |D ]  \  }}|j                  |      |j                  |      }}|j                           | |      } |||      }|j                          t        j                   j                   j#                  | j	                         d	
       |j%                          ||j'                         z  }|j)                  d      \  }}||j+                  d      z  }||j-                  |      j/                         j'                         z  } |j%                          d|z  |z  }t1        | ||      }|	j3                  |
dz   |||t5        |      z  d       ||kD  r&|}t        j6                  | j9                         d       t;        d|
dz    d| d|t5        |      z  dd|dd|dd|dd        |	|fS )u  
    Generic training loop for both SM-ODE and baseline models.

    FIX #3: Resets singularity at the start of every epoch so memory
            accumulates only within each epoch (not across epochs).
    FIX #4: Singularity modules skip updates during evaluation
            (handled inside SingularityMemory.update_singularity).
    FIX #6: Renamed from train_with_singularity — works for any model.
    gMbP?g-C6?)lrweight_decay)T_maxr   g        r   )r+   g      ?)max_normrR         Y@)epoch	train_acctest_acclosszsmode_cifar10_best.pthzEpoch /z	 | Loss: z.4fz
 | Train: .2fz
% | Test: z
% | Best: %)r   CrossEntropyLossoptimAdamWr(   lr_schedulerCosineAnnealingLRr[   trainhasattrr   r+   to	zero_gradbackwardr)   utilsclip_grad_norm_stepitemmaxr2   eqsumevaluateappendlensave
state_dictprint)modeltrainloader
testloaderr%   epochs	criterion	optimizer	schedulerbest_acchistoryr   running_losscorrecttotalinputstargetsoutputsr   rb   	predictedr   r   s                         r   train_modelr   4  s`    ##%IE,,.5tLI""44Yf4MIHGv -Z 534))[5K5K)L* 	:OFG$ii/F1CGF!FmGWg.DMMO HHNN**5+;+;+=*LNNDIIK'L";;q>LAyW\\!_$Ey||G,0027799G!	:$ 	7NU*	E:v6QY"  3{#33	
 	 hHJJu'')+CDuQwiq	,s;?O2OPS1T U!#j#jRUVWY 	ZY-Z^ Hr   c                    | j                          d}d}t        j                         5  |D ]  \  }}|j                  |      |j                  |      }} | |      }|j	                  d      \  }}	||j                  d      z  }||	j                  |      j                         j                         z  } 	 ddd       d|z  |z  S # 1 sw Y   xY w)z
    Evaluate model.

    FIX #4: Singularity update_singularity() now checks self.training,
            so memory is never contaminated by test data.
    r   rR   Nr   )	evalr)   rA   r   r   r2   r   r   r   )
r   r   r%   r   r   r   r   r   rb   r   s
             r   r   r   w  s     
JJLGE	 :) 	:OFG$ii/F1CGFFmG";;q>LAyW\\!_$Ey||G,0027799G	:: '>E!!: :s   BCCc                   *     e Zd ZdZd fd	Zd Z xZS )BaselineCIFAR10z<Standard ResNet-like architecture without singularity memoryc                    t         |           t        j                  t        j                  dddd      t        j
                  d      t        j                  d      t        j                  dddd      t        j
                  d      t        j                  d      t        j                  d      t        j                  dddd      t        j
                  d      t        j                  d      t        j                  dddd      t        j
                  d      t        j                  d      t        j                  d      t        j                  dd	dd      t        j
                  d	      t        j                  d      t        j                  d	d	dd      t        j
                  d	      t        j                  d      t        j                  d            | _	        t        j                  d	|      | _        y )
NrQ   r|   rR   rS   TrU   rX   rw   rE   )r   r	   r   r   r\   r]   r^   	MaxPool2dAdaptiveAvgPool2dfeaturesr   r   )r   r   r   s     r   r	   zBaselineCIFAR10.__init__  sR   IIaQ*NN2GGD!IIb"a+NN2GGD!LLOIIb#q!,NN3GGD!IIc31-NN3GGD!LLOIIc31-NN3GGD!IIc31-NN3GGD!  #/
2 ))C5r   c                     | j                  |      }|j                  |j                  d      d      }| j                  |      S )Nr   r/   )r   rj   r2   r   )r   rl   s     r   rv   zBaselineCIFAR10.forward  s7    MM!FF166!9b!q!!r   )r   rx   rM   s   @r   r   r     s    F6:"r   r   c            	         t        j                  t         j                  j                         rdnd      } t	        d|         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                  j                  d	d
d
|      }t        j                  j                  d	dd
|      }t        |dd
d      }t        |ddd      }t	        d       t	        d       t	        d       t	        d       t	        d       t	        d       t	                t	        d       t        dd      j!                  |       }t#        d |j%                         D              }t	        d|d       t'        |||| d      \  }	}
t	        d       t)        d      j!                  |       }t#        d |j%                         D              }t	        d |d       t'        |||| d      \  }}t	        d!       t	        d"       t	        d       t	        d#|
d$d%       t	        d&|d$d%       t	        d'|
|z
  d$d%       t	        d(       t	        d)       t	        d*       t	        d+       t	        d,       y )-NcudacpuzUsing device:     rW   rS   )gHPs?gec]?g~jt?)gۊe?ggDio?g|?5^?z../dataT)rootr   download	transformFrw   rX   )r+   shufflenum_workersz<============================================================z.SINGULARITY MEMORY ODE (SM-ODE) CIFAR-10 MODELz>
Theory: Memory stored in singularity, not lost between statesz:         -> Predictions use: f(laws) + alpha * memory_biaszA         -> Multiple blocks share singularity (like entanglement)z9
--- Training SM-ODE Model (with Singularity Memory) ---
r   )r   r   c              3   <   K   | ]  }|j                           y wr   numel.0ps     r   	<genexpr>zmain.<locals>.<genexpr>  s     CQqwwyC   zSM-ODE Parameters: ,2   )r   z2
--- Training Baseline Model (no Singularity) ---
)r   c              3   <   K   | ]  }|j                           y wr   r   r   s     r   r   zmain.<locals>.<genexpr>  s     FQqwwyFr   zBaseline Parameters: z=
============================================================zRESULTS SUMMARYz$
SM-ODE Model (Singularity Memory): r   r   z"Baseline Model (No Memory):       z#Improvement:                       z$
--- Singularity Memory Analysis ---zMemory dimension: 128z#Memory decay: 0.95 (high retention)zShared singularity blocks: 3z;Theory: Memory accumulates across batches within each epoch)r)   r%   r   is_availabler   
transformsCompose
RandomCropRandomHorizontalFlipToTensor	NormalizetorchvisiondatasetsCIFAR10r   rz   r   r   r(   r   r   )r%   transform_traintransform_testtrainsettestsetr   r   smode_modeltotal_paramssmode_history
smode_bestbaseline_modelbaseline_historybaseline_bests                 r   mainr     s   \\EJJ$;$;$=&5IF	N6(
#$ !((b!,'')57OP	* O  ''57OP) N
 ##++dT_ , H ""**edn + G X#tQRSKGUPQRJ	(O	
:;	(O	
KL	
FG	
MN	G 

GH2#>AA&IKC+*@*@*BCCL	Q/
01 +[*fR!M: 

@A$477?NF.*C*C*EFFL	!,q!1
23&1Z'#m
 
/	
	(O	1*S1A
CD	.}S.A
CD	/
]0J3/Oq
QR 

12	!#	/1	(*	GIr   __main__)r   )rI   r)   torch.nnr   torch.nn.functional
functionalre   torch.optimr   r   torchvision.transformsr   torch.utils.datar   numpynpModuler   rO   rz   r   r   r   r   rF    r   r   <module>r     s          + ' mS		 mSh?CRYY ?CLW5299 W5|@F"4#"bii #"THJV zF r   