
    ijs                     
   d dl Zd dlmZ d dlmZ d dlmZ  ed      d   j                  dd      j                  ej                        dz  Z ed	      d   d
z  j                  e      Z ed      d   j                  dd      j                  ej                        dz  Z ed      d   d
z  j                  e      Z edej$                   dej$                           edej$                   dej$                          d Zd Zd Zd Zd@dZd Zd Z G d d      Z G d d      Z G d d      ZdAdZedk(  rd Zd!Z d"Z!d#Z"d$Z#dZ$ ed%        ed&        ed%        ed'e         ed(e!         ed)e"         ed*e#         ed+e$         ed,        ed-        ed%        edee e#e!e"d./      Z% ed0        ed1 e&e%jN                  jP                  e%jN                  jR                  e%jN                  jT                  fD  cg c]  } | jV                   c}                ed%        ed2        ed3        ed%       e%jY                  ee      d$z  Z- ed4e-d5d6       e%j]                  ede# ede#       \  Z/Z0 ed7e/d8        ed%        ee%eee$9        ed:       e%jN                  jP                  e%jN                  jb                  e%jN                  jR                  e%jN                  jd                  e%jN                  jT                  e%jN                  jf                  e%jN                  jh                  e%jN                  jj                  e%jN                  jl                  e%jN                  jn                  e%jN                  jp                  e%jN                  jr                  e%jN                  jt                  d;e%jv                  jx                  e%jv                  jz                  e%jv                  j|                  e%jv                  j~                  e%jv                  j                  e%jv                  j                  d<e%j                  e%j                  e%j                  e%j                  e%j                  e%j                  e%j                  e%j                  e%j                  e%j                  e%j                  e%j                  e%j                  e%j                  d=d>ZP ej                  dBi eP  ed?       yyc c} w )C    N)read)Axes3Dz../X_train.wav     g     o@z../y_train.wav	   z../X_test.wavz../y_test.wavzTrain: z
, Labels: zTest: c                 .    t        j                  d|       S )Nr   )npmaximumxs    C/home/per/Documents/python/circle_diff_eq/quantum3d_vertex_model.pyrelur      s    ::a    c                 6    t        j                  | dkD  dd      S )Nr         ?g        )r
   wherer   s    r   relu_derivativer      s    88AE3$$r   c           	      d    ddt        j                  t        j                  | dd             z   z  S )Nr   ii  )r
   expclipr   s    r   sigmoidr      s+    #4 556677r   c                 (    t        |       }|d|z
  z  S )Nr   )r   )r   ss     r   sigmoid_derivativer      s    
AA;r   c                     t        j                  | t        j                  | |d      z
        }|t        j                  ||d      z  S )NTaxiskeepdims)r
   r   maxsum)r   r   exp_xs      r   softmaxr#   !   s:    FF1rvvadT::;E266%dT:::r   c                 ,    t        j                  |       S )Nr
   tanhr   s    r   r&   r&   %   s    771:r   c                 8    dt        j                  |       dz  z
  S )Nr      r%   r   s    r   tanh_derivativer)   (   s    rwwqz1}r   c                   0    e Zd ZdZddZd Zd Zd Zd Zy)	QuantumVertexGeneratoru   
    Converts batch of images to 100 × 12 × 3 vertex points.
    Uses learned weights for feature extraction and vertex generation.
    c                 r   t         j                  j                  |       || _        || _        || _        t         j                  j                  |d      dz  | _        t         j                  j                  |d      t        j                  d|z        z  | _	        t        j                  d      | _        t         j                  j                  dd      t        j                  d      z  | _        t        j                  d      | _        t         j                  j                  dd      t        j                  d      z  | _        t        j                  d      | _        t         j                  j                  d|d	z        d
z  | _        t        j                  |d	z        | _        t         j                  j                  d|dz        d
z  | _        t        j                  |dz        | _        t         j                  j                  dd      d
z  | _        t        j                  d      | _        i | _        y )N         ?          @         ?@         ?   {Gz?)r
   randomseednum_vertices
batch_size
input_sizerandnbase_verticessqrtW1zerosb1W2b2W3b3W_rotb_rotW_scaleb_scaleW_transb_transcache)selfr;   r9   r:   r8   s        r   __init__zQuantumVertexGenerator.__init__5   s   
		t($$  YY__\1=C ))//*c2RWWS:=M5NN((3- ))//#s+bggi.@@((3- ))//#r*RWWY-??((2, YY__R)9:TA
XXlQ./
 yyr<!+;<tCxxq 01 yyr1-4xx{ 
r   c                    |j                   d   }|| j                  d<   || j                  z  | j                  z   }t	        |      }|| j                  d<   || j
                  z  | j                  z   }t	        |      }|| j                  d<   || j                  z  | j                  z   }t	        |      }|| j                  d<   || j                  z  | j                  z   }t        || j                  z  | j                  z         }t        || j                  z  | j                   z         }|| j                  d<   || j                  d<   || j                  d<   |j#                  || j$                  d	      }	|	t&        j(                  j+                  |	d
d      dz   z  }	|j#                  || j$                  d      }
| j-                  | j.                  |	      }||
dz   z  |ddt&        j0                  ddf   z   }|| j                  d<   |S )z
        Args:
            X: (batch, 784) normalized images
        Returns:
            vertices: (batch, num_vertices, 3) 3D vertex points
        r   x0h1h2h3quaternionsscalestranslationr5   r(   Tr   :0yE>r-   r.   Nvertices)shaperL   r?   rA   r   rB   rC   rD   rE   rF   rG   r   rH   rI   r&   rJ   rK   reshaper9   r
   linalgnorm_rotate_vertices_by_quaternionr=   newaxis)rM   Xr:   rQ   rR   rS   rT   rU   rV   qscales_reshapedrotated_verticesrX   s                r   forwardzQuantumVertexGenerator.forward[   s    WWQZ
 

4 [477""X

4 $''\DGG#"X

4 $''\DGG#"X

4 4::o

2dll*T\\9:2,t||;<$/

=!%

8$/

=!
D,=,=qAD9D@A ..T5F5FJ>>t?Q?QSTU#'<=Arzz[\L\@]]!)

:r   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 = [w, x, y, z]r   r   r(   r-   r.   )r
   array)rM   vr`   wr   yzvxvyvzresult_xresult_yresult_zs                r   _rotate_by_quaternionz,QuantumVertexGenerator._rotate_by_quaternion   s4   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xx8X677r   c                 R   |dddddf   }|dddddf   }|dddddf   }|dddddf   }|t         j                  dddf   }|t         j                  dddf   }|t         j                  dddf   }	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d      S )z4Rotate all base vertices for a batch of quaternions.Nr   r   r(   r-   r.   r   )r
   r^   stack)rM   r=   rT   rg   r   rh   ri   rj   rk   rl   rm   rn   ro   s                r   r]   z5QuantumVertexGenerator._rotate_vertices_by_quaternion   s}   1a 1a 1a 1a 2::q!+,2::q!+,2::q!+,1qsSB.!A#!)r1AAQqS1Q3YRTDTTU1qsb(AaC!A#IOr+AAQqS1Q3YRTDTTU1qsb(AaC!A#I+;;qsQqSy3RT>TTUxx8X6Q??r   c                 	   |j                   d   }t        j                  || j                  dz  f      }t        j                  || j                  dz  f      }t        j                  |df      }t        j                  | j
                        }t        |      D ]  }t        | j                        D ]~  }| j                  d   ||dz  |dz   dz  f   }	| j                  d   ||dz  |dz   dz  f   }
t        j                  j                  |	      dz   }|	|z  }|d   |d   |d   |d   f\  }}}}| j
                  |   }t        j                  d      }|||f   }d||d   z  ||d   z  z   ||d   z  z   z  |d<   d||d   z  ||d   z  z   ||d   z  z
  z  |d<   d| |d   z  ||d   z  z   ||d   z  z   z  |d<   d||d   z  ||d   z  z
  ||d   z  z   z  |d<   ||||dz  |dz   dz  f<   ||z  t        |
      z  |||dz  |dz   dz  f<   ||xx   |z  cc<   ||xx   ||
d	z   z  z  cc<     || j                  j                  z  }|| j                  j                  z  }|| j                  j                  z  }||z   |z   }|t        | j                  d
         z  }|| j                   j                  z  }|t        | j                  d         z  }|| j"                  j                  z  }|t        | j                  d         z  }|| j$                  j                  z  }| j                  d   j                  |t        | j                  d         z  z  t        j&                  |t        | j                  d         z  d      | j                  d   j                  |t        | j                  d         z  z  t        j&                  |t        | j                  d         z  d      | j                  d   j                  |t        | j                  d
         z  z  t        j&                  |t        | j                  d
         z  d      | j                  d
   j                  |z  t        j&                  |d      | j                  d
   j                  |z  t        j&                  |d      | j                  d
   j                  |z  t        j&                  |d      |d}|S )ze
        Backprop through vertex generation.
        grad_vertices: (batch, num_vertices, 3)
        r   r5   r-   rT   r   rU   rW   r(   r.   rS   rR   rQ   rP   rr   )r?   rA   rB   rC   rD   rE   rF   rG   rH   rI   rJ   rK   r=   )rY   r
   r@   r9   
zeros_liker=   rangerL   r[   r\   r   rF   TrH   rJ   r   rD   rB   r?   r!   )rM   grad_verticesr:   grad_quaternionsgrad_scalesgrad_translationgrad_base_verticesbrf   r`   scale_vq_normq_normalizedrg   r   rh   ri   v_basegrad_qgvgrad_h3_rotgrad_h3_scalegrad_h3_transgrad_h3grad_h2grad_h1grad_x0grad_weightss                               r   backwardzQuantumVertexGenerator.backward   s   
 #((+
 88Z1B1BQ1F$GHhh
D,=,=,ABC88ZO4  ]]4+=+=>z" 	>A4,,- >JJ}-a1ac1Wn=**X.q!A#qsAg+~>*T1 6z)!_l1o|AP\]^P__
1a++A. !"1a4( 2a51RU7!2Qr!uW!<=q	2a51RU7!2Qr!uW!<=q	!BqEAbeG!3a1g!=>q	2a51RU7!2Qr!uW!<=q	39 AaC1aK0 /16k<Nw<W.WAqsAaC7{N+ !#r)# #1%w})==%=>	>D '5#dllnn4(4<<>>9-= ODJJt,<==DGGII%ODJJt,<==DGGII%ODJJt,<==DGGII% **T"$$/$**TBR2S(ST&&?4::d3C#DD1M**T"$$/$**TBR2S(ST&&?4::d3C#DD1M**T"$$/$**TBR2S(ST&&?4::d3C#DD1MZZ%''*::VV,15zz$'))K7vvk2zz$')),<<vv.Q7/
  r   N)r      d   *   )	__name__
__module____qualname____doc__rN   rc   rp   r]   r    r   r   r+   r+   /   s$    
$L,\8@"Qr   r+   c                   ,    e Zd ZdZddZd Zd ZddZy)	QuantumVertexODEzf
    ODE function that evolves vertex points through time.
    Uses learned weights for dynamics.
    c                 P   t         j                  j                  |       || _        || _        ||z  | _        t         j                  j                  | j
                  dz   |      t        j                  d| j
                  dz   z        z  | _        t        j                  |      | _
        t         j                  j                  ||      t        j                  d|z        z  | _        t        j                  |      | _        t         j                  j                  || j
                        dz  | _        t        j                  | j
                        | _        t         j                  j                  | j                  d      dz  | _        t         j                  j                  | j                  d      dz  | _        t         j                  j                  | j                  d      dz  | _        t         j                  j                  d| j                        dz  | _        i | _        y )Nr   r0   r6   r3   )r
   r7   r8   r9   
vertex_dim	state_dimr<   r>   W_dyn1r@   b_dyn1W_dyn2b_dyn2W_dyn3b_dyn3W_attn_qW_attn_kW_attn_v
W_attn_outrL   )rM   r9   r   
hidden_dimr8   s        r   rN   zQuantumVertexODE.__init__  sp   
		t($%
2 iioodnnq&8*EPSW[WeWehiWiPjHkkhhz*iiooj*=jHX@YYhhz*iiooj$..ADHhht~~. 		<tC		<tC		<tC))//"doo>E
r   c                    |j                  d      }|j                  | j                  | j                        }|| j                  z  }|| j                  z  }|| j
                  z  }||j                  z  t        j                  d      z  }t        |      }||z  }	|	| j                  z  }
t        j                  ||gg      }t        || j                  z  | j                  z         }t        || j                  z  | j                   z         }|| j"                  z  | j$                  z   }|
j'                         }||dz  z   }|S )z
        Compute dx/dt for vertex evolution.
        
        Args:
            t: current time (scalar)
            state: (state_dim,) flattened vertex positions
        Returns:
            d_state/dt: (state_dim,)
        r   r3   皙?)rZ   r9   r   r   r   r   rw   r
   r>   r#   r   concatenater   r   r   r   r   r   r   flatten)rM   tstaterX   QKVattn_scoresattn_weightsattn_outattn_verticescombinedrQ   rR   dynamics	attn_flats                   r   r   zQuantumVertexODE.dynamics  s"    b! ==!2!2DOOD t}}$t}}$t}}$ !##g+{+!# 4??2 >>51#,/ (T[[(4;;67"t{{"T[[01#dkk1 "))+	i#o-r   c                    |j                  d| j                  | j                        }|| j                  z  }|| j                  z  }|| j
                  z  }t        j                  |t        j                  |dd            t        j                  d      z  }t        |d      }||z  }	|	| j                  z  }
t        j                  |j                  d   df|      }t        j                  ||gd      }t        || j                   z  | j"                  z         }t        || j$                  z  | j&                  z         }|| j(                  z  | j*                  z   }||
j                  |j                  d   d      dz  z  }|S )z
        Batched ODE dynamics.

        Args:
            t: current time (scalar)
            state_batch: (batch, state_dim)
        Returns:
            d_state/dt: (batch, state_dim)
        r   r   r(   g      P@rr   r   r   )rZ   r9   r   r   r   r   r
   matmulswapaxesr>   r#   r   fullrY   r   r   r   r   r   r   r   r   )rM   r   state_batchrX   r   r   r   r   r   r   r   time_columnr   rQ   rR   r   s                   r   dynamics_batchzQuantumVertexODE.dynamics_batchG  sJ    &&r4+<+<dooNt}}$t}}$t}}$ii2;;q!Q#782774=H{3!# 4??2gg{003Q7;>>;"<1E(T[[(4;;67"t{{"T[[01#dkk1M))+*;*;A*>CcIIr   c                    t        j                  |d   |d   |      }t        j                  ||j                  d   | j                  f      }||d<   t        d|      D ]  }||dz
     }||   ||dz
     z
  }||dz
     }	| j                  ||	      }
| j                  |d|z  z   |	d|z  |
z  z         }| j                  |d|z  z   |	d|z  |z  z         }| j                  ||z   |	||z  z         }|	|dz  |
d|z  z   d|z  z   |z   z  z   ||<    |S )a%  
        Solve the ODE for a batch with fixed-step RK4.

        Args:
            initial_state_batch: (batch, state_dim)
            t_span: (t0, tf) time span
            num_steps: number of trajectory samples
        Returns:
            trajectory: (num_steps, batch, state_dim)
        r   r   r.   g      @r(   )r
   linspaceemptyrY   r   rv   r   )rM   initial_state_batcht_span	num_stepst_eval
trajectoryir   dtr   k1k2k3k4s                 r   solve_batchzQuantumVertexODE.solve_batchf  sH    VAYq	9=XXy*=*C*CA*FWX
+
1q)$ 
	IAq1uAVAE]*Bq1u%E$$Q.B$$Qr\538b=3HIB$$Qr\538b=3HIB$$QVUR"W_=B!R#X"qt)ad2BR2G$HHJqM
	I r   N)r   r-   r1   r   )r   g      ?   )r   r   r   r   rN   r   r   r   r   r   r   r   r      s    
6&P>r   r   c                   B    e Zd ZdZ	 	 d
dZd Zd ZddZd Zd Z	d Z
y	)Quantum3DClassifierz
    Complete classifier: Vertex Generation -> ODE Evolution -> Classification
    All from scratch with numpy and solve_ivp.
    c                    t         j                  j                  |       || _        || _        || _        || _        || _        t        ||||      | _	        t        |dd|dz         | _        t         j                  j                  |dz  dz  d      t        j                  d|dz  dz  z        z  | _        t        j                  d      | _        t         j                  j                  dd      t        j                  d      z  | _        t        j                  d      | _        t         j                  j                  |dz  d	z  d      t        j                  d|dz  d	z  z        z  | _        t        j                  d      | _        t         j                  j                  dd
      t        j                  d      z  | _        t        j                  d
      | _        t         j                  j                  |      dz  | _        t         j                  j                  d|z   d      t        j                  dd|z   z        z  | _        t        j                  d      | _        t         j                  j                  dd      t        j                  d      z  | _        t        j                  d      | _        t         j                  j                  d|      dz  | _        t        j                  |      | _        i | _        y )N)r;   r9   r:   r8   r-   r1   r   )r9   r   r   r8   r/   r0   r2      r3   r4   r6      )r
   r7   r8   num_classesr9   r:   r   num_time_stepsr+   
vertex_genr   oder<   r>   W_traj1r@   b_traj1W_traj2b_traj2W_stat1b_stat1W_stat2b_stat2
time_embedW_cls1b_cls1W_cls2b_cls2W_cls3b_cls3rL   )rM   r;   r9   r   r:   r   r   r8   s           r   rN   zQuantum3DClassifier.__init__  s>   
		t&($, 1!%!	
 $%	
 yy|a'7!';SABGGCS_bcScfgSgLhDiixx}yysC027793EExx} yy|a'7!';SABGGCS_bcScfgSgLhDiixx}yysB/"'')2DDxx| ))//.9D@ iioocN&:C@2773RUXfRfKgChhhhsmiiooc3/"'')2DDhhsmiiooc;7$>hh{+
r   c                 n    |d   }|d   }|j                  d      }t        j                  |||gd      S )z%Extract features from ODE trajectory.r   r   rr   r   )meanr
   r   )rM   r   initialfinal	mean_trajs        r   compute_trajectory_featuresz/Quantum3DClassifier.compute_trajectory_features  s>     Q-2OOO+	~~wy9BBr   c                    |j                  d      }|j                  d      }|j                  d      }|j                  d      }t	        j
                  ||||||z
  gd      }t	        j                  |ddt        j                  ddf   | j                  d      }|j                  |j                  d   d      S )z)Compute statistics from vertex positions.r   rr   Nr   r   )r   stdr    minr
   r   repeatr^   r9   rZ   rY   )rM   rX   r   r   max_vmin_vstatss          r   compute_vertex_statisticsz-Quantum3DClassifier.compute_vertex_statistics  s     }}!}$lll"!$!$c5%GaP 		%2::q 0143D3D1M}}U[[^R00r   c           	         |j                   d   }| j                  j                  |      }|j                  |d      }| j                  j                  || j                  | j                        }| j                  |      }t        || j                  z  | j                  z         }t        || j                  z  | j                  z         }|d   j                  || j                  d      }	| j                  |	      }
t        |
| j                   z  | j"                  z         }t        || j$                  z  | j&                  z         }t)        j*                  ||t)        j,                  | j.                  |df      gd      }t        || j0                  z  | j2                  z         }t        || j4                  z  | j6                  z         }|| j8                  z  | j:                  z   }|r|||fS |S )aH  
        Forward pass through the entire network.
        
        Args:
            X: (batch, 784) normalized images
            return_trajectory: if True, return ODE trajectory for visualization
        Returns:
            logits: (batch, num_classes)
            trajectory: optional (num_steps, batch, state_dim)
        r   r   )r   r   r-   r   rr   )rY   r   rc   rZ   r   r   r   r   r   r   r   r   r   r   r9   r   r   r   r   r   r
   r   tiler   r   r   r   r   r   r   )rM   r_   return_trajectoryr:   rX   vertices_flatr   traj_featuresh_trajfinal_verticesvertex_statsh_statr   h_clslogitss                  r   rc   zQuantum3DClassifier.forward  s    WWQZ
 ??**1- ((R8XX));;)) * 

 88D mdll2T\\ABft||+dll:; $B//
D<M<MqQ 55nE lT\\1DLL@Aft||+dll:; >>662774??ZYZO3\"]def X+dkk9:UT[[(4;;67$t{{2:x//r   c                    g }t        dt        |      | j                        D ]R  }| j                  |||| j                  z          }|j	                  t        j                  t        |      d             T t        j                  |d      S )zPredict class labels.r   r   rr   )	rv   lenr:   rc   appendr
   argmaxr#   r   )rM   r_   predictionsstartr   s        r   predictzQuantum3DClassifier.predict  sw    1c!fdoo6 	CE\\!E%$//*A"BCFryyqAB	C ~~k22r   c                 T    | j                  |      }t        j                  ||k(        S )zCompute accuracy.)r  r
   r   )rM   r_   rh   r  s       r   scorezQuantum3DClassifier.score  s#    ll1oww{a'((r   c                 R   | j                  |      }t        |      }t        j                  t	        |      | j
                  f      }d|t        j                  t	        |            |f<   t        j                  |t        j                  |dz         z         t	        |      z  }||fS )zCompute cross-entropy loss.r   g&.>)	rc   r#   r
   r@   r   r   aranger!   log)rM   r_   y_truer   probabilitiesy_onehotlosss          r   compute_lossz Quantum3DClassifier.compute_loss  s    a 88S[$*:*:;<363v;'/0 x"&&)=">>??#f+MV|r   N)r   r   
   r   r   r   r   )F)r   r   r   r   rN   r   r   rc   r  r  r  r   r   r   r   r     s7    
 EGJL1fC12h3)
r   r      c                    t        j                  |      }||   }||   }| j                  |d      \  }}}	t        j                  t	        |      d      }
|j
                  d   }t        j                  | j                  d   | j                  d   |      }t        j                  d      }t        t        |d            D ]  }|j                  d	d
|dz  dz         }|j                  ||   j                  dd      d       |j                  d||    d|
|    d       |j!                  d       |j                  d	d
|dz  dz   d      }|	|   }|j#                  |dddf   |dddf   |dddf   t        d      ddd       |j                  dd        |j                  d	d
dd      }d}t        j$                  j'                  t        j                  dd|            }t)        d|d	z        }t        d||      D ]]  }|||f   j                  | j*                  d      }||   }|j#                  |dddf   |dddf   |dddf   d|d d!d||   g"       _ |j-                          |j                  d#       |j/                  d$       |j1                  d%       |j3                  d&       |j                  d	d
d'      }|dd|f   j                  || j*                  d      }t        d| j*                  d      D ]#  }|dd|df   }|j5                  ||d(| )       % |j5                  ||dddddf   j7                  d      d*dd+,       |j/                  d-       |j1                  d.       |j                  d/       |j-                  d       |j9                  dd01       |j                  d	d
d2d      }|dd|ddf   }t        |dz
        D ]9  }|j5                  |||dz   df   |||dz   df   |||dz   df   ||   d3       ; |j#                  |d4   |d5   |d6   d7dd8d9:       |j#                  |d;   |d<   |d=   d>dd?d@:       |j-                          |j                  dA       |j/                  d$       |j1                  d%       |j3                  d&       |j                  d	d
dBd      }g dC} t        j$                  j;                  t        j                  dDdEt=        |                   }!t?        |!|       D ]  \  }"}|dd|ddf   }#|j5                  |#dddf   |#dddf   |#dddf   |"dd(| ,       |j#                  |#d4   |#d5   |#d6   |"dFG       |j#                  |#d;   |#d<   |#d=   |"dHd@I        |j                  dJ       |j/                  d$       |j1                  d%       |j3                  d&       |j-                  d       |j                  d	d
dK      }$|dddddf   }%|%j                  d      }&|%j)                  d      }'|%j7                  d      }(|$jA                  ||&|'dLdMdNO       |$j5                  ||(dPdd+,       |$j                  dQ       |$j/                  d-       |$j1                  d.       |$j-                  d       |$j9                  dd01       |j                  d	d
dR      })|j7                  d      }*|)j5                  ||*dddf   dS)       |)j5                  ||*dddf   dT)       |)j5                  ||*dddf   dUdV       |)j                  dW       |)j/                  d-       |)j1                  dX       |)j-                  d       |)j9                  dd01       t        jB                          t        jD                  dYdZd[\       t        jF                          y)]z2Visualize generated 3D vertices and ODE evolution.T)r   r   rr   r   )      )figsizer  r5      r(      gray)cmapzTrue: z
Pred: r  )fontsizeoff3d)
projectionNr   viridisr   g?)cr  r   alphazInitial Vertices)r     r-   zt=.2f2   )labelr   r  r  zVertex Evolution (Sample 0)r_   YZ   zVertex )r#  blackzMean Z)color	linewidthr#  Timez
Z PositionzVertical Motion Over Timeg333333?)r  r   )r(  r)  )r   r   )r   r   )r   r(   greenStarto)r  r   r#  marker)r   r   )r   r   )r   r(   redEndr   zVertex 0 Trajectory)   r  )r   r-   r  r   r   g?   )r(  r   7   )r(  r   r.  zSelected Vertex Paths   skybluegffffff?zZ range)r(  r  r#  navyzVertical Envelope   zCOM XzCOM YzCOM Z)r#  r)  zCenter of Mass MotionPositionzquantum_vertices_numpy.png   tight)dpibbox_inches)$r
   r  rc   r  r#   rY   r   r   pltfigurerv   r   add_subplotimshowrZ   	set_titler   scattercmplasmar    r9   legend
set_xlabel
set_ylabel
set_zlabelplotr   gridr  r   zipfill_betweentight_layoutsavefigshow)+modelX_testy_testnum_samplesindices	X_samples	y_samplesr   r   initial_verticesr  r   	time_axisfigr   ax1ax2v_initax3
sample_idxcolorsstep_stridet_idxv_trajt_valueax4sample_trajrf   z_valsax5v0_trajax6
vertex_idsline_colorsr(  coordsax7z_valuesz_minz_maxz_meanax8center_of_masss+                                              r   visualize_verticesrt  *  s    ii$GwIwI ,1==VZ=+[(FJ())GFO!4K  #IELLOU\\!_iHI **X
&C 3{A&' 6ooaAaC!G,

9Q<''B/f
=y|nH[^4DEPRS ooaAaC!Go=!!$F1a4L&A,q!t2YY#S 	 	B(156 //!QT/
:CJVV]]2;;q!Y78Faa(Kq)[1 QE:-.66u7I7I1ME"F1a4L&A,q!tgc]+r 	 	QQ
 JJLMM/0NN3NN3NN3 //!Q
#CQ
]+33Iu?Q?QSTUK1e((!, 9Q1W%FGA3-89 HHYAq!G,11q19TU]eHfNN6NN< MM-.JJJHHTH //!Qt/
4CJ*+G9q=! .1Q3"GAacE1H$5wq1uax7HQi1 	 	.. KKwt}gdmCws  <KK#U3  8JJLMM'(NN3NN3NN3 //!QT/
:CJ&&..S#s:!GHKZ0 `qQ1W%1vad|VAqD\RS]defdg[hiF4L&,tERPF5M6%=&-uPR[^_	`
 MM)*NN3NN3NN3JJJ //!Q
#C1a7#HLLaL ELLaL E]]]"FYuITQZ[HHYfHJMM%&NN6NN< JJJHHTH //!Q
#C %%1%-NHHYq!t,GH<HHYq!t,GH<HHYq!t,GqHIMM)*NN6NN:JJJHHTHKK,#7KHHJr   __main__r   r  r   r   r   z<============================================================z<Quantum 3D Vertex Classifier Prototype (NumPy + batched RK4)zVertices per sample: zt_span: zTime steps: zEval batch size: zVisualization samples: z5Mode: forward-pass visualization and diagnostics onlyz Trainable reference: run ex04.pyr   )r;   r9   r   r:   r   r   r8   z
Model created successfully!zTotal parameters: ~z=
============================================================zPROTOTYPE EVALUATIONzRandom-weight Test Accuracy: r!  %zSample Batch Loss: z.4f)rS  z
Saving prototype weights...)r?   rA   rB   rC   rD   rE   r=   rF   rG   rH   rI   rJ   rK   )r   r   r   r   r   r   )r   r   r   r   r   r   r   r   r   r   r   r   r   r   )r   r   
classifierz/Weights saved to quantum_classifier_weights.npz)r   )r  )zquantum_classifier_weights.npz)Rnumpyr
   scipy.io.wavfiler   matplotlib.pyplotpyplotr=  mpl_toolkits.mplot3dr   rZ   astypefloat64X_traininty_trainrQ  rR  printrY   r   r   r   r   r#   r&   r)   r+   r   r   rt  r   NUM_VERTICESNUM_CLASSEST_SPANNUM_TIME_STEPSEVAL_BATCH_SIZEVIS_SAMPLESrP  r!   r   r?   rB   rD   sizer  final_test_accr  sample_loss_rA   rC   rE   r=   rF   rG   rH   rI   rJ   rK   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   weightssavez)ps   0r   <module>r     s    !  '
 
 
#
+
+B
4
;
;BJJ
G%
O !!$q(
0
0
5	o	q	!	)	)"c	2	9	9"**	E	M


"Q
&	.	.s	3 j8 9 v||nJv||n5 6
%8;I I^C CR] ]F{B zLKFNOK	(O	
HI	(O	!,
01	HVH
	L(
)*	o.
/0	#K=
12	
AB	
,-	(O  !"%E 

)*	e6F6F6I6I5K[K[K^K^`e`p`p`s`s5t$uQVV$u vw
xy	(O 
/	
 !	(O[[036N	).)=Q
?@''/?(@&IY/BZ[NK	C0
12	(O uff+F 

)* ""%%""%%""%%""%%""%%""%%"--;;%%++%%++''//''//''//''//
  ii&&ii&&ii&&ii&&ii&&ii&&
 }}}}}}}}}}}}}}}}llllllllllll
1(GR BHH99	
;<{ B %vs   T