+
    Pj                     L   R t ^ RIt^ RIt^ RIHt ^ RIHt ^ RIHt R t	R R lt
 ! R R]P                  4      t ! R	 R
]P                  4      t ! R R]P                  4      t ! R R]P                  4      t ! R R]P                  4      t ! R R]P                  4      tR# )zk
Taken from
https://github.com/LTH14/mar/blob/fe470ac24afbee924668d8c5c83e9fec60af3a73/models/diffloss.py

N)Self)FlowLMConfigc                 0    V ^V,           ,          V,           #     )xshiftscales   &&&o/Users/ahmed/devFolder/Ultron/claude-voice/gateway/.venv/lib/python3.14/site-packages/pocket_tts/modules/mlp.pymodulater      s    E	?U""    c                d    V ^8  d   QhR\         P                  R\         P                  R\        /# )   r   alphaeps)torchTensorfloat)formats   "r   __annotate__r      s)       ell  r   c                    V P                  4       VP                  4       8  g   Q hV P                  pW P                  RRR7      ,           pWP                  V4      \        P
                  ! V4      ,          ,          P                  V4      pV# )r   Tdimkeepdim)r   dtypevartor   rsqrt)r   r   r   x_dtyper   ys   &&&   r   	_rms_normr"      sh    557eiik!!!ggG
"d+
+C	
hhsmekk#..	/33G<AHr   c                   T   a a ] tR t^t oRV3R lV 3R llltV3R lR ltRtVtV ;t# )RMSNormc                &   < V ^8  d   QhRS[ RS[/# )r   r   r   )intr   )r   __classdict__s   "r   r   RMSNorm.__annotate__   s     T TC Te Tr   c                   < \         SV `  4        W n        V3p\        P                  ! \
        P                  ! VR RR7      4      V n        R# )g      ?T)requires_gradN)super__init__r   nn	Parameterr   fullr   )selfr   r   alpha_shape	__class__s   &&& r   r,   RMSNorm.__init__   s7    f\\%**[#T"RS
r   c                4   < V ^8  d   QhRS[ P                  /# )r   r   r   r   )r   r'   s   "r   r   r(   #   s     2 2 2r   c                B    \        WP                  V P                  4      # N)r"   r   r   )r0   r   s   &&r   forwardRMSNorm.forward#   s    JJ11r   )r   r   )gh㈵>)	__name__
__module____qualname____firstlineno__r,   r8   __static_attributes____classdictcell____classcell__r2   r'   s   @@r   r$   r$      s      T T2 2 2r   r$   c                   @   a a ] tR t^'t oRtRV 3R lltR tRtVtV ;t	# )	LayerNormzJReimplementation of LayerNorm because the default one doesn't support jvp.c                   < \         SV `  4        W n        V'       da   \        P                  ! \
        P                  ! V4      4      V n        \        P                  ! \
        P                  ! V4      4      V n	        R # R # r7   )
r+   r,   r   r-   r.   r   onesweightzerosbias)r0   channelsr   elementwise_affiner2   s   &&&&r   r,   LayerNorm.__init__*   sM    ,,uzz(';<DKU[[%:;DI r   c                $   VP                  RRR7      pVP                  RRRR7      pW,
          \        P                  ! W0P                  ,           4      ,          p\        V R4      '       d$   WP                  ,          V P                  ,           pV# )r   Tr   F)r   unbiasedr   rF   r   )meanr   r   sqrtr   hasattrrF   rH   )r0   r   rN   r   s   &&  r   r8   LayerNorm.forward1   si    vv"dv+eeUDe9XC((N334""KK$))+Ar   )rH   r   rF   )ư>T
r:   r;   r<   r=   __doc__r,   r8   r>   r?   r@   rA   s   @@r   rC   rC   '   s     T< r   rC   c                   L   a a ] tR t^:t oRtRV3R lV 3R llltR tRtVtV ;t	# )TimestepEmbedderz4Embeds scalar timesteps into vector representations.c                ,   < V ^8  d   QhRS[ RS[ RS[ /# )r   hidden_sizefrequency_embedding_size
max_period)r&   )r   r'   s   "r   r   TimestepEmbedder.__annotate__=   s%     
 

:=
QT
r   c                  < \         SV `  4        \        P                  ! W!R R7      \        P                  ! 4       \        P                  ! WR R7      .pVP                  \        V4      4       \        P                  ! V!  V n        W n	        V^,          pV P                  R\        P                  ! \        P                  ! V4      ) \        P                  ! ^ VR7      ,          V,          4      4       R# )TrH   freqs)startendN)r+   r,   r-   LinearSiLUappendr$   
SequentialmlprY   register_bufferr   expmathlogarange)r0   rX   rY   rZ   blockshalfr2   s   &&&&  r   r,   TimestepEmbedder.__init__=   s     	II.$GGGIIIkT:

 	gk*+==&)(@%'1,UYY 44u||!QU7VVY]]^	
r   c                8   WP                   P                  VP                  4      ,          p\        P                  ! \        P
                  ! V4      \        P                  ! V4      .RR7      pV P                  ^,          '       d   Q hV P                  V4      pV# )r   r   r   )	r^   r   r   r   catcossinrY   re   )r0   targs	embeddingt_embs   &&   r   r8   TimestepEmbedder.forwardN   sj    ::==))IIuyy		$@bI	11A5566#r   )rY   re   )   i'  rS   rA   s   @@r   rV   rV   :   s     >
 
" r   rV   c                   <   a a ] tR t^Vt oRtV 3R ltR tRtVtV ;t	# )ResBlockzt
A residual block that can optionally change the number of channels.
:param channels: the number of input channels.
c           
       < \         SV `  4        Wn        \        VR R7      V n        \
        P                  ! \
        P                  ! WRR7      \
        P                  ! 4       \
        P                  ! WRR7      4      V n	        \
        P                  ! \
        P                  ! 4       \
        P                  ! V^V,          RR7      4      V n
        R# )rR   )r   Tr]   N)r+   r,   rI   rC   in_lnr-   rd   ra   rb   re   adaLN_modulation)r0   rI   r2   s   &&r   r,   ResBlock.__init__\   s     xT2
==IIht4GGIIIht4
 !#GGIryy1x<dC!
r   c                    V P                  V4      P                  ^RR7      w  r4p\        V P                  V4      W44      pV P	                  V4      pWV,          ,           # )   ro   r   )r}   chunkr   r|   re   )r0   r   r!   	shift_mlp	scale_mlpgate_mlphs   &&&    r   r8   ResBlock.forwardk   sU    )-)>)>q)A)G)Gr)G)R&	hTZZ]I9HHQKa<r   )r}   rI   r|   re   rS   rA   s   @@r   rz   rz   V   s     

   r   rz   c                   <   a a ] tR t^rt oRtV 3R ltR tRtVtV ;t	# )
FinalLayerz#
The final layer adopted from DiT.
c           	       < \         SV `  4        \        VR RR7      V n        \        P
                  ! WRR7      V n        \        P                  ! \        P                  ! 4       \        P
                  ! V^V,          RR7      4      V n	        R# )FrR   )rJ   r   Tr]   N)
r+   r,   rC   
norm_finalr-   ra   linearrd   rb   r}   )r0   model_channelsout_channelsr2   s   &&&r   r,   FinalLayer.__init__w   s_    #NuRVWii4H "GGIryy^1C$O!
r   c                    V P                  V4      P                  ^RR7      w  r4\        V P                  V4      W44      pV P	                  V4      pV# )r   ro   r   )r}   r   r   r   r   )r0   r   cr	   r
   s   &&&  r   r8   FinalLayer.forward   sK    ,,Q/55aR5@T__Q'6KKNr   )r}   r   r   rS   rA   s   @@r   r   r   r   s     
 r   r   c                   h   a a ] tR t^t oRtRV 3R llt]V3R lR l4       tV3R lR ltRt	Vt
V ;t# )	SimpleMLPAdaLNa[  Taken from https://arxiv.org/abs/2406.11838.

The MLP for Diffusion Loss.
:param in_channels: channels in the input Tensor.
:param model_channels: base channel count for the model.
:param out_channels: channels in the output Tensor.
:param cond_channels: channels in the condition.
:param num_res_blocks: number of residual blocks per downsample.
c                  < \         S
V `  4        Wn        W n        W0n        WPn        W`n        V^8w  g   Q h\        P                  ! \        V4       Uu. uF  p\        V4      NK  	  up4      V n        \        P                  ! WB4      V n        \        P                  ! W4      V n        . p\        V4       F  p	VP                  \!        V4      4       K  	  \        P                  ! V4      V n        \%        W#4      V n        R# u upi )r   N)r+   r,   in_channelsr   r   num_res_blocksnum_time_condsr-   
ModuleListrangerV   
time_embedra   
cond_embed
input_projrc   rz   
res_blocksr   final_layer)r0   r   r   r   cond_channelsr   r   _r   ir2   s   &&&&&&&   r   r,   SimpleMLPAdaLN.__init__   s     	&,(,,"""--7<^7LM7L!n-7LM
 ))MB))K@
~&Ah~67 ' --
3%nC Ns   Dc                2   < V ^8  d   QhRS[ RS[RS[RS[/# )r   cfg
latent_dimcond_dimreturn)r   r&   r   )r   r'   s   "r   r   SimpleMLPAdaLN.__annotate__   s+     
 
| 
 
PS 
X\ 
r   c           	     j    VP                   pVP                  pVP                  p^p\        W%W#WgR7      # )r   )r   )flowr   depthr   )clsr   r   r   configflow_dim
flow_depthr   s   &&&&    r   from_pydantic_config#SimpleMLPAdaLN.from_pydantic_config   s6    ::\\
*

 	
r   c          
         < V ^8  d   QhRS[ P                  RS[ P                  RS[ P                  RS[ P                  RS[ P                  /# )r   r   srs   r   r   r5   )r   r'   s   "r   r   r      sI     & &&"',,&38<<&DILL&	&r   c                  a a W#.oS P                  V4      p\        S4      S P                  8X  g!   Q RS P                   R\        S4       24       hS P                  ^8w  g   Q h\        V V3R l\	        S P                  4       4       4      S P                  ,          pS P                  V4      pWQ,           pS P                   F  pV! WF4      pK  	  S P                  WF4      # )z
Apply the model to an input batch.
:param c: conditioning from AR transformer.
:param s: start time tensor.
:param t: target time tensor.
:param x: an [N x C] Tensor of inputs.
:return: an [N x C] Tensor of outputs.
z	Expected z time conditions, got c              3   d   <"   T F%  pSP                   V,          ! SV,          4      x  K'  	  R # 5ir7   )r   ).0r   r0   tss   & r   	<genexpr>)SimpleMLPAdaLN.forward.<locals>.<genexpr>   s(     N3Ma"2a5))3Ms   -0)r   lenr   sumr   r   r   r   )	r0   r   r   rs   r   
t_combinedr!   blockr   s	   f&&&&   @r   r8   SimpleMLPAdaLN.forward   s     VOOA2w$--- 	
++,,B3r7)L	
- ""a'''N59L9L3MNNQUQdQdd 	 OOAN__EaA % %%r   )
r   r   r   r   r   r   r   r   r   r   r   )r:   r;   r<   r=   rT   r,   classmethodr   r8   r>   r?   r@   rA   s   @@r   r   r      s4     D@ 
 
& & &r   r   )rT   rh   r   torch.nnr-   typing_extensionsr   pocket_tts.utils.configr   r   r"   Moduler$   rC   rV   rz   r   r   r   r   r   <module>r      s       " 0#2bii 2		 &ryy 8 ryy  8 (Q&RYY Q&r   