
    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
mZ ddlZddlZddlmZmZmZmZ ddlmZ ddlZdZdZdZd	Zd
ZdZd	ZdZdZdez  Z e G d d             Z! e!       Z" G d d      Z# G d de
      Z$ G d dejJ                        Z& G d dejJ                        Z' G d dejJ                        Z( G d dejJ                        Z) G d dejJ                        Z* G d d ejJ                        Z+ G d! d"      Z,	 	 	 	 	 d0d#ejJ                  d$ed%ed&e-d'e.d(e/d)e/d*e/fd+Z0d1d,Z1d1d-Z2d. Z3e4d/k(  r e3       \  Z5Z6Z7yy)2a  
=============================================================
MMA-CCT: Multi-Messenger Astronomy - Conditional Collapse Theory
Neutrino Detection System based on CCT-ODE Framework

Architecture: Transformer with Cross-Channel Sensitivity Tracking
Training: Particle Statistics (Fermi-Dirac) + Sensitivity Amplification
Detection: Rare Event Hypersensitivity + Entropic Collapse

Author: CCT-ODE Framework
Version: 1.0
=============================================================
    N)Dataset
DataLoader)DictListTupleOptional)	dataclassg   JxޱAgmjݤ<g    .Ag    eA)皙?      Y@g     @@gh㈵>      ?g?c                       e Zd ZU dZeZeed<   eZ	eed<   e
Zeed<   dZeed<   dZeed<    ej                  d	ej                   z  d
z        Zeed<   dZeed<   dZeed<   y)PhysicsConstantsz*Container for physics simulation constantsc	mev_to_ev	gev_to_evgq=
ףp?ice_densitygHzG?refractive_index)      cherenkov_angler   absorption_lengthg      9@scattering_lengthN)__name__
__module____qualname____doc__C_LIGHTr   float__annotations__	MEV_TO_EVr   	GEV_TO_EVr   r   r   mathcospir   r   r        ex01.pyr   r   .   sp    4Au Iu  Iu K"e"%TXXb477lS&89OU9$u$#u#r&   r   c                      e Zd ZdZdZdZdZdZdZddZ		 	 	 	 dde
dee   d	ee   d
eej                     deeej                  f   f
dZdee   d	ee   d
eej                     deeej                  f   fdZdeeej                  f   fdZd	edededefdZded	ed
ej                  dedej                  f
dZdej                  d	ed
ej                  dededej                  fdZy)NeutrinoSimulatoraF  
    Realistic neutrino detection simulator based on IceCube-like detector.
    
    Simulates:
    - Cherenkov light emission from charged particles
    - Photon propagation and arrival times
    - Background noise (cosmic rays, atmospheric muons)
    - Flavor oscillation effects
    - Energy-dependent event signatures
    r         c                 (    || _         t        | _        y N)devicePHYSICSphysics)selfr.   s     r'   __init__zNeutrinoSimulator.__init__V   s    r&   N	is_signalflavorenergy	directionreturnc                 T    |r| j                  |||      }|S | j                         }|S )ai  
        Simulate a single neutrino or background event.
        
        Returns:
            Dictionary with:
            - 'hits': [num_hits, 4] tensor (x, y, z, t)
            - 'features': [num_features] tensor (energy, direction, time, etc.)
            - 'label': int (0=background, 1=signal)
            - 'metadata': dict with event properties
        )_simulate_neutrino_simulate_background)r1   r3   r4   r5   r6   events         r'   simulate_eventz NeutrinoSimulator.simulate_eventZ   s7      ++FFIFE  --/Er&   c           	         |%t        j                  ddd      j                         }|7t        j                  d      j                         }d|dz  z  }t	        |d      }|t        j                  d      j                         t
        j                  z  }t        j                  d      j                         dz  t
        j                  z  }t        j                  t        j                  |      t        j                  |      z  t        j                  |      t        j                  |      z  t        j                  |      g      }t        j                  ddd      j                         }| j                  |||      }| j                  ||||      }| j                  |||||      }	||	t        j                  d      |||||dd	S )
z$Simulate a real neutrino interactionr      r*   r*   
   r   r+   )r4   r5   r6   interaction_typenum_hitshitsfeatureslabelmetadata)torchrandintitemrandminr"   r$   tensorsinr#   _calculate_hits_generate_cherenkov_hits_extract_features)
r1   r4   r5   r6   thetaphirA   rB   rD   rE   s
             r'   r9   z$NeutrinoSimulator._simulate_neutrinoq   s    >]]1a.335F >ZZ]'')FFQJ'F'F JJqM&&(4772E**Q-$$&*TWW4C$((3-/$((3-/& I !==At499; ''8HI,,Xvy&Q ))$	6K[\  \\!_  &$4$	
 	
r&   c                    t        j                  ddd      j                         }t        j                  |      dz
  dz  }t        j                  |      dz
  dz  }t        j                  |      dz
  dz  }t        j                  |      dz  }t        j                  ||||gd      }t        j                  d      j                         dz  }t        j                  d      j                         t
        j                  z  }t        j                  d      j                         dz  t
        j                  z  }	t        j                  t        j                  |      t        j                  |	      z  t        j                  |      t        j                  |	      z  t        j                  |      g      }
t        j                  d	      }||d
<   |
|dd |dz  |d<   ||t        j                  d
      d|dddS )z(Simulate background event (muons, noise)   2   r?         ?  r*   dimr+      r      r   
background)r4   r5   typerC   )rH   rI   rJ   rK   stackr"   r$   rM   rN   r#   zeros)r1   rB   xyztrD   r5   rR   rS   r6   rE   s               r'   r:   z&NeutrinoSimulator._simulate_background   s    ==B-224 ZZ!C'4/ZZ!C'4/ZZ!C'4/ JJx 4'{{Aq!Q<Q/ A##%+ 

1""$tww.jjm  "Q&0LLHHUOdhhsm+HHUOdhhsm+HHUO"
 	 ;;r?!1&  \\!_ $		
 		
r&   rA   c                     |dk(  rd|z  nd}|dkD  r d}t        |dz  d      }||z  dz  |z  dz  }nd}d	}t        j                  j                  ||z        }	t	        |	d
      S )z.Calculate number of Cherenkov photons detectedr   g      ?        rW   i^  rU   rX   ư>{Gz?r*   )rL   nprandompoissonmax)
r1   r5   r4   rA   E_leptondNdxtrack_lengthnum_photons
efficiencydetecteds
             r'   rO   z!NeutrinoSimulator._calculate_hits   s~    
 %5$94&=sc>Dx!|T2L/D0<?$FKK 
 99$$[:%=>8Qr&   rB   c                    g }|| j                   k(  rt        |dz  d      }d|dz  z   }n>|| j                  k(  rt        |dz  d      }d|dz  z   }nt        |dz  d	      }d
|dz  z   }t        |      D ]  }t	        j
                  d      j                         |z  }	t	        j
                  d      j                         dz  t        j                  z  }
|dt	        j
                  d      j                         z   z  }||	z  }t	        j                  |d   |d    dg      }|t	        j                  |      dz   z  }t        j                  j                  ||      }|||t        j                  |
      z  |t        j                  |
      z  z   z  z   }t	        j                  |      }|t        j                   z  t        j"                  dz  z  }|t	        j
                  d      j                         dz  z  }|j%                  |d   |d   |d   |g        t	        j                  |t        j&                        }|j(                  d   dkD  r|ddddf   dz  |ddddf<   |S )z&Generate realistic Cherenkov ring hitsr@   rX   rV   rU   r+         r>   i,  (   r*   rW   r   rg   rh   g&.>dtypeNg     @@)FLAVOR_MUONrL   FLAVOR_ELECTRONrangerH   rK   rJ   r"   r$   rM   normlinalgcrossr#   rN   r/   r   r   appendfloat32shape)r1   rB   r5   r6   r4   rD   rp   ring_radiusit_trackrS   r_offsetbase_posperp_xperp_yposdisttimehits_tensors                      r'   rP   z*NeutrinoSimulator._generate_cherenkov_hits   sA   
  T%%%v{D1Lvz/Kt+++vz3/Lvz/K vz3/Lvz/K x 	8Ajjm((*\9G **Q-$$&*TWW4C"cEJJqM,>,>,@&@AH !7*H\\9Q<)A,"DEFuzz&1D89F\\''	6:FX$((3-)?&488TW=BX)XYYC ::c?D'111WYY5EFDEJJqM&&(2--DKKQQQ67+	8. ll4u}}= Q!#!,QU!3e!;K2A2r&   rD   c                    t        j                  d      }t        j                  |dz         dz  |d<   ||dd |j                  d   dz  |d<   |j                  d   dkD  r|d	d	d	df   j                  d
      }||dd |d	d	df   j                         }|dz  |d<   |d	d	d	df   |z
  }	|	j                  |	z  t        |	j                  d   d      z  }
t         j                  j                  |
      }|d   |j                         dz   z  }||d<   |dz  |d<   |dz  |d<   |S t        j                  d      |dd t        j                  d      |dd |S )z)Extract physics features from hit patternr[   rh   g      @r   r*   r\   r   r>   NrY   rU      r]   	   r@   r      )rH   ra   r"   log10r   meanstdTrm   r~   eigvalshsum)r1   rD   r5   r6   r4   rA   rE   centroidtime_spreadcenteredcoveigenvaluestrack_scores                r'   rQ   z#NeutrinoSimulator._extract_features  sp   
 ;;r? jj$/#5!1jjme+::a=1ArrE{''A'.H$HQqM q!t*..*K%-HQK ArrE{X-H**x'#hnnQ.?*CCC,,//4K%b/[__->-EFK%HQK "E>HRL+c1HRL  "KKNHQqM"[[^HQrNr&   )cpu)FNNN)r   r   r   r   r{   rz   
FLAVOR_TAUINTERACTION_CCINTERACTION_NCr2   boolr   intr   rH   Tensorr   strr<   r9   r:   rO   rP   rQ   r%   r&   r'   r)   r)   A   s|   	 OKJ NN 05/315;? '}' #+5<<"8 EIellIZD[..
# .
#+E?.
&.u||&<.
AEc5<<FWAX.
`,
d3+<&= ,
\ e  S  C  TW  .3 3e 3-2\\3*-327,,3j$ell $E $%*\\$;>$,/$49LL$r&   r)   c                   *    e Zd ZdZddZd Zd Zd Zy)NeutrinoDatasetz'Dataset for neutrino detection trainingc                 |    || _         || _        || _        t        |      | _        g | _        | j                          y r-   )num_samplessignal_ratior.   r)   	simulatordata_generate_data)r1   r   r   r.   s       r'   r2   zNeutrinoDataset.__init__E  s:    &(*62 	r&   c                    t        | j                  | j                  z        }| j                  |z
  }t        |      D ]9  }| j                  j                  d      }| j                  j                  |       ; t        |      D ]  }t        j                  ddd      j                         }dt        j                  d      j                         dz  z  }t        |d      }t        j                  d      j                         t        j                  z  }t        j                  d      j                         d	z  t        j                  z  }t        j                  t        j                   |      t        j"                  |      z  t        j                   |      t        j                   |      z  t        j"                  |      g      }	| j                  j                  d
|||	      }| j                  j                  |        t        j$                  t'        | j                              j)                         }
|
D cg c]  }| j                  |    c}| _        yc c}w )z2Generate dataset with realistic event distributionF)r3   r   r>   r?   r@   r*   r   r+   T)r3   r4   r5   r6   N)r   r   r   r|   r   r<   r   r   rH   rI   rJ   rK   rL   r"   r$   rM   rN   r#   randpermlentolist)r1   num_signalsnum_background_r;   r4   r5   rR   rS   r6   indicesr   s               r'   r   zNeutrinoDataset._generate_dataO  s    $**T->->>?))K7 ~& 	$ANN11E1BEIIU#	$
 {# 	$A]]1a.335FEJJqM..0145F'FJJqM&&(4772E**Q-$$&*TWW4C$((3-/$((3-/& I NN11# 2 E IIU#)	$. ..TYY0779+23aTYYq\3	3s   I c                     | j                   S r-   )r   r1   s    r'   __len__zNeutrinoDataset.__len__u  s    r&   c                     | j                   |   }d}|d   }|j                  d   |k  r@t        j                  ||j                  d   z
  d      }t        j                  ||gd      }n|d | }||d   |d   dS )	Nru   rD   r   r\   rY   rE   rF   )rD   rE   rF   )r   r   rH   ra   cat)r1   idxr;   max_hitsrD   paddings         r'   __getitem__zNeutrinoDataset.__getitem__x  s    		# V}::a=8#kk(TZZ]":A>G99dG_!4D	?D j)7^
 	
r&   N)'  r
   r   )r   r   r   r   r2   r   r   r   r%   r&   r'   r   r   B  s    1$4L 
r&   r   c                   4     e Zd ZdZddedef fdZd Z xZS )PositionalEncodingz#Positional encoding for transformerd_modelmax_lenc                 &   t         |           t        j                  ||      }t        j                  d|t        j
                        j                  d      }t        j                  t        j                  d|d      j                         t        j                  d       |z  z        }t        j                  ||z        |d d dd df<   t        j                  ||z        |d d dd df<   | j                  d|j                  d             y )Nr   rx   r*   r+   g     @pe)superr2   rH   ra   aranger   	unsqueezeexpr"   logrN   r#   register_buffer)r1   r   r   r   positiondiv_term	__class__s         r'   r2   zPositionalEncoding.__init__  s    [['*<<7%++>HHK99U\\!Wa8>>@TXXgEVDVY`D`abii8 341add7ii8 341add7T2<<?3r&   c                 P    || j                   d d d |j                  d      f   z   S )Nr*   )r   sizer1   rb   s     r'   forwardzPositionalEncoding.forward  s&    4771jqvvayj=)))r&   )i  )r   r   r   r   r   r2   r   __classcell__r   s   @r'   r   r     s    -4 4c 4*r&   r   c                   :     e Zd ZdZddededef fdZddZ xZS )	MultiHeadAttentionzg
    Multi-head attention with sensitivity tracking.
    Implements the CCT-ODE collapse operator.
    r   n_headsdropoutc                    t         |           ||z  dk(  sJ || _        || _        ||z  | _        t        j                  ||      | _        t        j                  ||      | _        t        j                  ||      | _	        t        j                  ||      | _
        t        j                  |      | _        t        j                  |      | _        g | _        g | _        y )Nr   )r   r2   r   r   d_knnLinearW_qW_kW_vW_oDropoutr   	LayerNorm
layer_normsensitivity_historycollapse_history)r1   r   r   r   r   s       r'   r2   zMultiHeadAttention.__init__  s     A%%%g%99Wg.99Wg.99Wg.99Wg.zz'*,,w/ $&  "r&   c                 &   |j                  d      }|j                  d      }| j                  |      }| j                  |      }	| j                  |      }
|j	                  ||| j
                  | j                        j                  dd      }|	j	                  |d| j
                  | j                        j                  dd      }	|
j	                  |d| j
                  | j                        j                  dd      }
t        j                  ||	j                  dd            t        j                  | j                        z  }||j                  |dk(  d      }t        j                  |d      }| j                  |      }|r:|j!                         }| j"                  j%                  |j'                                nd}t        j                  ||
      }|j                  dd      j)                         j	                  ||| j*                        }| j-                  |      }|rj|ht        j.                  |t        j0                  |d	z         z  d      j3                          }| j4                  j%                  |j7                                ||fS )
z
        Forward pass with optional sensitivity tracking.
        
        Returns:
            - output: [batch, seq, d_model]
            - sensitivity: [batch, n_heads, seq, seq] (optional)
        r   r*   r+   r]   Ng    erY   绽|=)r   r   r   r   viewr   r   	transposerH   matmulr"   sqrtmasked_fillFsoftmaxr   cloner   r   detach
contiguousr   r   r   r   r   r   rJ   )r1   querykeyvaluemaskreturn_sensitivity
batch_sizeseq_lenQKVscoresattention_weightssensitivitycontextoutputentropy_befores                    r'   r   zMultiHeadAttention.forward  s!    ZZ]
**Q- HHUOHHSMHHUO FF:wdhh?II!QOFF:r4<<:DDQJFF:r4<<:DDQJ aR!45		$((8KK''	48F IIf"5 LL):; +113K$$++K,>,>,@AK ,,0!4 ##Aq)446;;JQUQ]Q]^'" "3"?#ii(9EIIFWZ_F_<`(`fhinnppN!!(()<)<)>?{""r&   r
   NF	r   r   r   r   r   r   r2   r   r   r   s   @r'   r   r     s'    
# #c #E #(3#r&   r   c                   8     e Zd ZdZddededef fdZd Z xZS )CCTODEFeedForwardz&Feed-forward network with ODE dynamicsr   d_ffr   c                     t         |           t        j                  ||      | _        t        j                  ||      | _        t        j                  |      | _        t        j                         | _	        y r-   )
r   r2   r   r   linear1linear2r   r   GELU
activation)r1   r   r  r   r   s       r'   r2   zCCTODEFeedForward.__init__  sO    yy$/yyw/zz'*'')r&   c           	      ~    | j                  | j                  | j                  | j                  |                        S r-   )r  r   r  r
  r   s     r'   r   zCCTODEFeedForward.forward  s+    ||DLLa)IJKKr&   i   r
   r  r   s   @r'   r  r    s&    0$ $3 $ $Lr&   r  c            	       >     e Zd ZdZddedededef fdZd	dZ xZS )
CCTODELayerzf
    Single CCT-ODE layer combining attention and FFN.
    Tracks sensitivity matrices per layer.
    r   r   r  r   c                 R   t         |           t        |||      | _        t	        |||      | _        t        j                  |      | _        t        j                  |      | _	        t        j                  |      | _        t        j                  |      | _        d d d d d| _        y )N)r5   r   r6   r4   )r   r2   r   	attentionr  feedforwardr   r   norm1norm2r   dropout1dropout2layer_sensitivity)r1   r   r   r  r   r   s        r'   r2   zCCTODELayer.__init__  s    +GWgF,WdGD\\'*
\\'*


7+

7+ 	"
r&   c                 \   | j                  | j                  |      | j                  |      | j                  |      ||      \  }}|| j                  |      z   }|r|| j                  d<   | j	                  | j                  |            }|| j                  |      z   }|| j                  fS )aJ  
        Forward pass with optional sensitivity tracking.
        
        Args:
            x: [batch, seq, d_model]
            track_sensitivity: Whether to store sensitivity matrices
            
        Returns:
            - output: [batch, seq, d_model]
            - sensitivity_dict: dict of sensitivity matrices
        )r   r  )r  r  r  r  r  r  r  )r1   rb   r   track_sensitivityattn_outr   ff_outs          r'   r   zCCTODELayer.forward  s     !%JJqM4::a=$**Q-%6 !/ !
+ h'' 2=D"";/ !!$**Q-0f%%$((((r&   r  r  r  r   s   @r'   r  r    s.    

 
c 
 
e 
()r&   r  c                   p     e Zd ZdZ	 	 	 	 	 	 	 	 ddededededededed	ef fd
ZddZd Zd Z	d Z
 xZS )NeutrinoTransformera*  
    CCT-ODE Transformer for Neutrino Detection.
    
    Implements:
    - Hit-based processing (Transformer on hit coordinates)
    - Feature-based processing (MLP on physics features)
    - Sensitivity matrix tracking
    - Entropic collapse detection
    - Hypersensitivity for rare events
    r   r   n_layersr  r   
n_features	n_classesr   c	                    t         
|           || _        || _        || _        || _        || _        t        j                  t        j                  d|dz        t        j                  |dz        t        j                         t        j                  |dz  |      t        j                  |      t        j                               | _        t        ||      | _        t        j                  t!        |      D 	cg c]  }	t#        ||||       c}	      | _        t        j                  t        j                  ||      t        j                  |      t        j                         t        j                  ||      t        j                  |      t        j                               | _        t        j(                  |||d      | _        t        j                  |      | _        g g g g d| _        t        j                  t        j                  |dz  |      t        j                  |      t        j                         t        j0                  |      t        j                  ||dz        t        j                  |dz        t        j                         t        j0                  |      t        j                  |dz  |      	      | _        t        j                  t        j                  ||dz        t        j                  |dz        t        j                         t        j                  |dz  d      t        j4                               | _        y c c}	w )Nr\   r+   T)batch_firsthit_attentionfeature_attentioncross_attentionlayer_entropiesr*   )r   r2   r   r   r   r   r"  r   
Sequentialr   r   r  hit_embeddingr   hit_positional_encoding
ModuleListr|   r  hit_transformer_layersfeature_encoderMultiheadAttentionfusionfusion_normsensitivity_matricesr   
classifierSigmoidrare_event_detector)r1   r   r   r   r  r   r!  r"  r   r   r   s             r'   r2   zNeutrinoTransformer.__init__F  sm    	  "  ]]IIaA&LLA&GGIIIglG,LL!GGI
 (:'8'L$&(mm8_5
 $85
 '#  "}}IIj'*LL!GGIIIgw'LL!GGI 
 ++Wg4
 <<0  !#!!	%
! --IIgk7+LL!GGIJJwIIgw!|,LLA&GGIJJwIIglI.

  $&==IIgw!|,LLA&GGIIIglA&JJL$
 w5
s   :Mc                 ,   |j                   d   }|j                         j                  d      dkD  j                         }| j	                  |      }| j                  |      }|}t        | j                        D ]u  \  }}	 |	||j                  d      j                  d      |      \  }}
|s5|
j                  d      G| j                  d
   j                  |
d   j                                w | j                  ||      }| j                  |      }| j                  |j                  d      |||dk(        \  }}|j!                  d      }| j#                  ||z         }|r,| j                  d   j                  |j                                t%        j&                  ||gd      }| j)                  |      }| j+                  |      }|r.| j-                         }| j                  d   j                  |       ||||||r| j                  dS d	dS )a  
        Forward pass of the neutrino transformer.
        
        Args:
            hits: [batch, max_hits, 4] - Hit coordinates
            features: [batch, n_features] - Physics features
            track_sensitivity: Whether to compute sensitivity matrices
            
        Returns:
            - logits: [batch, n_classes]
            - sensitivity_dict: dict of sensitivity matrices
            - rare_event_prob: [batch, 1]
        r   r]   rY   rh   r*   r+   )r   r  r  Nr&  )r   r   r   key_padding_maskr(  r)  )logitshit_representationfeature_representationfused_representationrare_event_probr3  )r   absr   r   r+  r,  	enumerater.  r   getr3  r   r   _attention_poolr/  r1  squeezer2  rH   r   r4  r6  _compute_entropy)r1   rD   rE   r  r   hit_maskhits_embeddedr:  	layer_idxlayerr  
hit_pooledr;  fused
cross_attnjoint_representationr9  r=  total_entropys                      r'   r   zNeutrinoTransformer.forward  s6    ZZ]
 HHJNNrN*T188: **4044]C + )$*E*E F 
	Iu49"''*44Q7"351 1 !%6%:%:;%G%S))/:AA%k299;
	 ))*<hG
 "&!5!5h!? !KK(2215"$&!m	 ( 
z a   )?!?@%%&78??
@Q@Q@ST  %yy*e)<"E!56 2259  113M%%&78??N ",&<$).ARD$=$=
 	
 Y]
 	
r&   c                     |j                  d      j                         }||j                  dd      dz   z  }||z  j                  d      }|S )a   
        Attention-weighted pooling over sequence dimension.
        
        Args:
            x: [batch, seq, d_model]
            mask: [batch, seq] - 1 for valid, 0 for padding
            
        Returns:
            pooled: [batch, d_model]
        r]   r*   T)rZ   keepdimrh   rY   )r   r   r   )r1   rb   r   weightspooleds        r'   rA  z#NeutrinoTransformer._attention_pool  sS     ..$**,W[[Q[=DEg+""q")r&   c           	         d}| j                   d   fD ]r  }|D ]k  }t        j                  |d      }t        j                  |t        j
                  |dz         z  d       }||j                         j                         z  }m t |S )zDCompute entropy of attention distributions (CCT-ODE collapse metric)rg   r&  r]   rY   r   )r3  r   r   rH   r   r   r   rJ   )r1   rL  sensitivity_listr   probentropys         r'   rC  z$NeutrinoTransformer._compute_entropy!  s     !%!:!:?!K L 	7/ 7yy"5 99TEIIdUl,C%CLL!4!4!66	7	7 r&   c                     g g g g d| _         y)z(Reset sensitivity matrices for new batchr%  N)r3  r   s    r'   reset_sensitivity_trackingz.NeutrinoTransformer.reset_sensitivity_tracking/  s      !#!!	%
!r&   )   r         ru   r[   r+   r
   )T)r   r   r   r   r   r   r2   r   rA  rC  rV  r   r   s   @r'   r  r  :  s    	 d
d
 d
 	d

 d
 d
 d
 d
 d
La
F$
r&   r  c                   8     e Zd ZdZd fd	ZddZd Zd Z xZS )
CCTODELossz
    Conditional Collapse Theory loss function.
    
    Combines:
    - Cross-entropy for classification
    - Sensitivity amplification for rare events
    - Fermi-Dirac regularization for particle statistics
    - Entropy collapse bonus
    c                 ~    t         |           || _        || _        || _        t        j                         | _        y r-   )r   r2   rare_event_weightsensitivity_weight
fermi_tempr   CrossEntropyLossce_loss)r1   r]  r^  r_  r   s       r'   r2   zCCTODELoss.__init__H  s5    !2"4$**,r&   c                    | j                  |d   |      }|dk(  j                         }| j                   |d   j                         |z  j	                         z  }d}|| j                  |      }| j                  |      }||z   | j                  |z  z   d|z  z   }	|j                         |j                         t        |t        j                        r|j                         n||j                         |	j                         d}
|	|
fS )ae  
        Compute CCT-ODE loss.
        
        Args:
            outputs: dict from neutrino transformer
            targets: [batch] tensor of class labels
            sensitivity_matrices: optional dict of sensitivity matrices
            
        Returns:
            total_loss: scalar
            loss_dict: dict of individual loss components
        r9  r*   r=  rg   ri   )cross_entropyrare_event_bonussensitivity_lossfermi_regularizationtotal)ra  r   r]  rB  r   _compute_sensitivity_loss_compute_fermi_regularizationr^  rJ   
isinstancerH   r   )r1   outputstargetsr3  cesignal_mask
rare_bonusre  
fermi_loss
total_loss	loss_dicts              r'   r   zCCTODELoss.forwardO  s    \\'(+W5 !|**,,,,%&..0;>
$&

 +#==>RS 77@
*_t'>'>AQ'QQTX[eTee
  WWY * 1;EFVX]XdXd;e 0 5 5 7k{$.OO$5__&
	 9$$r&   c                     |j                  d      syd}d}|d   D ]e  }t        j                  |d      }t        j                  |t        j
                  |dz         z  d       }|j                          }||z  }|dz  }g |dkD  r| |z  S y)z)Encourage high sensitivity to rare eventsr&  rg   r   r]   rY   r   r*   )r@  r   r   rH   r   r   r   )r1   r3  total_concentrationcountr   rS  rT  concentrations           r'   rh  z$CCTODELoss._compute_sensitivity_lossy  s     $''8 "/@ 
	K99[b1D yy		$,(?!?RHHG %\\^OM=0QJE
	 19''%//r&   c                     |d   }t        j                  |d      }t        j                  |      \  }}|dd |dd z
  }t        j                  |dd |dd z
        j                         }|S )z
        Fermi-Dirac inspired regularization.
        
        Encourages weight distribution to follow quantum statistics.
        r<  r]   rY   r*   N)rH   r}   sortr   relur   )r1   rk  	fused_repr5   energy_sortedr   energy_diffrp  s           r'   ri  z(CCTODELoss._compute_fermi_regularization  s     23	 I2. !::f-q#AB'-*<< VVK,{12>?DDF
r&   )      @r
   r   r-   )	r   r   r   r   r2   r   rh  ri  r   r   s   @r'   r[  r[  =  s    -(%T4r&   r[  c                       e Zd ZdZddZd Zy)AdaptiveSensitivityCallbackz
    Callback for adaptive sensitivity adjustment during training.
    
    Monitors entropy evolution and adjusts training parameters accordingly.
    c                 <    || _         || _        g | _        d| _        y )Nnormal)modelsensitivity_targetentropy_historyphase)r1   r  r  s      r'   r2   z$AdaptiveSensitivityCallback.__init__  s     
"4!
r&   c                    |r;d|v r7|d   j                  dg       }|r |d   }| j                  j                  |       t        | j                        dkD  rjt	        j
                  | j                  dd       }|| j                  dz  kD  rd| _        d	d
dS || j                  dz  k  rd| _        dddS d| _        dddS dddS )z Called after each training batchr3  r)  r]   r@   iNr+   	sensitiveg       @g      ?)sensitivity_scalelearning_rate_factorrW   periodicr  r   )r@  r  r   r   rj   r   r  r  )r1   	batch_idxrr  rk  	entropiescurrent_entropy
recent_avgs          r'   on_batch_endz(AdaptiveSensitivityCallback.on_batch_end  s     -8 67;;<MrRI"+B-$$++O< t##$r)!5!5cd!;<JD33a77(
-0#NNd55;;'
-0#NN%
-0#NN%(#FFr&   Nr  )r   r   r   r   r2   r  r%   r&   r'   r  r    s    Gr&   r  r  train_loader
val_loaderepochsr.   lrweight_decayr]  c                 
	   | j                  |      } t        |      }t        j                  | j	                         ||      }	t        j
                  j                  |	|      }
t        |       }g g g g g g d}t        d      }t        |      D ]  }| j                          g }d}d}t        |      D ]U  \  }}|d   j                  |      }|d   j                  |      }|d	   j                  |      }|	j                           | ||d
      } ||||d         \  }}|j                          t        j                  j                   j#                  | j	                         d       |	j%                          |j'                  |||      }|j)                  |j+                                |d   j-                  d      }|||k(  j/                         j+                         z  }||j1                  d      z  }| j3                          X |
j%                          t5        j6                  |      }||z  }| j9                          g }d}d} t        j:                         5  |D ]  }|d   j                  |      }|d   j                  |      }|d	   j                  |      } | ||d      } |||d      \  }}!|j)                  |j+                                |d   j-                  d      }|||k(  j/                         j+                         z  }| |j1                  d      z  }  	 ddd       t5        j6                  |      }"|| z  }#|d   j)                  |       |d   j)                  |"       |d   j)                  |       |d   j)                  |#       |d   j)                  |j<                         t?        d|dz    d|        t?        d|dd|d       t?        d|"dd|#d       t?        d |j<                          t?        d!|
jA                         d   d"       t?                |"|k  s|"}t        jB                  | jE                         d#       t?        d$|"dd%        |S # 1 sw Y   PxY w)&a  
    Train the neutrino detection model using CCT-ODE framework.
    
    Args:
        model: NeutrinoTransformer model
        train_loader: Training data loader
        val_loader: Validation data loader
        epochs: Number of training epochs
        device: 'cuda' or 'cpu'
        lr: Learning rate
        weight_decay: Weight decay for AdamW
        rare_event_weight: Weight for rare event detection loss
        
    Returns:
        history: dict of training metrics
    )r]  )r  r  )T_max)
train_lossval_loss	train_accval_accr   rare_event_aucinfr   rD   rE   rF   Tr  r3  r   )max_normr9  r*   rY   FNr  r  r  r  r   zEpoch /z  Train Loss: .4fz | Train Acc: z  Val Loss: z | Val Acc: z  Sensitivity Phase: z  LR: z.6fbest_neutrino_cct_model.ptu"     ✓ Saved best model (val_loss: ))#tor[  optimAdamW
parameterslr_schedulerCosineAnnealingLRr  r   r|   trainr?  	zero_gradbackwardrH   r   utilsclip_grad_norm_stepr  r   rJ   argmaxr   r   rV  rj   r   evalno_gradr  printget_last_lrsave
state_dict)$r  r  r  r  r.   r  r  r]  	criterion	optimizer	schedulersensitivity_callbackhistorybest_val_lossepochepoch_lossesepoch_correctepoch_totalr  batchrD   rE   labelsrk  lossrr  adjustpredsr  r  
val_lossesval_correct	val_totalr   r  r  s$                                       r'   train_neutrino_cctr    s}   4 HHVE->?IE,,.2LQI ""44Yf4MI 7u= G %LMv WH ), 7 	/Iu=##F+DZ(++F3H7^&&v.F !D(dCG (AW9XYOD) MMO HHNN**5+;+;+=*LNN *66y)WUF 		,H%,,,3Eevo22499;;M6;;q>)K ,,.?	/B 	 WW\*
!K/	 	


	]]_ 	,# ,V}''/ ,//7w**62h%H#GVT:a!!$))+.)00Q07446;;==V[[^+	,	, 77:&	) 	$$Z0
""8,##I.	!!'*%%&:&@&@A 	uQwiq)*z#.nYsOLMXcN,wsmDE%&:&@&@%ABCy,,.q1#678 m#$MJJu'')+GH6xnAFGoWHr NM	, 	,s   !CQ88R	c                 	   | j                          | j                  |      } g }g }g }g }g }t        j                         5  |D ]7  }|d   j                  |      }	|d   j                  |      }
|d   j                  |      } | |	|
d      }|d   j	                  d      }|d	   j                         }|j                  |j                         j                                |j                  |j                         j                                |j                  |j                         j                                |j                  |	j                                |j                  |
j                                : 	 d
d
d
       t        j                  |      }t        j                  |      }t        j                  |      }||k(  j                         }|dk(  |dk(  z  j                         }|dk(  |dk(  z  j                         }|dk(  |dk(  z  j                         }|dk(  |dk(  z  j                         }|||z   dz   z  }|||z   dz   z  }d|z  |z  ||z   dz   z  }	 ddlm}m}  |||      } |||      \  }}}t#        d       t#        d       t#        d       t#        d       t#        d|d       t#        d|d       t#        d|d       t#        d|d       t#        d|d       t#        d       t#        d| d|        t#        d| d|        t#        d       t#        d |j                                 t#        d!d|z
  j                                 t#        d"|        t#        d#|        t#        d$       d%D ]  }|d&k(  r*t        j                  |D cg c]
  }|d   d'k   c}      }nb|d(k(  r4t        j                  |D cg c]  }|d   d'k\  xr |d   dk   c}      }n)t        j                  |D cg c]
  }|d   dk\   c}      }|j                         dkD  s||   ||   k(  j                         } t#        d)|j%                          d*| dd+|j                          d,        t#        d       |||||||||d-|||||d.S # 1 sw Y   ,xY w# t         $ r d}d\  }}}Y Gw xY wc c}w c c}w c c}w )/z
    Comprehensive evaluation of neutrino detector.
    
    Computes:
    - Accuracy, Precision, Recall, F1
    - ROC-AUC (important for rare event detection)
    - Confusion matrix
    - Sensitivity analysis per event type
    rD   rE   rF   Fr  r9  r*   rY   r=  Nr   rh   r+   )roc_auc_score	roc_curverW   )NNN<============================================================z$NEUTRINO DETECTOR EVALUATION RESULTSz
Overall Metrics:z  Accuracy: r  z  Precision: z
  Recall: z  F1 Score: z  ROC-AUC (Rare Event): z
Confusion Matrix:z  True Positives: z | False Negatives: z  False Positives: z | True Negatives: z
Rare Event Detection:z  Signal events: z  Background events: z  Detected signals: z  Missed signals: z
Per-Energy Analysis:)lowmediumhighr  g      r    z energy: Accuracy = z (n=r  )tptnfpfn)accuracy	precisionrecallf1aucconfusion_matrixpredictionsr  
rare_probsfprtpr)r  r  rH   r  r  rB  extendr   r   rj   arrayr   r   sklearn.metricsr  r  ImportErrorr  
capitalize)!r  test_loaderr.   	all_preds
all_labelsall_rare_probsall_hitsall_featuresr  rD   rE   r  rk  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  
thresholds
energy_binfr   bin_accs!                                    r'   evaluate_neutrino_detectorr  f  s    
JJLHHVEIJNHL	 0  	0E=##F+DZ(++F3H7^&&v.FD(eDGH%,,,3E !23;;=JUYY[//12fjjl1134!!*.."2"9"9";<OODHHJ'/	00$ #I*%JXXn-N Z'--/H >jAo
.	3	3	5B>jAo
.	3	3	5B>jAo
.	3	3	5B>jAo
.	3	3	5Bb2gn%I27T>"F	
Y	9v#5#<	=B0<J7(^DS*
 
(O	
01	(O	 	L#
'(	M)C
)*	Jvcl
#$	LC
!"	$SI
./	!	rd"6rd
;<	t#6rd
;<	#%	jnn./
01	!1z>"6"6"8!9
:;	 
%&	rd
#$ 
"$/ 
d
88,?QQqTD[?@D8#88,OQQqTT\8adSj8OPD88,?QQqTS[?@D88:> *T*::@@BGBz,,.//CGC=PTUYU]U]U_T``abc
d 
(O #%RrD $ a0 0P  0/S*08 @O?s1    D>R)%R6 1S
 S
S
)R36S
Sc                 J   | j                          | j                  |      } g }t        j                         5  |D ]r  }|d   j                  |      }|d   j                  |      }|d   j                  |      } | ||d      }|d   r|j	                  |d          | j                          t 	 ddd       t        d       t        d	       t        d       t        d
       t        d       |r|D 	cg c]  }	d|	v s|	d    }
}	|
rwt        d       t        dt        |
              t        dt        j                  |
D cg c]'  }|D ]   }|j                         j                         " ) c}}      d       t        d       y# 1 sw Y   xY wc c}	w c c}}w )z
    Analyze the sensitivity matrices learned by the model.
    
    This reveals which aspects of the neutrino detection the model
    is most sensitive to, validating the CCT-ODE framework.
    rD   rE   rF   Tr  r3  Nr  zSENSITIVITY MATRIX ANALYSISz(
Analyzing layer-by-layer sensitivity...z<Higher sensitivity = Model is more responsive to that aspectr&  z
Hit Attention Sensitivities:z  Number of layers analyzed: z  Average attention entropy: r  )r  r  rH   r  r   rV  r  r   rj   r   rJ   )r  data_loaderr.   all_sensitivitiesr  rD   rE   r  rk  shit_attneses                r'   analyze_sensitivity_matricesr    s    
JJLHHVE	 /  	/E=##F+DZ(++F3H7^&&v.FD(dCG -.!((1G)HI,,.	// 
(O	
'(	(O	
56	
HI 0AZ1_XYEYAo&ZZ241#h-AB1"''U]:jrgi:jbc1668==?:j?:j2klo1pqr	(O=/ /. [
 ;ks   A8F6	F F,FFc                  r   i ddddddddd	d
dddddddddddddddddddddddd d!t         j                  j                         rd"nd#i} t        d$       t        d%       t        d&       t        d$       t        d'       | j	                         D ]  \  }}t        d(| d)|         t        d*| d!           t        d+       t        | d   | d   | d!   ,      }t        | d   | d   | d!   ,      }t        | d   | d   | d!   ,      }t        || d   d-d./      }t        || d   d0d./      }t        || d   d0d./      }t        d1t        |       d2       t        d3t        |       d2       t        d4t        |       d2       t        d5| d   d6z  d7d8       t        d9       t        | d   | d   | d   | d   | d	   | d   | d   | d   :      }	t        d; |	j                         D              }
t        d<|
d=       t        d>       t        d?       t        |	||| d   | d!   | d   | d   | d   @      }t        d?       t        dA       |	j                  t        j                  dB             t        dC       t        dD       t        |	|| d!   E      }t        dF       t        |	|| d!   E       t        dG       t        dH       t        d$       t        dIt!        |dJ         dK       t        dL|dM   dK       t        dN|dO   dK       t        dP       |	||fS )QzD
    Main function to train and evaluate the neutrino detector.
    r   rW  r   r   r   rX  r  rY  r   ru   r!  r[   r"  r+   r   r
   train_samplesiP  val_samplesr   test_samplesr   r  rv   r   @   r  -C6?r  ri   r]  r}  r.   cudar   r  z@MMA-CCT: Multi-Messenger Astronomy - Conditional Collapse TheoryzNeutrino Detection Systemz
Configuration:r  z: z	
Device: z
[1/6] Generating datasets...)r   r   r.   Tr   )r   shufflenum_workersFz	  Train: z samplesz  Val: z  Test: z  Signal ratio: d   z.1f%z
[2/6] Creating model...)r   r   r   r  r   r!  r"  r   c              3   V   K   | ]!  }|j                   s|j                          # y wr-   )requires_gradnumel).0ps     r'   	<genexpr>zmain.<locals>.<genexpr>Z  s     N1aooQWWYNs   ))z  Model parameters: ,z
[3/6] Training model...z<------------------------------------------------------------)r  r  r  r  r.   r  r  r]  z
[4/6] Loading best model...r  z  Best model loaded.z
[5/6] Evaluating model...)r.   z(
[6/6] Analyzing sensitivity matrices...z=
============================================================zTRAINING COMPLETEzBest validation loss: r  r  zFinal test accuracy: r  zFinal test AUC: r  z+
Model saved to: best_neutrino_cct_model.pt)rH   r  is_availabler  itemsr   r   r   r  r   r  r  load_state_dictloadr  r  rL   )CONFIGr   r   train_datasetval_datasettest_datasetr  r  r  r  
num_paramsr  resultss                r'   mainr  	  s5   31 	A 		
 	C 	b 	Q 	3 	 	u 	 	 	"  	b!" 	d#$ 	%& 	S'* 	EJJ335&5+F0 
(O	
LM	
%&	(O	
lln #
U3%r%!"# 
Jvh'(
)* 

*+#?+N+hM
 "=)N+hK
 #>*N+hL m|8LVZhijLKF<4HRWefgJ\f\6JTYghiK	Ic-()
23	GC$%X
./	HS&'x
01	VN3C7<A
>? 

%&y!y!
#F^
#,'%y!	E N(8(8(:NNJ	 A
/0 

%&	(O !hh$<N+ !45	G 
(O 

)*	%**%ABC	
 ! 

'((F8DTUG 

56 F8<LM	/	
	(O	"3wz':#;C"@
AB	!'*"5c!:
;<	WU^C0
12	
89'7""r&   __main__)rV   r  r  ri   r}  )r  )8r   rH   torch.nnr   torch.nn.functional
functionalr   torch.optimr  torch.utils.datar   r   r"   numpyrj   typingr   r   r   r   dataclassesr	   warningsr   EV_TO_JOULESr    r!   NEUTRINO_ENERGY_RANGENEUTRINO_FLUX_TOTALBACKGROUND_RATESIGNAL_BRANCHING_RATIODETECTOR_VOLUMEPHOTON_SPEEDr   r/   r)   r   Moduler   r   r  r  r  r[  r  r   r   r   r  r  r  r  r   r  r  r  r%   r&   r'   <module>r&     s        0   . . !  		 %    W}
	$ 	$ 	$ 
~ ~BG
g G
\* * M# M#`L		 L5)")) 5)x|
")) |
Fi i`'G 'G\ "H99HH H 	H
 H 	H H H^od*bx#v z"fE7G r&   