
    iv                        d dl Z d dlmZ d dlmZ d dlmZ d dlmZ d dl	m
Z
 d dlmZ d dlmZmZ d dl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Z"d'dZ#d(dZ$d Z%d)dZ&d*dZ'd Z(d+dZ)d Z*e+dk(  rp ejX                  d      Z-e-j]                  de/dd !       e-j]                  d"e/dd#!       e-j]                  d$e0dd%!       e-jc                         Z2 e*e2      \  Z3Z4Z5Z6yy# e$ r" d dlZ ej8                  g d       d dlmZ Y w xY w),    N)
DataLoader)FuncAnimationPillowWriter)Axes3D)Rotation)odeint)pipinstalltorchdiffeqz-qc                   0     e Zd ZdZd fd	Zd Zd Z xZS )QuantumVertexGeneratoru   
    Converts batch of images to 100×12 3D vertex points.
    
    Each sample → 12 learnable vertices in 3D space.
    The vertices represent "quantum states" that encode spatial information.
    c                    t         |           || _        || _        || _        t        j                  t        j                  |d      dz        | _	        t        j                  t        j                  dddd      t        j                  d      t        j                         t        j                  d      t        j                  dddd      t        j                  d      t        j                         t        j                  d      t        j                  dddd      t        j                  d      t        j                         t        j                  d	      t        j                          t        j"                  d
d      t        j                               | _        t        j"                  d|dz        | _        t        j"                  d|dz        | _        t        j"                  dd      | _        y )N         ?       padding   @      )r   r   i        )super__init__img_sizenum_vertices
batch_sizenn	Parametertorchrandnbase_vertices
SequentialConv2dBatchNorm2dReLU	MaxPool2dAdaptiveAvgPool2dFlattenLinearfeature_netrotation_predictorscale_predictortranslate_predictor)selfr   r   r   	__class__s       1/home/per/Documents/python/circle_diff_eq/ex04.pyr   zQuantumVertexGenerator.__init__   sX    ($  \\%++lA*F*LM ==IIaQ*NN2GGILLOIIb"a+NN2GGILLOIIb#q!,NN3GGI  (JJLIIgs#GGI
& #%))C1A"B!yylQ.>?#%99S!#4     c                    |j                   d   }| j                  |      }| j                  |      }t        j                  | j                  |            }t        j                  | j                  |            }g }t        |      D ]  }g }	t        | j                        D ]v  }
| j                  |
   }|||
dz  |
dz   dz  f   }||j                         dz   z  }| j                  ||      }|||
dz  |
dz   dz  f   }||dz   z  }|	j                  |       x t        j                  |	d      }	|	||   z   }	|j                  |	        t        j                  |d      }|S )z
        Args:
            x: (batch_size, 1, 28, 28) images
        Returns:
            vertices: (batch_size, num_vertices, 3) 3D vertex points
        r   r   r   g:0yE>r   r   dim)shaper,   r-   r!   sigmoidr.   tanhr/   ranger   r#   norm_rotate_by_quaternionappendstack)r0   xr   featuresquaternionsscalestranslationvertices_listbsample_verticesvv_baseqv_rotscale_vv_scaledverticess                    r2   forwardzQuantumVertexGenerator.forwardB   sz    WWQZ
 ##A& --h7t33H=>jj!9!9(!CD z" 	2A O4,,- 1++A.  1Q3!Qw;/D) 2261= !AaC1aK0 GcM2&&x01" $kk/qAO-A>O  1-	20 ;;}!4r3   c                    |d   |d   |d   |d   f\  }}}}|d   |d   |d   }	}}d||z  ||z  z   dz
  |z  ||z  ||z  z
  |z  z   ||z  ||z  z   |	z  z   z  }
d||z  ||z  z   |z  ||z  ||z  z   dz
  |z  z   ||z  ||z  z
  |	z  z   z  }d||z  ||z  z
  |z  ||z  ||z  z   |z  z   ||z  ||z  z   dz
  |	z  z   z  }t        j                  |
||g      S )z Rotate vector v by quaternion q.r   r   r   r   r   )r!   r>   )r0   rG   rI   wr?   yzvxvyvzresult_xresult_yresult_zs                r2   r<   z,QuantumVertexGenerator._rotate_by_quaternionq   s6    qT1Q41qt+
1a qT1Q41B!ac	C+qsQqSy"n<!ac	2~MO!ac	2~1qsS"(<<!ac	2~MO!ac	2~1qsB6!A#!)c/29MMO{{Hh9::r3   )      d   )__name__
__module____qualname____doc__r   rN   r<   __classcell__r1   s   @r2   r   r      s    !5F-^;r3   r   c                   *     e Zd ZdZd fd	Zd Z xZS )QuantumVertexODEFuncu   
    ODE function that evolves 100×12 vertex points through time.
    
    The dynamics model quantum-like interactions between vertices.
    c                    t         |           || _        || _        ||z  dz   | _        t        j                  t        j                  | j                  |      t        j                  |      t        j                         t        j                  d      t        j                  ||      t        j                  |      t        j                         t        j                  |||z              | _        t        j                  |dd      | _        t        j                  t        j                  |dz  dz   d      t        j                         t        j                  d|            | _        y )Nr   皙?r   T)	embed_dim	num_headsbatch_firstr   r   )r   r   
vertex_dimr   	input_dimr   r$   r+   	LayerNormSiLUDropoutdynamics_netMultiheadAttention	attentionr'   edge_mlp)r0   ri   r   
hidden_dimr1   s       r2   r   zQuantumVertexODEFunc.__init__   s   $( &
2Q6 MMIIdnnj1LL$GGIJJsOIIj*-LL$GGIIIj,";<	
 ..q^bc IIj1nq("-GGIIIb*%
r3   c                    |j                   d   }|j                  || j                  | j                        }t	        j
                  ||      }| j                  |||      \  }}t	        j                  |j                  d      |j                  d      |j                  |d      gd      }| j                  |      }	|	|j                  d      dz  z   }	|	S )z
        Args:
            t: current time (scalar)
            vertices_flat: (batch, num_vertices * vertex_dim) flattened vertex positions
        Returns:
            d_vertices/dt: (batch, num_vertices * vertex_dim)
        r   r   r5   333333?)r7   viewr   ri   r!   cdistrp   catflattenexpandrn   )
r0   tvertices_flatr   rM   pairwise_distattn_out_combineddynamicss
             r2   rN   zQuantumVertexODEFunc.forward   s     #((+
 !%%j$2C2CT__U Hh7 nnXxB! 99QQHHZ#
 	 $$X. h..q1C77r3   )r   rZ   r   )r\   r]   r^   r_   r   rN   r`   ra   s   @r2   rc   rc      s    
> r3   rc   c                   6     e Zd ZdZd fd	Zd Zd Zd Z xZS )Quantum3DClassifieru   
    Quantum 3D Vertex Classifier:
    - 100 sample batch → 100×12 3D vertices
    - Vertices evolve through ODE
    - Classification based on evolved states + statistics
    c                    t         |           || _        || _        || _        t        d|d      | _        t        d|d      | _        t        j                  t        j                  |dz  dz  d      t        j                         t        j                  dd            | _        t        j                  t        j                  |dz  dz  d      t        j                         t        j                  dd	            | _        t        j                  t        j                  d
d      t        j                         t        j                  d      t        j                  dd      t        j                         t        j                  d      t        j                  d|            | _        t        j"                  t%        j&                  d            | _        y )NrY   r[   )r   r   r   r   r   )ri   r   rr   r      r      ru   2   )r   r   num_classesr   t_spanr   
vertex_genrc   ode_funcr   r$   r+   r'   trajectory_encoderstat_encoderrm   
classifierr    r!   r"   
time_embed)r0   r   r   r   r1   s       r2   r   zQuantum3DClassifier.__init__   sU   &( 1%
 -%
 #%--IIlQ&*C0GGIIIc3#
 MMIIlQ&*C0GGIIIc2
 --IIh$GGIJJsOIIc3GGIJJsOIIc;'
 ,,u{{27r3   c                     |d   }|d   }|j                  d      }|j                  d      }t        j                  |||gd      S )zp
        Compute features from ODE trajectory.
        trajectory: (num_steps, batch, num_vertices * 3)
        r   rt   r5   )meanstdr!   rx   )r0   
trajectoryinitialfinal	mean_trajstd_trajs         r2   compute_trajectory_featuresz/Quantum3DClassifier.compute_trajectory_features  sT     Q- 2 OOO*	 >>a>(yy'5)4"==r3   c                    |j                  |j                  d   d      }|j                  d      }|j                  d      }|j	                  d      d   }|j                  d      d   }t        j                  ||||||z
  gd      }|j                  d      j                  d| j                  d      }|j                  |j                  d   d      }|S )zf
        Compute statistics from vertex positions.
        vertices: (batch, num_vertices, 3)
        r   rt   r   r5   )rv   r7   r   r   maxminr!   rx   	unsqueezerz   r   reshape)r0   rM   flatr   r   max_vmin_vstatss           r2   compute_vertex_statisticsz-Quantum3DClassifier.compute_vertex_statistics  s     }}X^^A.3 }}}#llql!#A&#A& 		4eUEEMBK "))"d.?.?Dekk!nb1r3   c                    |j                   d   }| j                  |      }|j                  |d      }t        j                  | j
                  d   | j
                  d   d|j                        }t        | j                  ||d      }| j                  |      }| j                  |      }|d   j                  || j                  d      }	| j                  |	      }
| j                  |
      }t        j                  ||| j                  j!                  d      j#                  |d      gd	      }| j%                  |      }|S )
z
        Args:
            x: (batch, 1, 28, 28) MNIST images
        Returns:
            logits: (batch, num_classes)
        r   rt   r   r   devicerk4methodr   r5   )r7   r   r   r!   linspacer   r   r   r   r   r   r   r   r   rx   r   r   rz   r   )r0   r?   r   rM   r|   t_pointsr   traj_featurestraj_encodedfinal_verticesvertex_statsstats_encodedr   logitss                 r2   rN   zQuantum3DClassifier.forward,  s1    WWQZ
 ??1% ((R8 >>$++a.$++a."QXXV MM	

 88D..}= $B//
D<M<MqQ 55nE)),7 99OO%%a(//
B?
 	 *r3   )
   rZ   r   r   )	r\   r]   r^   r_   r   r   r   rN   r`   ra   s   @r2   r   r      s    /8b>&.-r3   r   c                   0     e Zd ZdZd fd	Zd Zd Z xZS )Quantum3DClassifierLitez@
    Lighter version with more efficient vertex processing.
    c                    t         |           || _        || _        || _        || _        ||z  | _        d| _        t        d| j                        | _	        d| _
        t        j                  t        j                  dddd      t        j                  d      t        j                         t        j                   d      t        j                  dddd      t        j                  d      t        j                         t        j"                  d	      t        j$                         t        j&                  d
d      t        j                               | _        t        j&                  d| j                        | _        t        j&                  d| j                        | _        t        j                  t        j&                  | j                  dz   d      t        j.                  d      t        j0                         t        j&                  dd      t        j.                  d      t        j0                         t        j&                  d| j                              | _        t        j                  t        j&                  | j                  dz  d      t        j                         t        j4                  d      t        j&                  d|            | _        y )Nr   r   )posvelaccangles	angle_vel	angle_accr      r   r   r   )r   r   i   r   i  r   r   ru   )r   r   r   r   r   state_dim_per_vertex	state_dimposition_dimsliceposition_slicestate_groupsr   r$   r%   r&   r'   r(   r)   r*   r+   encodervertex_headvertex_scalerk   rl   ode_netrm   r   )r0   r   r   r   r   r1   s        r2   r   z Quantum3DClassifierLite.__init__a  s   &($8!%(<<#At'8'89
 }}IIaQ*NN2GGILLOIIb"a+NN2GGI  (JJLIIgs#GGI
  99S$..9IIc4>>: }}IIdnnq(#.LLGGIIIc3LLGGIIIc4>>*
 --IIdnnq(#.GGIJJsOIIc;'	
r3   c           	          | j                  t        j                  ||j                  |j                  d   d      gd            S )Nr   r   rt   r5   )r   r!   rx   rz   r7   )r0   r{   states      r2   r   z Quantum3DClassifierLite.ode_func  s5    ||EIIuahhu{{1~q.I&JPRSTTr3   c                 "   |j                   d   }| j                  |      }| j                  |      }t        j                  | j                  |            dz  dz   }||z  }t        j                  | j                  d   | j                  d   d|j                        }t        | j                  ||d      }|d   }|d	   }	|j                  d
      }
|j                  d
      }t        j                  ||	|
|gd	
      }| j                  |      S )Nr   r   r   r      r   r   r   rt   r5   )r7   r   r   r!   r8   r   r   r   r   r   r   r   r   rx   r   )r0   r?   r   r@   rM   rB   r   r   r   r   mean_tstd_tr   s                r2   rN   zQuantum3DClassifierLite.forward  s    WWQZ
 <<? ##H-t00:;a?#Ef$ >>$++a.$++a."QXXVDMM8XeL
 Q-2Q'1%99gufe<"Ex((r3   )r   rZ   r         ?   )r\   r]   r^   r_   r   r   rN   r`   ra   s   @r2   r   r   \  s    4
lU)r3   r   c                    |j                   }t        | d      r| j                  d   nd}t        j                  d|||      }t        | d      rx| j                  |      }|j                  |j                  d   d      }t        | j                  ||d	      }|}	|j                  ||j                  d   | j                  d
      }
n | j                  |      }| j                  |      }t        j                  | j                  |            dz  dz   }||z  }t        | j                  ||d	      }|j                  |j                  d   | j                  | j                        }|dddd| j                   f   }	|j                  ||j                  d   | j                  | j                        dddddd| j                   f   }
|	|
||fS )zDReturn initial vertices and ODE trajectory for either model variant.r   r   r   r   r   r   rt   r   r   r   r   r   N)r   hasattrr   r!   r   r   r   r7   r   r   r   r   r   r8   r   r   r   )modelimages	num_stepsr   t_endr   rM   r|   r   initial_verticesposition_trajectoryr@   rB   latent_verticess                 r2   get_vertex_trajectoryr     s   ]]F&uh7ELLOSE~~a	&AHul###F+ ((a"=ENNM8ER
#(00FLLOUM_M_abc==())(3u11(;<q@3F%.ENNM8ER
'//LLOU//1K1K
 +1a1E1E+EF(00v||A(:(:E<V<V

Q5''
') 0*hFFr3   c           	      N
   | j                          t        t        |            }|d   d| j                  |      |d   d| }}t	        j
                         5  t        | |d      \  }}}	}
ddd       t        j                  d      }t        t        |d            D ]=  }|j                  d	d
|dz  dz         }|j                  ||   j                         j                         d       |j                  d||   j!                          d       |j#                  d       |j                  d	d
|dz  dz   d      }|   j                         j%                         }|j'                  |dddf   |dddf   |dddf   t        d      dd       |j                  dd       |j)                  d       |j+                  d       |j-                  d       @ |j                  d	d
dd      }d}t/        j0                  dj2                  d   dz
  d	t4              }|D ]f  }|||f   j                         j%                         }|j'                  |dddf   |dddf   |dddf   d
|   j!                         dd d!"       h |j7                          |j                  d#       |j                  d	d
d$      }t        d| j8                  d%      D ]r  }|dd||f   j                         }t	        j:                  |d&      j%                         }|j=                  
j                         j%                         |d'| (       t |j)                  d)       |j+                  d*       |j                  d+       |j7                  d       |j?                  d,       |j                  d	d
d-d      }|dd|df   j                         j%                         }t        j@                  jC                  t/        j0                  dd|j2                  d               }t        |j2                  d   dz
        D ]9  }|j=                  |||dz   df   |||dz   df   |||dz   df   ||   d.       ; |j'                  |d/   |d0   |d1   d2dd34       |j'                  |d5   |d6   |d7   d8dd94       |j7                          |j                  d:       t        jD                          t        jF                  d;d<d=>       t        jH                          y# 1 sw Y   xY w)?z.Visualize generated 3D vertices and evolution.r   Nr   r   r   )r   rZ   figsize   r      r   graycmapzDigit: r   fontsizeoff3d
projectionrZ   viridisr[   )cr   szInitial VerticesXYZ)      )dtypezt=.2fr   g?)labelr   alphazVertex Evolution (Sample 0)   r   r5   zVertex )r   TimezDistance from originzVertex Distance Over TimeTr   )color	linewidth)r   r   )r   r   r   greenStart)r   r   r   )rt   r   )rt   r   )rt   r   redEndzVertex 0 Trajectory in 3Dzquantum_vertices.png   tight)dpibbox_inches)%evalnextitertor!   no_gradr   pltfigurer:   r   add_subplotimshowcpusqueeze	set_titleitemaxisnumpyscatter
set_xlabel
set_ylabel
set_zlabelnpr   r7   intlegendr   r;   plotgridcmplasmatight_layoutsavefigshow)r   test_loaderr   num_samplessamplesr   labelsrM   r   r   r   figiax1ax2v_initax3
sample_idx	t_indicest_idxv_trajax4rG   valsdistsax5v0_trajcolorss                               r2   visualize_verticesr2    s    
JJL4$%GQZ-008'!*\k:RFF	 h5J5RXdf5g2%q(h **X
&C 3{A&' ooaAaC!G,

6!9==?**,6
:q	 012R@ ooaAaC!Go=!"((*F1a4L&A,q!tb	PY]`a(15sss" //!QT/
:CJA288;a?#NI M$UJ%67;;=CCEF1a4L&A,q!thuo224S9:b 	 	MM JJLMM/0 //!Q
#C1e((!, E"1j!#3488:

4Q'--/%%'smDE NN6NN)*MM-.JJJHHTN //!Qt/
4C!!Z"23779??AGVV]]2;;q!-@-F-Fq-IJKF&,,Q/!34 .1Q3"GAacE1H$5wq1uax7HQi1 	 	.. KKwt}gdmw#U\K]KK%3V[K\JJLMM-.KK&CWEHHJyh hs   TT$quantum_vertices.gif<   c           	      R	    !"  j                          t        t        |            }|d   d| j                  |      |d   d| }}t	        j
                         5    |      }	|	j                  d      }
t         ||      \  }}}}ddd       d}|   j                         j                         j                         }dd|f   j                         j                         j                         !j                         j                         j                         t         d      rj                  ||j                  d    j                   j                        dd|f   j                         j                         j                         }t         j"                  j%                  |ddddddf   d	
      "n!t!        j&                  | j                  f      "t)        |j)                         !j)                               }t+        |j+                         !j+                               }t+        dd||z
  dz   z        }t-        j.                  d      }|j1                  dd	d      }|j1                  dd	d	d      }|j3                  ||   j                         j                         j5                         d       |j7                  d||   j9                          d
|   j9                          d       |j;                  d       t,        j<                  j?                  t!        j@                  dd j                              }g  |D ],  }|jC                  g g g |dd      \  } jE                  |       . |jG                  !ddddf   !ddddf   !dddd	f   |d      |j7                  d      |jI                  d       |jK                  d        |jM                  d!       |jO                  ||z
  ||z          |jQ                  ||z
  ||z          |jS                  ||z
  ||z            !"fd"}tU        |||d#d$%      }|jW                  |tY        d&'      (       t-        jZ                  |       t]        d)|        y# 1 sw Y   xY w)*z5Export an animation of the learned ODE vertex motion.r   Nr   r5   r   r   r   r   r   )r  g      ?re   ư>)r   r   r   r   r   r   r   zTrue: z	 | Pred:    r   r   r   g333333?)r   r   r   P   )r   r    r   r   r   c           
         |    }dd|    j                         dz   z  z  z   }|d d df   |d d df   |d d df   f_        j                  |       t        
      D ]L  \  }}d | dz   |d d f   }|j	                  |d d df   |d d df          |j                  |d d df          N 	j                  d|    dd	t        d
d       d       	g
S )N(   r8  r6  r   r   r   z12-Vertex ODE Motion | t=r   z	 | state=r   r   zD/vertex)r   
_offsets3d	set_sizes	enumerateset_dataset_3d_propertiesset_textgetattr)	frame_idxcoordssizes
vertex_idxlinehistoryr   r  t_nptitletrail_linestraj_npvelocity_mags         r2   updatez animate_vertices.<locals>.updateT  s   #R<	2l6F6F6H46OPQQ$QTlF1a4L&A,G%  )+ 6 	2Jny1}nj!;<GMM'!Q$-A7""71a4=1	2
 	'Y'< =U$:A>?xI	
 ---r3   Z   F)framesintervalblitrZ   )fps)writerzAnimation saved to )/r  r  r  r  r!   r  argmaxr   detachr
  r  r   r   r7   r   r   r  linalgr;   onesr   r   r  r  r  r	  r  r  r  r  r  r   r   r  r=   r  r  r  r  set_xlimset_ylimset_zlimr   saver   closeprint)#r   r  r   r  output_pathr   r   r   r!  r   predictionsr   r   full_trajectoryr   r(  
initial_np	latent_npxyz_minxyz_maxpadr"  ax_imgax_3dvertex_colorsr   rG  rN  animr  rI  rJ  rK  rL  rM  s#   `                            @@@@@@r2   animate_verticesrk    s   	JJL4$%GQZ-008'!*\k:RFF	 
vmmm*K`6YL
H-
 J!*-446::<BBDJ!!Z-0779==?EEGG??  "((*Du,-#++v||A(:(:E<V<V

Z-##% 	 yy~~i1ac	&:~Cww	5+=+=>?*.."GKKM2G*.."GKKM2G
dC7W,t34
5C
**W
%C__Q1%FOOAq!O5E
MM&$++-113;;=FMK


#((*+9[5L5Q5Q5S4TU   KKFFNN2;;q!U5G5G#HIMK !

2r2Uc
M4 ! mm1a1a1a

  G OOBE	S	S	S	NN7S='C-0	NN7S='C-0	NN7S='C-0. ." fY%PDIIk,2"6I7IIcN	}
-.W
 
s   .RR&c                 r   t        j                  ddd      \  }\  }}|j                  | dd       |j                  dd	       |j	                  d
       |j                  d       |j                  d       |j                  |ddd       |j                  |ddd       |j                  dd	       |j	                  d
       |j                  d       |j                          |j                  d       t        j                          t        j                  dd       t        j                          y)zPlot training progress.r   r   )rZ   r   r   zb-)r   zTraining Lossr   r   EpochLossTTrain)r   r   zr-TestAccuracyzAccuracy (%)ztraining_progress.pngr   )r   N)r  subplotsr  r  r  r  r  r  r  r  r  )train_losses
train_accs	test_accsr"  r$  r%  s         r2   visualize_trainingrv  k  s     ll1a9OC#sHH\41H-MM/BM/NN7NN6HHTNHHZWH:HHYFaH8MM*rM*NN7NN>"JJLHHTNKK'S1HHJr3   c                 J   | j                  |      } t        j                         }t        j                  | j                         dd      }t        j                  j                  ||      }g }g }	g }
t        d|        t        d       t        |      D ]  }| j                          d}d}d}t        |      D ]K  \  }\  }}|j                  |      |j                  |      }}|j                           | |      } |||      }|j                          t        j                  j                  j!                  | j                         d	       |j#                          ||j%                         z  }|j'                  d
      \  }}||j)                  d      z  }||j+                  |      j-                         j%                         z  }|dz  dk(  st        d|d
z    d| dt/        |       d|j%                         d       N |j#                          d|z  |z  }t1        | ||      }|t/        |      z  }|j3                  |       |	j3                  |       |
j3                  |       t        d|d
z    d| d|dd|dd|dd       t        d        ||	|
fS )z Train the Quantum 3D classifier.gMbP?g{Gz?)lrweight_decay)T_maxzTraining on <============================================================g        r   g      ?r   r   z  Epoch z	 | Batch /z	 | Loss: z.4f      Y@zEpoch z
 | Train: r   z
% | Test: %z<------------------------------------------------------------)r  r   CrossEntropyLossoptimAdamW
parameterslr_schedulerCosineAnnealingLRr^  r:   trainr>  	zero_gradbackwardr!   utilsclip_grad_norm_stepr  r   sizeeqsumlenevaluater=   )r   train_loaderr  epochsr   	criterion	optimizer	schedulerrs  rt  ru  epochrunning_losscorrecttotal	batch_idxdatatargetoutputlossr   	predicted	train_acctest_accavg_losss                            r2   train_modelr    s    HHVE##%IE,,.5tLI""44Yf4MILJI	L
!"	(Ov $)2<)@ 	n%I~f776?FIIf,=&D!4[FVV,DMMOHHNN**5+;+;+=sCNNDIIK'L!::a=LAyV[[^#Ey||F+//16688G3!#q	9+Qs<?P>QQZ[_[d[d[fgjZklm!	n$ 	 7NU*	E;7#l"33H%)$"uQwiq	(3z)TWXbcklobppqrshI$L Y..r3   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 on test set.r   r   Nr}  )	r  r!   r  r  r   r  r  r  r  )
r   r  r   r  r  r  r  r  r   r  s
             r2   r  r    s    	JJLGE	 9' 	9LD&776?FIIf,=&D4[F!::a=LAyV[[^#Ey||F+//16688G	99 '>E!!9 9s   BCCc           	          | j                   j                  | j                  | j                  | j                  t        | dd      | j                         d}t        j                  ||       t        d|        y)z"Save the trained model checkpoint.r   r   )
model_typer   r   r   r   
state_dictzCheckpoint saved to N)
r1   r\   r   r   r   rB  r  r!   r\  r^  )r   path
checkpoints      r2   save_checkpointr    sj     oo..((**,, '/Eq I&&(J 
JJz4 	 
'(r3   c           	      ~   t        j                  | |      }|d   }|dk(  r2t        |d   |d   t        |d         |j	                  dd      	      }n4|d
k(  r!t        |d   |d   t        |d               }nt        d|       |j                  |d          |j                  |      }|j                          |S )zLoad a saved model checkpoint.)map_locationr  r   r   r   r   r   r   r   r   r   r   r   )r   r   r   z&Unsupported model type in checkpoint: r  )
r!   loadr   tuplegetr   
ValueErrorload_state_dictr  r  )r  r   r  r  r   s        r2   load_checkpointr    s    Dv6JL)J..'"=1#N3H-.!+0F!J	
 
,	,#"=1#N3H-.
 A*NOO	*\23HHVE	JJLLr3   c                    t         j                  j                         rdnd}d}d}d}d}d}t        d       t        d	|        t        d
|        t        d|        t        d|        t        d       t	        j
                  t	        j                         g      }t        j                  ddd|      }t        j                  ddd|      }	t        ||dd      }
t        |	|dd      }t        dt        |       dt        |	              t        d       | j                  r/t        | j                  |      }t        d| j                          nt        d|d|      }t        d       t        d       t        d| d|j                   d       t        d       t        d        t        d       | j                  rg g g }}}n t        ||
|||!      \  }}}t        ||       t        d"       t        d#       t        d       t!        |||      }t        d$|d%d&       t        d       |rt#        |||       t%        |||       t'        |||d| j(                  | j*                  '       ||||fS )(Ncudar
  r[   r   rZ   r   zex04_model.ptz+Trainable reference implementation: ex04.pyzDevice: zBatch Size: zVertices per sample: zState dim per vertex: r{  z../dataT)rootr  download	transformFr   )r   shufflenum_workerszTrain: z	 | Test: z
Loaded checkpoint: r   r   r  z
Model Architecture:u"   Input: 100 × 1 × 28 × 28 imagesu   Output: 100 × u    × z latent statez9State layout: pos, vel, acc, angles, angle_vel, angle_acczODE t_span: (0, 1.5))r  r   z=
============================================================zFINAL RESULTSzFinal Test Accuracy: r   r~  )r  r_  r   )r!   r  is_availabler^  
transformsComposeToTensordatasetsMNISTr   r  r  r   r   r  r  r  rv  r2  rk  animation_outputanimation_steps)argsDEVICE
BATCH_SIZEEPOCHSNUM_VERTICESSTATE_DIM_PER_VERTEXCHECKPOINT_PATHr  train_datasettest_datasetr  r  r   rs  rt  ru  final_test_accs                    r2   mainr    s   zz..0VeFJFL%O	
78	HVH
	L
%&	!,
01	"#7"8
9:	(O ""$ I NN	XabM>>yXabLm
D^_`L\j%]^_K	GC&'y\1B0C
DE	(O 4 4f=%d&:&:%;<= (%!5	
 

!"	.0	OL>e.H.H-I
WX	
EF	 "	(O.0"b)j /:<&/
+j) 	/ 
/	/	(Oe[&9N	!.!5Q
78	(O <Y?uk62))&& ,
I55r3   __main__z1Train or visualize the ex04 quantum vertex model.)descriptionz--load-checkpointr9  z8Load a saved checkpoint instead of training a new model.)typedefaulthelpz--animation-outputz Path for the exported animation.z--animation-stepsz0Number of ODE frames to render in the animation.)r   )r   )r   r3  r4  )r   r  )r  )r
  )7r!   torch.nnr   torch.optimr  torch.utils.datar   torchvision.datasetsr  torchvision.transformsr  matplotlib.pyplotpyplotr  matplotlib.animationr   r   mpl_toolkits.mplot3dr   r  r  scipy.spatial.transformr   argparser   r   ImportError
subprocessrunModuler   rc   r   r   r   r2  rk  rv  r  r  r  r  r  r\   ArgumentParserparseradd_argumentstrr  
parse_argsparsed_argsr   lossesr  r   r3   r2   <module>r     sz      ' ' +  < '  , #"f;RYY f;RF299 FRP")) PfU)bii U)pG:CLQ/h05/p"")6P6f z$X$$1deF
+#rW  Y
,3@V?  A
+#rO  Q##%K)-k):&E69h o  #JNN:;""#s   D7 7#EE