
    -ij              '          d dl Z d dlmZ d dlmc mZ d dlZd dlmZm	Z	m
Z
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 G d dej                        Z	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dZdede j$                  dededededededede	e   dedede	e   dedededed e
eef   f$d!Zed"k(  r e j2                  d#       d$Zd%Zd&Z e j:                  d d'e j<                  z  e      Zd( e j@                  e      z  d) e j@                  d*ez        z  z   d+ e j@                  d,ez        z  z   d- e jB                  e      z  z   Z"e"jG                  d       jG                  d.      jI                  ed.d.      Z% e&d/        e&d0        e&d/        eed1d2d34      Z' e&d5        ee'e%d6d7dd8      Z( e&d9        e jR                         5   e'e%dd&       \  Z*Z+ddd       eez  Z, e-d: e'j]                         D              Z/ e&d;d/         e&d<        e&d=e, d>        e&d?e/ d@        e&dAe,e/z  dBdC        e&dD        e&dEe+dF   d    ja                  d3G               e&dHe+dI   d    ja                  d3G               e&dJe+dK   d    ja                  d3G               e&d;d/         e&dL        e&d/        eed%d&d3d*d,M      Z1 e1e%      Z2 e&dNe%jf                           e&dOe2jf                          e1ji                         Z5 e&dP        e&dQe5dR            e&dSe5dT            e&dAe5dU   dBdC        e&dV        e jR                         5   e1e%dW      \  Z2Z6ddd        e&dXe6jf                           e&dY       yy# 1 sw Y   xY w# 1 sw Y   5xY w)[    N)TupleOptionalDictListc            
       P    e Zd ZdZ	 	 	 ddededef fdZdej                  fdZ
dej                  fdZdej                  fd	Zd
ej                  dej                  fdZdej                  dej                  dej                  d
ej                  dej                  f
dZ xZS )ParametricFourierBasisu   
    Learnable parametric Fourier basis functions.
    Instead of storing coefficients explicitly, we learn parameters
    that generate the Fourier coefficients on demand.
    
    F(t) = Σ θ_k * φ_k(t) where φ_k are basis functions
    n_harmonicslearnable_frequenciesinitial_freq_scalec                 J   t         |           || _        t        j                  t        j                  |      |z        | _        t        j                  t        j                  d|dz         j                               | _
        t        j                  t        j                  |      dz        | _        t        j                  t        j                  |            | _        t        j                  t        j                  |      dz  t        j                   z        | _        y )N   皙?   )super__init__r	   nn	Parametertorchones	freq_basearangefloatfreq_multiplierrandnamp_base	amp_scalerandmathpi
phase_base)selfr	   r
   r   	__class__s       0/home/per/Documents/python/linearize/pfourier.pyr   zParametricFourierBasis.__init__   s     	& JJ{#&88
  "||LLK!O,224 

 U[[%=%CDejj&=> ,,JJ{#a'$''1
    returnc                 4    | j                   | j                  z  S )z4Get frequencies from parameters - O(n) not infinite.)r   r   r!   s    r#   get_frequenciesz&ParametricFourierBasis.get_frequencies*   s    ~~ 4 444r$   c                 Z    t        j                  | j                        | j                  z  S )z3Get amplitudes from parameters - O(n) not infinite.)Fsoftplusr   r   r'   s    r#   get_amplitudesz%ParametricFourierBasis.get_amplitudes.   s    zz$--(4>>99r$   c                 b    t        j                  | j                        t        j                  z  S )zGet phases from parameters.)r   tanhr    r   r   r'   s    r#   
get_phasesz!ParametricFourierBasis.get_phases2   s    zz$//*TWW44r$   tc                 h   |j                         dk(  r|j                  d      }|j                  d   }|j                  d   }| j                         }| j	                         }| j                         }|j                  d      }|t        j                  ||z  |z         z  }|j                  d      S )z
        Generate Fourier series value at time t.
        
        Args:
            t: Time points, shape (batch, seq_len) or (seq_len,)
            
        Returns:
            Fourier series values, same shape as t
        r   r   dim)	r4   	unsqueezeshaper(   r,   r/   r   sinsum)	r!   r0   
batch_sizeseq_lenomegaAphi
t_expanded	harmonicss	            r#   forwardzParametricFourierBasis.forward6   s     557a<AAWWQZ
''!* $$&!oo [[_
		J$
 
	
 }}}$$r$   frequencies
amplitudesphasesc                    |j                         dk(  r|j                  d      }|j                         dk(  r|j                  d      }|j                         dk(  r|j                  d      }|j                         dk(  r|j                  d      }|j                  d   dk(  ro|j                  d   dkD  r]|j                  |j                  d   d      }|j                  |j                  d   d      }|j                  |j                  d   d      }|j                  d      }|j                  d      }|j                  d      }|j                  d      }|t	        j
                  ||z  |z         z  j                  d      S )a  
        Reconstruct a time-domain signal from explicit Fourier parameters.

        This is the inverse transform for this parameterization, similar in
        spirit to an inverse FFT, but using learned continuous frequencies
        instead of discrete FFT bins.

        Args:
            frequencies: Frequency parameters, shape (n_harmonics,) or (batch, n_harmonics)
            amplitudes: Amplitude parameters, same shape as frequencies
            phases: Phase parameters, same shape as frequencies
            t: Time points, shape (seq_len,) or (batch, seq_len)

        Returns:
            Reconstructed time-domain signal, shape (batch, seq_len)
        r   r   r2   r3   )r4   r5   r6   expandr   r7   r8   )	r!   rA   rB   rC   r0   r>   freq_expandedamp_expandedphase_expandeds	            r#   inversezParametricFourierBasis.inverseU   sT   . 557a<AA??!%//2K>>q #--a0J::<1%%a(FQ1$a%,,QWWQZ<K#**1771:r:J]]1771:r2F[[_
#--a0!++A.))!,uyyJ&7 
 
323;	r$   )
   T      ?)__name__
__module____qualname____doc__intboolr   r   r   Tensorr(   r,   r/   r@   rI   __classcell__r"   s   @r#   r   r      s     &*$'	

  $
 "	
45 5: :5ELL 5% %%,, %>-\\- LL- 	-
 <<- 
-r$   r   c                        e Zd ZdZddededef fdZdej                  fdZ	 ddej                  d	e	dej                  fd
Z
edefd       Z xZS )LowRankFourierLayeru   
    Low-rank parametric Fourier layer.
    W_ij = Σ_k u_ik * v_jk (rank-r factorization of weight matrix)
    
    Compression: O(mn) → O(r(m+n))
    	input_dim
output_dimrankc                 j   t         |           || _        || _        || _        t        j                  t        j                  ||      dz        | _	        t        j                  t        j                  ||      dz        | _
        t        j                  t        j                  |            | _        y )N{Gz?)r   r   rW   rX   rY   r   r   r   r   UVr   basis_scale)r!   rW   rX   rY   r"   s       r#   r   zLowRankFourierLayer.__init__   s    "$	 ekk*d;dBCekk)T:TAB <<

4(89r$   r%   c                     | j                   | j                  j                  z  }|| j                  j	                  d      z  S )z
        Reconstruct full weight matrix from parametric factors.
        Only computed when needed (on-demand generation).
        r   )r\   r]   Tr^   r5   )r!   Ws     r#   get_weight_matrixz%LowRankFourierLayer.get_weight_matrix   s6     FFTVVXX4##--a000r$   xgenerate_weightsc                     |r0| j                         }t        j                  ||j                        S t        j                  t        j                  || j                        | j
                  j                        S )a  
        Forward pass with optional on-demand weight generation.
        
        Args:
            x: Input tensor (batch, seq, input_dim)
            generate_weights: If True, reconstructs full matrix
                            If False, uses factorized form directly
        )rb   r   matmulr`   r]   r\   )r!   rc   rd   ra   s       r#   r@   zLowRankFourierLayer.forward   sZ     &&(A<<133'' <<Q' r$   c                     | j                   | j                  z  }| j                  | j                   | j                  z   z  | j                  z   }||z  S )z6How much this layer is compressed vs explicit storage.)rW   rX   rY   )r!   explicit
parametrics      r#   compression_ratioz%LowRankFourierLayer.compression_ratio   sE     >>DOO3YY$..4??"BCdiiO
*$$r$   )   F)rL   rM   rN   rO   rP   r   r   rR   rb   rQ   r@   propertyr   rj   rS   rT   s   @r#   rV   rV      s{    :# :3 :c :15<< 1 "'<<  
	0 %5 % %r$   rV   c            	            e Zd ZdZ	 	 ddedededef fdZdej                  dej                  d	ej                  fd
Z xZ	S )FourierGeneratorLayeru   
    Neural network generator that produces Fourier-based weights.
    
    W_ij = NN_θ(i, j) — a small network generates weight values
    from their indices in the weight matrix.
    
    Compression: O(mn) → O(θ) where θ << mn
    rW   rX   
hidden_dimn_fourier_termsc           
         t         |           || _        || _        || _        | j                  dt        j                  d|dz         j                                | j                  dt        j                  d|dz         j                                t        j                  t        j                  ||      dz        | _        t        j                  t        j                  d|      t        j                         t        j                  ||      t        j                         t        j                  |d            | _        y )Nfreq_ir   freq_jr[   r   )r   r   rW   rX   rq   register_bufferr   r   r   r   r   r   fourier_coeffs
SequentialLinearTanhneural_correction)r!   rW   rX   rp   rq   r"   s        r#   r   zFourierGeneratorLayer.__init__   s     	"$. 	LLOa/0668	
 	LLOa/0668	
 !llKK9D@

 "$IIa$GGIIIj*-GGIIIj!$"
r$   	indices_i	indices_jr%   c                    d|j                         z  | j                  z  dz
  }d|j                         z  | j                  z  dz
  }|j                  d      }|j                  d      }t	        j
                  || j                  j                  d      z  t        j                  z        }t	        j                  || j                  j                  d      z  t        j                  z        }t	        j                  d|| j                  |      }	t	        j                  ||gd      }
| j                  |
      j                  d      }|	|z   S )a  
        Generate weight values for given matrix indices.
        
        Args:
            indices_i: Row indices (output dimension)
            indices_j: Column indices (input dimension)
            
        Returns:
            Weight values for each (i,j) pair
        r   r   r2   r   z...k,kl,...l->...r3   )r   rX   rW   r5   r   r7   rs   r   r   cosrt   einsumrv   stackrz   squeeze)r!   r{   r|   i_normj_norm
i_expanded
j_expandedi_freqj_freqfourier_outindex_featuresrz   s               r#   r@   zFourierGeneratorLayer.forward   s!    Y__&&81<Y__&&7!; %%b)
%%b)
 ..q11DGG;
 ..q11DGG;

 ll#6@S@SU[\ ff%52> 22>BJJ2N...r$   )       )
rL   rM   rN   rO   rP   r   r   rR   r@   rS   rT   s   @r#   ro   ro      sb      #
#
 #
 	#

 #
J"/ "/%,, "/5<< "/r$   ro   c                        e Zd ZdZ	 	 	 ddedededededef fdZ	 dd	ej                  d
edej                  fdZ	de
eef   fdZ xZS )ParametricTimeSeriesModelz
    Complete parametric time series model using learned Fourier series.
    
    Instead of storing explicit weights, uses parametric functions
    to generate weights on demand.
    rW   rp   rX   n_layersrY   rq   c           
      <   t         |           || _        || _        || _        || _        |dz   | _        t        |      | _        t        j                  | j                  |      | _        t        j                  t        |      D cg c]  }t        |||       c}      | _        t!        |||      | _        t        j$                  t'        j(                  |            | _        t        j$                  t'        j(                  |      dz        | _        y c c}w )Nr   )rY   )rq   g?)r   r   rW   rp   rX   r   input_with_fourier_dimr   input_basisr   rx   input_projection
ModuleListrangerV   hidden_layersro   output_projectionr   r   r   ode_gainode_damping)	r!   rW   rp   rX   r   rY   rq   ir"   s	           r#   r   z"ParametricTimeSeriesModel.__init__   s     	"$$ &/!m# 2/B "		$*E*Ez R  ]] 8_,
   ,
  "7
O"

 UZZ%9:<<

8(<s(BC!,
s   Drc   return_generated_weightsr%   c                    |j                   \  }}}t        j                  dd||j                        }| j	                  |      j                  |d      }t        j                  ||j                  d      gd      }| j                  |      }t        | j                        D ]L  \  }	}
 |
|d      }|| j                  |	   || j                  |	   |z  z
  z  z   }t        j                  |      }N t        j                  ||j                        j!                  d|d      j#                  ||| j$                        }t        j                  | j$                  |j                        j!                  dd| j$                        j#                  ||| j$                        }| j'                  ||      }||z  j)                  dd	      }|r||fS |S )
a?  
        Forward pass with parametric weight generation.
        
        Args:
            x: Input time series (batch, seq_len, input_dim)
            return_generated_weights: If True, return generated weights too
            
        Returns:
            Output predictions and optionally generated weights
        r   r   devicer2   r3   F)rd   T)r4   keepdim)r6   r   linspacer   r   repeatcatr5   r   	enumerater   r   r   r*   gelur   viewrE   rp   r   r8   )r!   rc   r   r9   r:   _r0   fourier_featuresh	layer_idxlayerh_newseq_idx
hidden_idxoutput_weightsoutputs                   r#   r@   z!ParametricTimeSeriesModel.forwardG  s    "#
GQ NN1a:++A.55j!D IIq*44R89rB!!!$ !*$*<*< = 		Iu!e4E DMM),((3a77 A q	A		 ,,wqxx8==a!LSS
 \\$//!((CHHAt_ff


 //D n$))b$)?#>))r$   c                    | j                   | j                  z  | j                  z  }t        d | j                  D              }|| j
                  j                  dz  z  }|| j                  j                  j                         z  }|| j                  j                  j                         j                         z  }|||t        |d      z  dS )z@Calculate storage cost of parametric vs explicit representation.c              3      K   | ]T  }|j                   j                         |j                  j                         z   |j                  j                         z    V y wN)r\   numelr]   r^   ).0r   s     r#   	<genexpr>z=ParametricTimeSeriesModel.get_storage_cost.<locals>.<genexpr>  sE       
 GGMMOeggmmo-0A0A0G0G0II 
s   AA   r   )rh   ri   rj   )rW   rp   r   r8   r   r   r	   r   rv   r   rz   
state_dict__len__max)r!   explicit_paramsparametric_paramss      r#   get_storage_costz*ParametricTimeSeriesModel.get_storage_cost~  s    ..4??:T]]J  
++ 
 
 	T--99A==T33BBHHJJT33EEPPRZZ\\ (+!037H!3L!L
 	
r$   )   r   rJ   rl   )rL   rM   rN   rO   rP   r   r   rR   rQ   r@   r   strr   rS   rT   s   @r#   r   r     s     !%D%D %D 	%D
 %D %D %DT */5<<5 #'5 
	5n
$sCx. 
r$   r   c            	           e Zd ZdZ	 	 	 ddedededef fdZdej                  dee	ej                  f   fd	Z
d
ee	ej                  f   dej                  dej                  fdZdd
ee	ej                  f   deej                     dej                  fdZddej                  deej                     deej                  ef   fdZ xZS )AdaptiveFourierEncoderu  
    Encodes arbitrary time series into parametric Fourier representation.
    
    Given a time series x(t), finds parameters θ such that:
    x(t) ≈ Σ_k θ_k * φ_k(t)
    
    The encoder learns to compress arbitrary inputs into fixed-size parameter vectors.
    rW   
latent_dimn_fourier_basisencoder_layersc                 J   t         |           || _        || _        || _        t        j                  |||d      | _        t        |      | _	        t        j                  ||      | _        t        j                  ||      | _        t        j                  ||      | _        y )NT)
input_sizehidden_size
num_layersbatch_first)r   r   rW   r   r   r   GRUencoderr   fourier_basisrx   latent_to_freqlatent_to_amplatent_to_phase)r!   rW   r   r   r   r"   s        r#   r   zAdaptiveFourierEncoder.__init__  s     	"$. vv "%	
 4OD !ii
ODYYz?C!yy_Er$   rc   r%   c                 0   | j                  |      \  }}|d   }|t        j                  | j                  |            t	        j
                  | j                  |            t	        j
                  | j                  |            t        j                  z  dS )z
        Encode time series into parametric representation.
        
        Args:
            x: Time series (batch, seq_len, input_dim)
            
        Returns:
            Dictionary of parametric components
        r2   )latentrA   rB   rC   )
r   r*   r+   r   r   r.   r   r   r   r   )r!   rc   r   h_nzs        r#   encodezAdaptiveFourierEncoder.encode  s{     a3G ::d&9&9!&<=**T%7%7%:;jj!5!5a!89DGGC	
 	
r$   paramsr0   c                 P    | j                   j                  |d   |d   |d   |      S )a	  
        Decode parametric representation back to time series.
        
        Args:
            params: Dictionary from encode()
            t: Time points (seq_len,) or (batch, seq_len)
            
        Returns:
            Reconstructed time series
        rA   rB   rC   )r   rI   )r!   r   r0   s      r#   decodezAdaptiveFourierEncoder.decode  s8     !!))=!< 8	
 	
r$   c                 0   ||d   j                         dkD  r|d   j                  d   nd}t        | dd      }t        j                  dd||d   j
                        }|dkD  r!|j                  d      j                  |d      }| j                  ||      S )z
        Alias for decode(), matching the inverse-transform naming you asked for.

        If `t` is omitted, a default evenly spaced grid over [0, 1] is used.
        rA   r   r   _default_inverse_stepsd   r   r2   )	r4   r6   getattrr   r   r   r5   rE   r   )r!   r   r0   batchn_stepss        r#   rI   zAdaptiveFourierEncoder.inverse  s     96<]6K6O6O6QTU6UF=)//2[\Ed$<cBGq!WVM5J5Q5QRAqyKKN))%4{{61%%r$   c                     |Tt        j                  dd|j                  d   |j                        }|j	                         dkD  r|j                  d      }| j                  |      }| j                  ||      }||fS )z
        Full encode-decode cycle.
        
        Args:
            x: Input time series
            t: Time points (if None, uses linspace)
            
        Returns:
            Reconstructed time series and parameters
        r   r   r   r   )r   r   r6   r   r4   r5   r   r   )r!   rc   r0   r   reconstructions        r#   r@   zAdaptiveFourierEncoder.forward  sk     9q!QWWQZAAuuw{KKNQVQ/v%%r$   )r      r   r   )rL   rM   rN   rO   rP   r   r   rR   r   r   r   r   r   rI   r   r@   rS   rT   s   @r#   r   r     s	    !FF F 	F
 F8
 
c5<<.?)@ 
,
T#u||"34 
 
%,, 
$&d3#45 &(5<<:P &\a\h\h && &(5<<*@ &ERWR^R^`dRdLe &r$   r   -C6?Tmodeltime_seriesepochslr
lambda_reglambda_freq_reglambda_amp_reglambda_phase_reg	grad_clippatience	min_deltause_inversescheduler_patiencescheduler_factorscheduler_min_lrscheduler_onverboser%   c                 2
   t         j                  j                  | j                         |      }|t	        d|dz        }t         j                  j
                  j                  |d|||      }g g g g g g g d}t        j                  dd	|j                  d	   |j                  
      }|j                         dkD  r|j                  d      }|j                  d   d	k(  r|d   n|j                  d      }t        d      }d}d}t        |      D ]  }|j                           | ||      \  }}|r| j!                  ||      n| j#                  ||      }|j                         dk(  r#|j                  d   d	k(  r|j                  d      }t%        j&                  ||      }t)        d | j                         D              }|d   j+                         j-                         }|d   j+                         j-                         } |d   j+                         j-                         }!|||z  z   ||z  z   || z  z   ||!z  z   }"|"j/                          t         j0                  j2                  j5                  | j                         |       |j7                          |dk(  r|j9                         n|"j9                         }#|j7                  |#       |d   j;                  |"j9                                |d   j;                  |j9                                |d   j;                  |j9                                |d   j;                  |j9                                |d   j;                  | j9                                |d   j;                  |!j9                                |d   j;                  |j<                  d   d          |"j9                         |
z   |k  r`|"j9                         }| j?                         jA                         D $%ci c]$  \  }$}%|$|%jC                         jE                         & }}$}%d}n|d	z  }|r|dz  dk(  s||d	z
  k(  rtG        d| d|"j9                         dd |j9                         dd!|j9                         dd"|j9                         dd#| j9                         dd$|!j9                         dd%|#dd&|j<                  d   d   d'       |	||	k\  s|rtG        d(| d)|dd*        n || jI                  |       |S c c}%}$w )+a  
    Train the parametric Fourier encoder on a time series.
    
    This implements the retro-compression idea: given explicit data,
    find the parametric function that best approximates it.
    
    Args:
        model: AdaptiveFourierEncoder to train
        time_series: Training data (batch, seq_len, input_dim)
        epochs: Number of training epochs
        lr: Learning rate
        lambda_reg: L2 regularization strength
        lambda_freq_reg: Penalty on learned frequencies
        lambda_amp_reg: Penalty on learned amplitudes
        lambda_phase_reg: Penalty on learned phases
        grad_clip: Gradient clipping threshold
        patience: Early stopping patience. Disable if None.
        min_delta: Minimum improvement required to reset patience
        use_inverse: If True, train through model.inverse(params, t)
        scheduler_patience: Plateau scheduler patience. Defaults to a conservative value.
        scheduler_factor: LR decay factor for plateau scheduler
        scheduler_min_lr: Minimum learning rate
        scheduler_on: Metric used for LR scheduling: "loss" or "reconstruction"
        verbose: Print progress
        
    Returns:
        Training history dictionary
    )r   N   rJ   min)modefactorr   min_lr)lossreconstruction_error
param_costfrequency_costamplitude_cost
phase_costr   r   r   r   r   r2   ).r   infr   c              3   X   K   | ]"  }|j                         j                          $ y wr   )squaremeanr   ps     r#   r   z+train_parametric_fourier.<locals>.<genexpr>^  s     Gq*Gs   (*rA   rB   rC   r   r   r   r   r   r   r   r   r   zEpoch z: Loss=z.4fz, Recon=z, Param=z, Freq=z, Amp=z, Phase=z, Sched=z, LR=z.2ezEarly stopping at epoch z (best loss z).)%r   optimAdamW
parametersr   lr_schedulerReduceLROnPlateaur   r6   r   r4   r5   r   r   r   	zero_gradrI   r   r*   mse_lossr8   r   r   backwardr   utilsclip_grad_norm_stepitemappendparam_groupsr   itemsdetachcloneprintload_state_dict)&r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   	optimizer	schedulerhistoryr0   target	best_loss
best_state
bad_epochsepochr   r   r   
recon_lossr   r   r   r   r   scheduler_metrickvs&                                         r#   train_parametric_fourierr"    s   ^ !!%"2"2"4!<I! Vr\2((::# ; I  "G 	q![..q1+:L:LMA1KKN$/$5$5b$9Q$>[ KDWDWXZD[FeIJJv A +q)	65@vq1ellSY[\F]1$)=)=b)AQ)F+33B7N ZZ7
 GE4D4D4FGG
.557<<>-446;;=H%,,.335
 :%&./ ~-. +	, 	 	&&u'7'7'99E0<@P0P:??,VZV_V_Va'( 	tyy{+&'..z/@A$$Z__%67 !(()<)<)>? !(()<)<)>?$$Z__%67Y33A6t<=99;"Y.		I<A<L<L<N<T<T<VWDAq!QXXZ--//WJWJ!OJq(EVaZ,?wtyy{3&7 8#*3/ 0#*3/ 0&++-c2 3%**,S1 2#*3/ 0)#. /,,Q/5c:<	 J($:0|Ic?RTUVCAF j)N3 Xs   )T__main__*   i'  r   r      g       @g      ?rk   g      ?rJ   r   r2   z<============================================================z4PARAMETRIC FOURIER SERIES - LEARNED FROM TIME SERIESr   r   r   )rW   r   r   r   z
[Training parametric model...]i  r[   )r   r   r   r   z
[Evaluating...]c              #   <   K   | ]  }|j                           y wr   )r   r  s     r#   r   r     s     B!AGGIBs   
zCOMPRESSION RESULTS:z  Explicit storage: z valuesz  Parametric storage: z parametersz  Compression ratio: z.1frc   z
Learned Fourier parameters:z  Frequencies: rA   )decimalsz  Amplitudes:  rB   z  Phases:      rC   z"FULL PARAMETRIC TIME SERIES MODEL:)rW   rp   rX   r   rY   rq   zInput shape: zOutput shape: z
Storage analysis:z  Explicit parameters: rh   z  Parametric parameters: ri   rj   z)
[Testing on-demand weight generation...])r   zGenerated weights shape: z*(Weights generated on-the-fly, not stored))i  gMbP?r   r   r   r   rK   Ngư>TNg?gh㈵>r   T)7r   torch.nnr   torch.nn.functional
functionalr*   r   typingr   r   r   r   Moduler   rV   ro   r   r   rR   rP   r   rQ   r   r"  rL   manual_seedr:   r9   rW   r   r   r0   r7   r   rc   r5   rE   r   r  r   r  no_gradr   r   r   r8   r  r   round
full_modelr   r6   r   storageweights r$   r#   <module>r5     s'        . .{RYY {|<%")) <%~Q/BII Q/hv
		 v
ro&RYY o&j ! ""(,!"(#T!TT T 		T
 T T T T T smT T T !T T T  !T" #T$ 
#t)_%Tv zEbGJI 	q!ehh,0AieiilieiiA	yuyya  	! 	kekk'""	#  ++a.**2.55j"bIK	(O	
@A	(O #	E 

,-&G 

	 8!&{2A!78 	)OBu/?/?/ABB	Bvh-	
 !	  1
9:	"#4"5[
AB	!/2C"CC!H
JK	)+	OF=1!4::A:FG
HI	OF<039919EF
GH	OF8,Q/55q5AB
CD 
Bvh-	
./	(O*J $F	M+++,
-.	N6<<.
)* ))+G	!	#GJ$7#8
9:	%gl&;%<
=>	!'*=">s!C1
EF 
68	 Q$[4PQ	%gmm_
56	68A ^8 8\Q Qs   :OOOO