+
    Lj                       ^ RI HtHt ^ RIt^ RIHt ^ RIHu Ht ^RI	H
t
Ht ^RIHtHtHt ^RIHt ^RIHtHtHtHtHtHt ^RIHtHtHt ^RIHt ^R	IH t H!t!H"t"H#t#H$t$ ]! 4       '       d   ^ RI%t&MRt&]PN                  ! ](4      t) ! R
 R4      t* ! R R4      t+R R lt,] ! R R]PZ                  4      4       t.] ! R R]PZ                  4      4       t/] ! R R]PZ                  4      4       t0 ! R R]PZ                  4      t1] ! R R]PZ                  4      4       t2 ! R R]PZ                  4      t3] ! R R]PZ                  4      4       t4 ! R R]PZ                  4      t5R# )     )AnyCallableN)	deprecatelogging)is_torch_npu_availableis_torch_xla_availableis_xformers_available)maybe_allow_in_graph)GEGLUGELUApproximateGELUFP32SiLULinearActivationSwiGLU)	AttentionAttentionProcessorJointAttnProcessor2_0)SinusoidalPositionalEmbedding)AdaLayerNormAdaLayerNormContinuousAdaLayerNormZeroRMSNormSD35AdaLayerNormZeroXc                   Z   a  ] tR t^'t o ]V 3R lR l4       tV 3R lR ltR tR tRt	V t
R# )	AttentionMixinc                6   < V ^8  d   QhRS[ S[S[3,          /# )   return)dictstrr   )format__classdict__s   "F/app/.local/lib/python3.14/site-packages/diffusers/models/attention.py__annotate__AttentionMixin.__annotate__)   s      c+=&=!>     c                b   a / pR V3R lloV P                  4        F  w  r#S! W#V4       K  	  V# )z
Returns:
    `dict` of attention processors: A dictionary containing all attention processors used in the model with
    indexed by its weight name.
c                    V ^8  d   QhR\         R\        P                  P                  R\        \         \
        3,          /# )r   namemodule
processors)r    torchnnModuler   r   )r!   s   "r#   r$   4AttentionMixin.attn_processors.<locals>.__annotate__2   s6     	 	c 	588?? 	X\]`bt]tXu 	r&   c                    < \        VR 4      '       d   VP                  4       W  R2&   VP                  4        F  w  r4S! V  RV 2WB4       K  	  V# )get_processor
.processor.)hasattrr1   named_children)r)   r*   r+   sub_namechildfn_recursive_add_processorss   &&&  r#   r8   CAttentionMixin.attn_processors.<locals>.fn_recursive_add_processors2   sZ    v//282F2F2H
V:./#)#8#8#:+tfAhZ,@%T $; r&   )r5   )selfr+   r)   r*   r8   s   &   @r#   attn_processorsAttentionMixin.attn_processors(   s=     
	 	 !//1LD'jA 2 r&   c                F   < V ^8  d   QhRS[ S[S[S[ 3,          ,          /# )r   	processor)r   r   r    )r!   r"   s   "r#   r$   r%   @   s(      A  A,>cK]F]A^,^  Ar&   c           	     ,  a \        V P                  P                  4       4      p\        V\        4      '       d/   \        V4      V8w  d   \        R\        V4       RV RV R24      hR V3R lloV P                  4        F  w  r4S! W4V4       K  	  R# )a  
Sets the attention processor to use to compute attention.

Parameters:
    processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
        The instantiated processor class or a dictionary of processor classes that will be set as the processor
        for **all** `Attention` layers.

        If `processor` is a dict, the key needs to define the path to the corresponding cross attention
        processor. This is strongly recommended when setting trainable attention processors.

z>A dict of processors was passed, but the number of processors z0 does not match the number of attention layers: z. Please make sure to pass z processor classes.c                X    V ^8  d   QhR\         R\        P                  P                  /# )r   r)   r*   )r    r,   r-   r.   )r!   s   "r#   r$   7AttentionMixin.set_attn_processor.<locals>.__annotate__U   s&     	T 	Tc 	T588?? 	Tr&   c                   < \        VR 4      '       dL   \        V\        4      '       g   VP                  V4       M#VP                  VP	                  V  R24      4       VP                  4        F  w  r4S! V  RV 2WB4       K  	  R# )set_processorr2   r3   N)r4   
isinstancer   rC   popr5   )r)   r*   r>   r6   r7   fn_recursive_attn_processors   &&&  r#   rF   FAttentionMixin.set_attn_processor.<locals>.fn_recursive_attn_processorU   ss    v//!)T22((3(($z7J)KL#)#8#8#:+tfAhZ,@%S $;r&   N)lenr;   keysrD   r   
ValueErrorr5   )r:   r>   countr)   r*   rF   s   &&   @r#   set_attn_processor!AttentionMixin.set_attn_processor@   s     D((--/0i&&3y>U+BPQTU^Q_P` a005w6QRWQXXkm 
	T 	T !//1LD'i@ 2r&   c                P   V P                   P                  4        F4  w  rR\        VP                  P                  4      9   g   K+  \        R4      h	  V P                  4        F?  p\        V\        4      '       g   K  VP                  '       g   K/  VP                  4        KA  	  R# )z
Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value)
are fused. For cross-attention modules, key and value projection matrices are fused.
AddedzQ`fuse_qkv_projections()` is not supported for models having added KV projections.N)r;   itemsr    	__class____name__rJ   modulesrD   AttentionModuleMixin_supports_qkv_fusionfuse_projections)r:   _attn_processorr*   s   &   r#   fuse_qkv_projections#AttentionMixin.fuse_qkv_projectionsb   s|    
 "&!5!5!;!;!=A#n66??@@ !tuu "> llnF&"677F<W<W<W'') %r&   c                    V P                  4        F?  p\        V\        4      '       g   K  VP                  '       g   K/  VP	                  4        KA  	  R# )u]   Disables the fused QKV projection if enabled.

> [!WARNING] > This API is 🧪 experimental.
N)rS   rD   rT   rU   unfuse_projections)r:   r*   s   & r#   unfuse_qkv_projections%AttentionMixin.unfuse_qkv_projectionso   s9    
 llnF&"677F<W<W<W))+ %r&    N)rR   
__module____qualname____firstlineno__propertyr;   rL   rY   r]   __static_attributes____classdictcell__r"   s   @r#   r   r   '   s3      . A  AD*, ,r&   r   c                   |  a  ] tR t^yt o Rt. tRtRtV 3R lR ltRV 3R lR llt	V 3R lR	 lt
V 3R
 lR ltR V 3R lR lltR!V 3R lR llt]P                  ! 4       R 4       t]P                  ! 4       R 4       tV 3R lR ltV 3R lR ltR"V 3R lR lltR!V 3R lR lltR"V 3R lR lltV 3R lR ltRtV tR# )#rT   NTFc                $   < V ^8  d   QhRS[ RR/# )r   r>   r   N)r   )r!   r"   s   "r#   r$   !AttentionModuleMixin.__annotate__   s     # #'9 #d #r&   c                r   \        V R4      '       d   \        V P                  \        P                  P
                  4      '       dk   \        V\        P                  P
                  4      '       gA   \        P                  RV P                   RV 24       V P                  P                  R4       Wn        R# )zu
Set the attention processor to use.

Args:
    processor (`AttnProcessor`):
        The attention processor to use.
r>   z-You are removing possibly trained weights of z with N)
r4   rD   r>   r,   r-   r.   loggerinfo_modulesrE   )r:   r>   s   &&r#   rC   "AttentionModuleMixin.set_processor   sx     D+&&4>>588??;;y%((//::KKGGWW]^g]hijMMk*"r&   c                $   < V ^8  d   QhRS[ RR/# )r   return_deprecated_lorar   r   bool)r!   r"   s   "r#   r$   ri      s     " "D "EY "r&   c                .    V'       g   V P                   # R# )z
Get the attention processor in use.

Args:
    return_deprecated_lora (`bool`, *optional*, defaults to `False`):
        Set to `True` to return the deprecated LoRA attention processor.

Returns:
    "AttentionProcessor": The attention processor in use.
N)r>   )r:   rp   s   &&r#   r1   "AttentionModuleMixin.get_processor   s     &>>! &r&   c                    < V ^8  d   QhRS[ /# )r   backend)r    )r!   r"   s   "r#   r$   ri      s     4 4S 4r&   c                $   ^RI Hp VP                  P                  4        Uu0 uF  q3P                  kK  	  ppW9  d'   \        RV: R2RP                  V4      ,           4      hV! VP                  4       4      pWP                  n	        R# u upi )   )AttentionBackendNamez	`backend=z ` must be one of the following: z, N)
attention_dispatchry   __members__valuesvaluerJ   joinlowerr>   _attention_backend)r:   rv   ry   xavailable_backendss   &&   r#   set_attention_backend*AttentionModuleMixin.set_attention_backend   sw    </C/O/O/V/V/XY/X!gg/XY,z
*JKdiiXjNkkll&w}}7,3) Zs   Bc                $   < V ^8  d   QhRS[ RR/# )r   use_npu_flash_attentionr   Nrq   )r!   r"   s   "r#   r$   ri      s     2 24 2D 2r&   c                n    V'       d   \        4       '       g   \        R4      hV P                  R4       R# )z
Set whether to use NPU flash attention from `torch_npu` or not.

Args:
    use_npu_flash_attention (`bool`): Whether to use NPU flash attention or not.
ztorch_npu is not available_native_npuN)r   ImportErrorr   )r:   r   s   &&r#   set_use_npu_flash_attention0AttentionModuleMixin.set_use_npu_flash_attention   s*     #)++!">??""=1r&   c                Z   < V ^8  d   QhRS[ RS[S[R,          R3,          R,          RR/# )r   use_xla_flash_attentionpartition_specN.r   )rr   tupler    )r!   r"   s   "r#   r$   ri      s;     2 2!%2 cDj#o.52
 
2r&   c                n    V'       d   \        4       '       g   \        R4      hV P                  R4       R# )a  
Set whether to use XLA flash attention from `torch_xla` or not.

Args:
    use_xla_flash_attention (`bool`):
        Whether to use pallas flash attention kernel from `torch_xla` or not.
    partition_spec (`tuple[]`, *optional*):
        Specify the partition specification if using SPMD. Otherwise None.
    is_flux (`bool`, *optional*, defaults to `False`):
        Whether the model is a Flux model.
ztorch_xla is not available_native_xlaN)r   r   r   )r:   r   r   is_fluxs   &&&&r#   set_use_xla_flash_attention0AttentionModuleMixin.set_use_xla_flash_attention   s*    " #)++!">??""=1r&   c                8   < V ^8  d   QhRS[ RS[R,          RR/# )r   'use_memory_efficient_attention_xformersattention_opNr   )rr   r   )r!   r"   s   "r#   r$   ri      s*     %7 %77;%7KSVZ?%7	%7r&   c                   V'       d   \        4       '       g   \        RRR7      h\        P                  P	                  4       '       g   \        R4      h \        4       '       dQ   RpVe   Vw  rEVP                  vr6\        P                  ! RRVR7      p\        P                  P                  WwV4      pT P                  R4       R# R#   \         d   pThRp?ii ; i)	ax  
Set whether to use memory efficient attention from `xformers` or not.

Args:
    use_memory_efficient_attention_xformers (`bool`):
        Whether to use memory efficient attention from `xformers` or not.
    attention_op (`Callable`, *optional*):
        The attention operation to use. Defaults to `None` which uses the default attention operation from
        `xformers`.
zeRefer to https://github.com/facebookresearch/xformers for more information on how to install xformersxformers)r)   zvtorch.cuda.is_available() should be True but is False. xformers' memory efficient attention is only available for GPU Ncudadevicedtype)rx   r   (   )r	   ModuleNotFoundErrorr,   r   is_availablerJ   SUPPORTED_DTYPESrandnxopsopsmemory_efficient_attention	Exceptionr   )	r:   r   r   r   op_fwop_bwrW   qes	   &&&      r#   +set_use_memory_efficient_attention_xformers@AttentionModuleMixin.set_use_memory_efficient_attention_xformers   s     3(**){#  ZZ,,.. / 

,.. $'3+7LE(-(>(>IE!KK
6O HH??aH **:61 3* ! Gs   A C CCCc                B   V P                   '       g/   \        P                  V P                  P                   R24       R# \        V RR4      '       d   R# V P                  P                  P                  P                  pV P                  P                  P                  P                  p\        V R4      '       Edz   V P                  '       Edg   \        P                  ! V P                  P                  P                  V P                   P                  P                  .4      pVP"                  ^,          pVP"                  ^ ,          p\$        P&                  ! WEV P(                  WR7      V n        V P*                  P                  P-                  V4       \        V R4      '       d   V P(                  '       dz   \        P                  ! V P                  P.                  P                  V P                   P.                  P                  .4      pV P*                  P.                  P-                  V4       EM\        P                  ! V P                  P                  P                  V P                  P                  P                  V P                   P                  P                  .4      pVP"                  ^,          pVP"                  ^ ,          p\$        P&                  ! WEV P(                  WR7      V n        V P0                  P                  P-                  V4       \        V R4      '       d   V P(                  '       d   \        P                  ! V P                  P.                  P                  V P                  P.                  P                  V P                   P.                  P                  .4      pV P0                  P.                  P-                  V4       \        V RR4      Ee   \        V R	R4      Ee   \        V R
R4      Ee   \        P                  ! V P2                  P                  P                  V P4                  P                  P                  V P6                  P                  P                  .4      pVP"                  ^,          pVP"                  ^ ,          p\$        P&                  ! WEV P8                  WR7      V n        V P:                  P                  P-                  V4       V P8                  '       d   \        P                  ! V P2                  P.                  P                  V P4                  P.                  P                  V P6                  P.                  P                  .4      pV P:                  P.                  P-                  V4       RV n        R# )zU
Fuse the query, key, and value projections into a single projection for efficiency.
zK does not support fusing QKV projections, so `fuse_projections` will no-op.Nfused_projectionsFis_cross_attention)biasr   r   use_bias
add_q_proj
add_k_proj
add_v_projT)rU   rk   debugrQ   rR   getattrto_qweightdatar   r   r4   r   r,   catto_kto_vshaper-   Linearr   to_kvcopy_r   to_qkvr   r   r   added_proj_biasto_added_qkvr   )r:   r   r   concatenated_weightsin_featuresout_featuresconcatenated_biass   &      r#   rV   %AttentionModuleMixin.fuse_projections   s    (((LL>>**++vw  4,e44!!&&--		  %%++4-..43J3J3J#(99dii.>.>.C.CTYYEUEUEZEZ-[#\ .44Q7K/55a8L;4==Y_mDJJJ##$89tZ((T]]]$)IItyy~~/B/BDIINNDWDW.X$Y!

%%&78 $)99dii.>.>.C.CTYYEUEUEZEZ\`\e\e\l\l\q\q-r#s .44Q7K/55a8L))KDMMZ`nDKKK$$%9:tZ((T]]]$)IItyy~~/B/BDIINNDWDWY]YbYbYgYgYlYl.m$n!  &&'89 D,-9lD1=lD1=#(99'',,doo.D.D.I.I4??KaKaKfKfg$  /44Q7K/55a8L "		0D0DV!D $$**+?@###$)II__))..0D0D0I0I4??K_K_KdKde%! !!&&,,->?!%r&   c                   V P                   '       g   R# \        V RR4      '       g   R# \        V R4      '       d   \        V R4       \        V R4      '       d   \        V R4       \        V R4      '       d   \        V R4       RV n        R# )zL
Unfuse the query, key, and value projections back to separate projections.
Nr   Fr   r   r   )rU   r   r4   delattrr   )r:   s   &r#   r\   'AttentionModuleMixin.unfuse_projections:  sw     ((( t0%88 4""D(#4!!D'"4((D.)!&r&   c                $   < V ^8  d   QhRS[ RR/# )r   
slice_sizer   Nint)r!   r"   s   "r#   r$   ri   T  s     & &c &d &r&   c                   \        V R4      '       d1   Ve-   WP                  8  d   \        RV RV P                   R24      hRpVe   V P                  R4      pVf   V P	                  4       pV P                  V4       R# )z
Set the slice size for attention computation.

Args:
    slice_size (`int`):
        The slice size for attention computation.
sliceable_head_dimNzslice_size z has to be smaller or equal to r3   sliced)r4   r   rJ   _get_compatible_processordefault_processor_clsrC   )r:   r   r>   s   && r#   set_attention_slice(AttentionModuleMixin.set_attention_sliceT  s     4-..:3Ij[r[rNr{:,6UVZVmVmUnnopqq	 !66x@I 224I9%r&   c                N   < V ^8  d   QhRS[ P                  RS[ P                  /# )r   tensorr   r,   Tensor)r!   r"   s   "r#   r$   ri   k  s#        r&   c                    V P                   pVP                  w  r4pVP                  W2,          W$V4      pVP                  ^ ^^^4      P                  W2,          WEV,          4      pV# )z
Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size // heads, seq_len, dim * heads]`.

Args:
    tensor (`torch.Tensor`): The tensor to reshape.

Returns:
    `torch.Tensor`: The reshaped tensor.
)headsr   reshapepermute)r:   r   	head_size
batch_sizeseq_lendims   &&    r#   batch_to_head_dim&AttentionModuleMixin.batch_to_head_dimk  s_     JJ	#)<< 
S
 7SQ1a+33J4KW\eVefr&   c                T   < V ^8  d   QhRS[ P                  RS[RS[ P                  /# )r   r   out_dimr   r,   r   r   )r!   r"   s   "r#   r$   ri   {  s*       s 5<< r&   c                B   V P                   pVP                  ^8X  d   VP                  w  rEp^pMVP                  w  rGrVVP                  WEV,          W6V,          4      pVP	                  ^ ^^^4      pV^8X  d&   VP                  WC,          WW,          Wc,          4      pV# )z
Reshape the tensor for multi-head attention processing.

Args:
    tensor (`torch.Tensor`): The tensor to reshape.
    out_dim (`int`, *optional*, defaults to `3`): The output dimension of the tensor.

Returns:
    `torch.Tensor`: The reshaped tensor.
)r   ndimr   r   r   )r:   r   r   r   r   r   r   	extra_dims   &&&     r#   head_to_batch_dim&AttentionModuleMixin.head_to_batch_dim{  s     JJ	;;!'-||$JI28,,/J7
i,?S\L\]1a+a<^^J$:G<OQTQabFr&   c                   < V ^8  d   QhRS[ P                  RS[ P                  RS[ P                  R,          RS[ P                  /# )r   querykeyattention_maskNr   r   )r!   r"   s   "r#   r$   ri     sC     - -\\-(--FKllUYFY-	-r&   c                T   VP                   pV P                  '       d!   VP                  4       pVP                  4       pVff   \        P                  ! VP
                  ^ ,          VP
                  ^,          VP
                  ^,          VP                   VP                  R7      p^ pMTp^p\        P                  ! VVVP                  RR4      VV P                  R7      p?V P                  '       d   VP                  4       pVP                  RR7      p?VP                  V4      pV# )a  
Compute the attention scores.

Args:
    query (`torch.Tensor`): The query tensor.
    key (`torch.Tensor`): The key tensor.
    attention_mask (`torch.Tensor`, *optional*): The attention mask to use.

Returns:
    `torch.Tensor`: The attention probabilities/scores.
r   r   )betaalphar   )r   upcast_attentionfloatr,   emptyr   r   baddbmm	transposescaleupcast_softmaxsoftmaxto)	r:   r   r   r   r   baddbmm_inputr   attention_scoresattention_probss	   &&&&     r#   get_attention_scores)AttentionModuleMixin.get_attention_scores  s        KKME))+C!!KKAA		!EKKX]XdXdM D*MD ==MM"b!**
 /557*22r2:),,U3r&   c          
      `   < V ^8  d   QhRS[ P                  RS[RS[RS[RS[ P                  /# )r   r   target_lengthr   r   r   r   )r!   r"   s   "r#   r$   ri     s=     ) )#ll);>)LO)Z])	)r&   c                l   V P                   pVf   V# VP                  R,          pWb8w  d   VP                  P                  R8X  dn   VP                  ^ ,          VP                  ^,          V3p\        P
                  ! WqP                  VP                  R7      p\        P                  ! W.^R7      pM\        P                  ! V^ V3RR7      pV^8X  d4   VP                  ^ ,          W5,          8  d   VP                  V^ R7      pV# V^8X  d%   VP                  ^4      pVP                  V^R7      pV# )a  
Prepare the attention mask for the attention computation.

Args:
    attention_mask (`torch.Tensor`): The attention mask to prepare.
    target_length (`int`): The target length of the attention mask.
    batch_size (`int`): The batch size for repeating the attention mask.
    out_dim (`int`, *optional*, defaults to `3`): Output dimension.

Returns:
    `torch.Tensor`: The prepared attention mask.
mpsr   r           )r}   r   )r   r   r   typer,   zerosr   r   Fpadrepeat_interleave	unsqueeze)	r:   r   r  r   r   r   current_lengthpadding_shapepaddings	   &&&&&    r#   prepare_attention_mask+AttentionModuleMixin.prepare_attention_mask  s!    JJ	!!!,2226*$$))U2 "0!5!5a!8.:N:Nq:QS` a++m;O;OXfXmXmn!&N+D!!L "#~=7IQT!Ua<##A&)??!/!A!A)QR!A!S
 	 \+55a8N+==iQ=ONr&   c                N   < V ^8  d   QhRS[ P                  RS[ P                  /# )r   encoder_hidden_statesr   r   )r!   r"   s   "r#   r$   ri     s&     % % %QVQ]Q] %r&   c                l   V P                   f   Q R4       h\        V P                   \        P                  4      '       d   V P                  V4      pV# \        V P                   \        P                  4      '       d8   VP                  ^^4      pV P                  V4      pVP                  ^^4      pV# Q h)z
Normalize the encoder hidden states.

Args:
    encoder_hidden_states (`torch.Tensor`): Hidden states of the encoder.

Returns:
    `torch.Tensor`: The normalized encoder hidden states.
zGself.norm_cross must be defined to call self.norm_encoder_hidden_states)
norm_crossrD   r-   	LayerNorm	GroupNormr   )r:   r  s   &&r#   norm_encoder_hidden_states/AttentionModuleMixin.norm_encoder_hidden_states  s     *u,uu*door||44$(OO4I$J! %$ 66 %:$C$CAq$I!$(OO4I$J!$9$C$CAq$I! %$ 5r&   )r   r>   r   r   r   )F)NFN)   )rR   r`   ra   rb   _default_processor_cls_available_processorsrU   r   rC   r1   r   r   r   r   r,   no_gradrV   r\   r   r   r   r  r  r  rd   re   rf   s   @r#   rT   rT   y   s     !# #(" "4 42 22 2.%7 %7N ]]_@& @&D ]]_' '2& &.   2- -^) )V% %r&   rT   c                p    V ^8  d   QhR\         P                  R\        P                  R\        R\        /# )r   ffhidden_states	chunk_dim
chunk_size)r-   r.   r,   r   r   )r!   s   "r#   r$   r$   
  s2      bii  QT be r&   c                 B   VP                   V,          V,          ^ 8w  d$   \        RVP                   V,           RV R24      hVP                   V,          V,          p\        P                  ! VP	                  WBR7       Uu. uF
  qP! V4      NK  	  upVR7      pV# u upi )r   z)`hidden_states` dimension to be chunked: z$ has to be divisible by chunk size: z[. Make sure to set an appropriate `chunk_size` when calling `unet.enable_forward_chunking`.r   )r   rJ   r,   r   chunk)r$  r%  r&  r'  
num_chunks	hid_slice	ff_outputs   &&&&   r#   _chunked_feed_forwardr-  
  s    9%
2a778K8KI8V7WW{  }G  |H  Hc  d
 	
 $$Y/:=J		(5(;(;J(;(VW(V9I(VWI  	Xs   Bc                   T   a a ] tR tRt oRtV3R lV 3R lltV3R lR ltRtVtV ;t	# )GatedSelfAttentionDensei  aX  
A gated self-attention dense layer that combines visual features and object features.

Parameters:
    query_dim (`int`): The number of channels in the query.
    context_dim (`int`): The number of channels in the context.
    n_heads (`int`): The number of heads to use for attention.
    d_head (`int`): The number of channels in each head.
c                2   < V ^8  d   QhRS[ RS[ RS[ RS[ /# )r   	query_dimcontext_dimn_headsd_headr   )r!   r"   s   "r#   r$   $GatedSelfAttentionDense.__annotate__%  s)      # C # s r&   c                  < \         SV `  4        \        P                  ! W!4      V n        \        WVR 7      V n        \        VRR7      V n        \        P                  ! V4      V n
        \        P                  ! V4      V n        V P                  R\        P                  ! \        P                  ! R4      4      4       V P                  R\        P                  ! \        P                  ! R4      4      4       RV n        R# ))r1  r   dim_headgegluactivation_fn
alpha_attnr
  alpha_denseTN)super__init__r-   r   linearr   attnFeedForwardr$  r  norm1norm2register_parameter	Parameterr,   r   enabled)r:   r1  r2  r3  r4  rQ   s   &&&&&r#   r>   GatedSelfAttentionDense.__init__%  s     ii7	6R	iw?\\),
\\),
bll5<<;L.MNr||ELL<M/NOr&   c                h   < V ^8  d   QhRS[ P                  RS[ P                  RS[ P                  /# )r   r   objsr   r   )r!   r"   s   "r#   r$   r5  6  s.     
 
 
U\\ 
ell 
r&   c                   V P                   '       g   V# VP                  ^,          pV P                  V4      pWP                  P	                  4       V P                  V P                  \        P                  ! W.^R7      4      4      RRV1R3,          ,          ,           pWP                  P	                  4       V P                  V P                  V4      4      ,          ,           pV# )rx   r   NNNN)rF  r   r?  r;  tanhr@  rB  r,   r   r<  r$  rC  )r:   r   rI  n_visuals   &&& r#   forwardGatedSelfAttentionDense.forward6  s    |||H771:{{4 $$&4::eii	WX>Y3Z)[\]_h`h_hjk\k)lll  %%'$''$**Q-*@@@r&   )r@  rF  r$  r?  rB  rC  
rR   r`   ra   rb   __doc__r>  rN  rd   re   __classcell__rQ   r"   s   @@r#   r/  r/    s#      "
 
 
r&   r/  c                   r   a a ] tR tRt oRtR
V3R lV 3R llltRV3R lR lltRV3R lR lltR	tVt	V ;t
# )JointTransformerBlockiC  a  
A Transformer block following the MMDiT architecture, introduced in Stable Diffusion 3.

Reference: https://huggingface.co/papers/2403.03206

Parameters:
    dim (`int`): The number of channels in the input and output.
    num_attention_heads (`int`): The number of heads to use for multi-head attention.
    attention_head_dim (`int`): The number of channels in each head.
    context_pre_only (`bool`): Boolean to determine if we should add some blocks associated with the
        processing of `context` conditions.
c                L   < V ^8  d   QhRS[ RS[ RS[ RS[RS[R,          RS[/# )r   r   num_attention_headsattention_head_dimcontext_pre_onlyqk_normNuse_dual_attention)r   rr   r    )r!   r"   s   "r#   r$   "JointTransformerBlock.__annotate__R  sS     O OO !O  	O
 O tO !Or&   c                  < \         S	V `  4        W`n        W@n        V'       d   R MRpV'       d   \	        V4      V n        M\        V4      V n        VR 8X  d   \        WRRRRR7      V n        M'VR8X  d   \        V4      V n        M\        RV R24      h\        \        R	4      '       d   \        4       pM\        R
4      h\        VRVVVVVRVVRR7      V n        V'       d   \        VRVVVRVVRR7	      V n        MRV n        \         P"                  ! VRRR7      V n        \'        WRR7      V n        V'       g2   \         P"                  ! VRRR7      V n        \'        WRR7      V n        MRV n        RV n        RV n        ^ V n        R# )ada_norm_continousada_norm_zeroFư>T
layer_norm)elementwise_affineepsr   	norm_typezUnknown context_norm_type: z>, currently only support `ada_norm_continous`, `ada_norm_zero`scaled_dot_product_attentionzYThe current PyTorch version does not support the `scaled_dot_product_attention` function.N)r1  cross_attention_dimadded_kv_proj_dimr7  r   r   rY  r   r>   rZ  rc  )	r1  rf  r7  r   r   r   r>   rZ  rc  rb  rc  gelu-approximate)r   dim_outr:  )r=  r>  r[  rY  r   rB  r   r   norm1_contextrJ   r4   r  r   r   r@  attn2r-   r  rC  rA  r$  norm2_context
ff_context_chunk_size
_chunk_dim)
r:   r   rW  rX  rY  rZ  r[  context_norm_typer>   rQ   s
   &&&&&&&  r#   r>  JointTransformerBlock.__init__R  s    	"4 04D0/.s3DJ)#.DJ 44!7U4S_"D /1!1#!6D-.?-@@~  1455-/Ik   $!'%-
	 "$(+)#
DJ DJ\\#%TJ
#BTU!#ceQU!VD)cN`aDO!%D"DO  r&   c                4   < V ^8  d   QhRS[ R,          RS[ /# r   r'  Nr   r   )r!   r"   s   "r#   r$   r\          t # r&   c                    Wn         W n        R # r  ro  rp  r:   r'  r   s   &&&r#   set_chunk_feed_forward,JointTransformerBlock.set_chunk_feed_forward      %r&   c                   < V ^8  d   QhRS[ P                  RS[ P                  RS[ P                  RS[S[S[3,          R,          RS[S[ P                  S[ P                  3,          /# )r   r%  r  tembjoint_attention_kwargsNr   )r,   FloatTensorr   r    r   r   r   )r!   r"   s   "r#   r$   r\    su     C4 C4((C4  %00C4 	C4
 !%S#X 5C4 
u||U\\)	*C4r&   c                
   T;'       g    / pV P                   '       d   V P                  WR 7      w  rVrxrpMV P                  WR 7      w  rVrxp	V P                  '       d   V P                  W#4      pMV P                  W#R 7      w  rrpV P                  ! RRVRV/VB w  ppVP                  ^4      V,          pVV,           pV P                   '       d6   V P                  ! RRX
/VB pXP                  ^4      V,          pVV,           pV P                  V4      pV^VR,          ,           ,          VR,          ,           pV P                  e-   \        V P                  WPP                  V P                  4      pMV P                  V4      pV	P                  ^4      V,          pVV,           pV P                  '       d   RpW!3# XP                  ^4      V,          pVV,           pV P                  V4      pV^XR,          ,           ,          XR,          ,           pV P                  e-   \        V P                  WP                  V P                  4      pMV P                  V4      pVXP                  ^4      V,          ,           pW!3# ))embr%  r  Nr_   rK  N)r[  rB  rY  rk  r@  r  rl  rC  ro  r-  r$  rp  rm  rn  )r:   r%  r  r}  r~  norm_hidden_statesgate_msa	shift_mlp	scale_mlpgate_mlpnorm_hidden_states2	gate_msa2r  
c_gate_msac_shift_mlpc_scale_mlp
c_gate_mlpattn_outputcontext_attn_outputattn_output2r,  context_ff_outputs   &&&&&                 r#   rN  JointTransformerBlock.forward  s    "8!=!=2"""kokuku lv lh)_h LP::Vc:KnH)   )-););<Q)X&[_[m[m% \n \X&Kj
 ,099 ,
,,
"<,
 %,
(( ((+k9%3"""::b4GbKabL$..q1L@L)L8M!ZZ6/1y7I3IJYW^M__'-dgg7I??\`\l\lmI 23I&&q)I5	%	1    $(!  %33 #-"6"6q"9<O"O$9<O$O!)-););<Q)R&)Cq;W^K_G_)`cnovcw)w&+$9OO%?RVRbRb%! %)OO4N$O!$9J<P<PQR<SVg<g$g!$33r&   )rp  ro  r@  rl  rY  r$  rn  rB  rk  rC  rm  r[  )FNFr   r  rR   r`   ra   rb   rQ  r>  ry  rN  rd   re   rR  rS  s   @@r#   rU  rU  C  s3     O Od 
C4 C4 C4r&   rU  c                   r   a a ] tR tRt oRtR
V3R lV 3R llltRV3R lR lltRV3R lR lltR	tVt	V ;t
# )BasicTransformerBlocki  ah  
A basic Transformer block.

Parameters:
    dim (`int`): The number of channels in the input and output.
    num_attention_heads (`int`): The number of heads to use for multi-head attention.
    attention_head_dim (`int`): The number of channels in each head.
    dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
    cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
    activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
    num_embeds_ada_norm (:
        obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`.
    attention_bias (:
        obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.
    only_cross_attention (`bool`, *optional*):
        Whether to use only cross-attention layers. In this case two cross attention layers are used.
    double_self_attention (`bool`, *optional*):
        Whether to use two self-attention layers. In this case no cross attention layers are used.
    upcast_attention (`bool`, *optional*):
        Whether to upcast the attention computation to float32. This is useful for mixed precision training.
    norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
        Whether to use learnable elementwise affine parameters for normalization.
    norm_type (`str`, *optional*, defaults to `"layer_norm"`):
        The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
    final_dropout (`bool` *optional*, defaults to False):
        Whether to apply a final dropout after the last feed-forward layer.
    attention_type (`str`, *optional*, defaults to `"default"`):
        The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`.
    positional_embeddings (`str`, *optional*, defaults to `None`):
        The type of positional embeddings to apply to.
    num_positional_embeddings (`int`, *optional*, defaults to `None`):
        The maximum number of positional embeddings to apply.
c          ,         < V ^8  d   QhRS[ RS[ RS[ RS[ R,          RS[RS[ R,          RS[R	S[R
S[RS[RS[RS[RS[RS[RS[RS[R,          RS[ R,          RS[ R,          RS[ R,          RS[ R,          RS[RS[/# )r   r   rW  rX  rf  Nr:  num_embeds_ada_normattention_biasonly_cross_attentiondouble_self_attentionr   norm_elementwise_affinerd  norm_epsfinal_dropoutattention_typepositional_embeddingsnum_positional_embeddings-ada_norm_continous_conditioning_embedding_dimada_norm_biasff_inner_dimff_biasattention_out_bias)r   r    rr   r   )r!   r"   s   "r#   r$   "BasicTransformerBlock.__annotate__  s    f ff !f  	f !4Zf f !4Zf f #f  $f f "&f f f  !f" #f$  #Tz%f& $':'f( 8;Tz)f* Tz+f, Dj-f. /f0 !1fr&   c                8  < \         SV `  4        Wn        W n        W0n        W@n        WPn        W`n        Wn        Wn	        Wn
        VV n        VV n        Wn        VR J;'       d    VR8H  V n        VR J;'       d    VR8H  V n        VR8H  V n        VR8H  V n        VR8H  V n        VR9   d   Vf   \'        RV RV R24      hWn        Wpn        V'       d   Vf   \'        R	4      hVR
8X  d   \-        VVR7      V n        MR V n        VR8X  d   \1        W4      V n        MRVR8X  d   \5        W4      V n        M:VR8X  d   \7        VVVVVR4      V n        M\8        P:                  ! WVR7      V n        \=        TTTTTV	'       d   TMR VVR7      V n        Vf	   V
'       du   VR8X  d   \1        W4      V n         M9VR8X  d   \7        VVVVVR4      V n         M\8        P:                  ! WV4      V n         \=        TV
'       g   TMR VVVVVVR7      V n!        M2VR8X  d   \8        P:                  ! WV4      V n         MR V n         R V n!        VR8X  d   \7        VVVVVR4      V n"        M2VR9   d   \8        P:                  ! WV4      V n"        MVR8X  d   R V n"        \G        VVVVVVR7      V n$        VR8X  g   VR8X  d   \K        WW#4      V n&        VR8X  d?   \8        PN                  ! \P        PR                  ! ^V4      VR,          ,          4      V n*        R V n+        ^ V n,        R # )Nr_  ada_normada_norm_singlera  ada_norm_continuous`norm_type` is set to w, but `num_embeds_ada_norm` is not defined. Please make sure to define `num_embeds_ada_norm` if setting `norm_type` to r3   \If `positional_embedding` type is defined, `num_positition_embeddings` must also be defined.
sinusoidalmax_seq_lengthrms_normrh  r1  r   r7  dropoutr   rf  r   out_biasr1  rf  r   r7  r  r   r   r  layer_norm_i2vgenr  r:  r  	inner_dimr   gatedzgated-text-imageg      ?r  r_  )r_  r  ra  )-r=  r>  r   rW  rX  r  rf  r:  r  r  r  r  r  r  use_ada_layer_norm_zerouse_ada_layer_normuse_ada_layer_norm_singleuse_layer_normuse_ada_layer_norm_continuousrJ   rd  r  r   	pos_embedr   rB  r   r   r-   r  r   attn1rC  rl  norm3rA  r$  r/  fuserrE  r,   r   scale_shift_tablero  rp  )r:   r   rW  rX  r  rf  r:  r  r  r  r  r   r  rd  r  r  r  r  r  r  r  r  r  r  rQ   s   &&&&&&&&&&&&&&&&&&&&&&&&r#   r>  BasicTransformerBlock.__init__  sR   4 	#6 "4#6 *,%:"'>$%:")B&$8! )<4(G'i'iYZiMi$#6d#B"_"_	U_H_)26G)G&'<7-6:O-O*55:M:U( 4KKT+UVX 
 ##6  &?&Gn  !L0:3OhiDN!DN 
"%c?DJ/))#CDJ///='DJ c[cdDJ%'7K 3QU-'	

 *.C J&)#C
333A+!
  \\#9PQ
"?T$7Z^)+#!1+	DJ --\\#9PQ
!
DJ --/='DJ EEc5LMDJ--DJ''"
 W$:L(L0K^sDJ ))%'\\%++a2ES2P%QD"  r&   c                4   < V ^8  d   QhRS[ R,          RS[ /# rt  r   )r!   r"   s   "r#   r$   r    ru  r&   c                    Wn         W n        R # r  rw  rx  s   &&&r#   ry  ,BasicTransformerBlock.set_chunk_feed_forward  r{  r&   c                p  < V ^8  d   QhRS[ P                  RS[ P                  R,          RS[ P                  R,          RS[ P                  R,          RS[ P                  R,          RS[S[S[3,          RS[ P                  R,          R	S[S[S[ P                  3,          R,          R
S[ P                  /	# )r   r%  r   Nr  encoder_attention_masktimestepcross_attention_kwargsclass_labelsadded_cond_kwargsr   )r,   r   
LongTensorr   r    r   )r!   r"   s   "r#   r$   r    s     x x||x t+x  %||d2	x
 !&t 3x ""T)x !%S#Xx &&-x  U\\ 12T9x 
xr&   c	                	   Ve*   VP                  RR 4      e   \        P                  R4       VP                  ^ ,          p	V P                  R8X  d   V P                  W4      p
EMV P                  R8X  d#   V P                  WWqP                  R7      w  rrpMV P                  R9   d   V P                  V4      p
MV P                  R8X  d   V P                  WR,          4      p
MV P                  R8X  dk   V P                  R ,          VP                  V	^R4      ,           P                  ^^R	7      w  pprrV P                  V4      p
V
^V,           ,          V,           p
M\        R
4      hV P                  e   V P                  V
4      p
Ve   VP                  4       M/ pVP                  RR 4      pV P                  ! V
3RV P                  '       d   TMR RV/VB pV P                  R8X  d   XP!                  ^4      V,          pMV P                  R8X  d
   XV,          pVV,           pVP"                  ^8X  d   VP%                  ^4      pVe   V P'                  VVR,          4      pV P(                  e   V P                  R8X  d   V P+                  W4      p
MlV P                  R9   d   V P+                  V4      p
MIV P                  R8X  d   Tp
M5V P                  R8X  d   V P+                  WR,          4      p
M\        R4      hV P                  e#   V P                  R8w  d   V P                  V
4      p
V P(                  ! V
3RVRV/VB pVV,           pV P                  R8X  d   V P-                  WR,          4      p
M"V P                  R8X  g   V P-                  V4      p
V P                  R8X  d&   V
^XR,          ,           ,          XR,          ,           p
V P                  R8X  d)   V P+                  V4      p
V
^X,           ,          X,           p
V P.                  e-   \1        V P2                  WP4                  V P.                  4      pMV P3                  V
4      pV P                  R8X  d   XP!                  ^4      V,          pMV P                  R8X  d
   XV,          pVV,           pVP"                  ^8X  d   VP%                  ^4      pV# )Nr   SPassing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.r  r_  )hidden_dtyper  pooled_text_embr  r   zIncorrect norm usedgligenr  r   rI  zIncorrect norm)ra  r  r   )r_  ra  r  r  )getrk   warningr   rd  rB  r   r  r   r)  rJ   r  copyrE   r  r  r  r   squeezer  rl  rC  r  ro  r-  r$  rp  )r:   r%  r   r  r  r  r  r  r  r   r  r  r  r  r  	shift_msa	scale_msagligen_kwargsr  r,  s   &&&&&&&&&           r#   rN  BasicTransformerBlock.forward  sk    "-%))'48Dtu #((+
>>Z'!%M!D^^.KO::DWDW LV LH) ^^BB!%M!:^^44!%MM^;_!`^^00&&t,x/?/?
Ar/RReA1eo KIy(y "&M!:!3q9}!E	!Q233>>%!%0B!C CYBd!7!<!<!>jl.228TBjj
;?;T;T;T"7Z^
 *
 %	
 >>_,",,Q/+=K^^00"[0K#m3")11!4M $ JJ}mF6KLM ::!~~+%)ZZ%H"#WW%)ZZ%>"#44 &3"#88%)ZZQb?c%d" !122~~)dnn@Q.Q%)^^4F%G"**"&;  6 )	K (-7M >>22!%MM^;_!`#44!%M!:>>_,!3q9W;M7M!NQZ[bQc!c>>..!%M!:!3q9}!E	!Q'-dgg7I??\`\l\lmI 23I>>_, **1-	9I^^00 9,I!M1")11!4Mr&   )rp  ro  r:  r  rX  r  rl  rf  r   r  r  r$  r  rB  rC  r  r  rd  rW  r  r  r  r  r  r  r  r  r  r  r  )r
  Nr8  NFFFFTra  h㈵>FdefaultNNNNNTTr  )NNNNNNNr  rS  s   @@r#   r  r    s4      Df fP 
x x xr&   r  c                   L   a a ] tR tRt oRtRV3R lV 3R llltR tRtVtV ;t	# )LuminaFeedForwardi;  a  
A feed-forward layer.

Parameters:
    hidden_size (`int`):
        The dimensionality of the hidden layers in the model. This parameter determines the width of the model's
        hidden representations.
    intermediate_size (`int`): The intermediate dimension of the feedforward layer.
    multiple_of (`int`, *optional*): Value to ensure hidden dimension is a multiple
        of this value.
    ffn_dim_multiplier (float, *optional*): Custom multiplier for hidden
        dimension. Defaults to None.
c          	      N   < V ^8  d   QhRS[ RS[ RS[ R,          RS[R,          /# )r   r   r  multiple_ofNffn_dim_multiplier)r   r   )r!   r"   s   "r#   r$   LuminaFeedForward.__annotate__J  s;        4Z	
 "DLr&   c                Z  < \         SV `  4        Ve   \        WB,          4      pW2V,           ^,
          V,          ,          p\        P                  ! VVRR7      V n        \        P                  ! VVRR7      V n        \        P                  ! VVRR7      V n        \        4       V n	        R # )NFr   )
r=  r>  r   r-   r   linear_1linear_2linear_3r   silu)r:   r   r  r  r  rQ   s   &&&&&r#   r>  LuminaFeedForward.__init__J  s     	).:;I$;a$?K#OP			

 		

 		

 J	r&   c                    V P                  V P                  V P                  V4      4      V P                  V4      ,          4      # r  )r  r  r  r  )r:   r   s   &&r#   rN  LuminaFeedForward.forwardh  s1    }}TYYt}}Q'784==;KKLLr&   )r  r  r  r  )   NrP  rS  s   @@r#   r  r  ;  s       <M Mr&   r  c                   n   a a ] tR tRt oRtR
V3R lV 3R llltV3R lR ltR
V3R lR lltR	tVt	V ;t
# )TemporalBasicTransformerBlockil  a  
A basic Transformer block for video like data.

Parameters:
    dim (`int`): The number of channels in the input and output.
    time_mix_inner_dim (`int`): The number of channels for temporal attention.
    num_attention_heads (`int`): The number of heads to use for multi-head attention.
    attention_head_dim (`int`): The number of channels in each head.
    cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
c                F   < V ^8  d   QhRS[ RS[ RS[ RS[ RS[ R,          /# )r   r   time_mix_inner_dimrW  rX  rf  Nr   )r!   r"   s   "r#   r$   *TemporalBasicTransformerBlock.__annotate__y  sA     3 33  3 !	3
  3 !4Z3r&   c                  < \         SV `  4        W8H  V n        \        P                  ! V4      V n        \        VVR R7      V n        \        P                  ! V4      V n        \        VVVRR7      V n
        Ve1   \        P                  ! V4      V n        \        VVVVR7      V n        MRV n        RV n        \        P                  ! V4      V n        \        VR R7      V n        RV n        RV n        R# )r8  )rj  r:  N)r1  r   r7  rf  )r1  rf  r   r7  r9  )r=  r>  is_resr-   r  norm_inrA  ff_inrB  r   r  rC  rl  r  r$  ro  rp  )r:   r   r  rW  rX  rf  rQ   s   &&&&&&r#   r>  &TemporalBasicTransformerBlock.__init__y  s     	/||C( !&!

 \\"45
(%' $	

 * &89DJ",$7)+	DJ DJDJ \\"45
0H  r&   c                .   < V ^8  d   QhRS[ R,          /# )r   r'  Nr   )r!   r"   s   "r#   r$   r    s      t r&   c                     Wn         ^V n        R# )rx   Nrw  )r:   r'  kwargss   &&,r#   ry  4TemporalBasicTransformerBlock.set_chunk_feed_forward  s    %r&   c                |   < V ^8  d   QhRS[ P                  RS[RS[ P                  R,          RS[ P                  /# )r   r%  
num_framesr  Nr   r   )r!   r"   s   "r#   r$   r    sD     7 7||7 7  %||d2	7
 
7r&   c                   VP                   ^ ,          pVP                   w  rVpWR,          pVR,          P                  WBWg4      pVP                  ^ ^^^4      pVP                  WF,          W'4      pTpV P                  V4      pV P                  e-   \        V P                  WP                  V P                  4      pMV P                  V4      pV P                  '       d	   W,           pV P                  V4      p	V P                  V	RR7      p
W,           pV P                  e,   V P                  V4      p	V P                  WR7      p
W,           pV P                  V4      p	V P                  e-   \        V P                  WP                  V P                  4      pMV P                  V	4      pV P                  '       d
   W,           pMTpVR,          P                  WFW'4      pVP                  ^ ^^^4      pVP                  WB,          Wg4      pV# )r   N)r  )NrK  )r   r   r   r  ro  r-  r  rp  r  rB  r  rl  rC  r  r$  )r:   r%  r  r  r   batch_frames
seq_lengthchannelsresidualr  r  r,  s   &&&&        r#   rN  %TemporalBasicTransformerBlock.forward  s    #((+
-:-@-@*(!/
%g.66zzd%--aAq9%--j.Ez\ ]3'1$**m__^b^n^noM JJ}5M;;;)4M!ZZ6jj!34jP#3 ::!!%M!:**%7*eK'7M "ZZ6'-dgg7I??\`\l\lmI 23I;;;%5M%M%g.66zzd%--aAq9%--j.Ez\r&   )rp  ro  r  rl  r$  r  r  rB  rC  r  r  r  r  rS  s   @@r#   r  r  l  s.     	3 3j 7 7 7r&   r  c                   H   a a ] tR tRt oRV3R lV 3R llltR tRtVtV ;t# )SkipFFTransformerBlocki  c                X   < V ^8  d   QhRS[ RS[ RS[ RS[ RS[RS[ R,          RS[R	S[/# )
r   r   rW  rX  kv_input_dimkv_input_dim_proj_use_biasrf  Nr  r  )r   rr   )r!   r"   s   "r#   r$   #SkipFFTransformerBlock.__annotate__  s_     (
 (
(
 !(
  	(

 (
 %)(
 !4Z(
 (
 !(
r&   c
           
       < \         S
V `  4        WA8w  d   \        P                  ! WAV4      V n        MR V n        \        VR4      V n        \        VVVVVVV	R7      V n        \        VR4      V n	        \        VVVVVVV	R7      V n
        R # )Nr`  )r1  r   r7  r  r   rf  r  )r1  rf  r   r7  r  r   r  )r=  r>  r-   r   	kv_mapperr   rB  r   r  rC  rl  )r:   r   rW  rX  r	  r
  r  rf  r  r  rQ   s   &&&&&&&&&&r#   r>  SkipFFTransformerBlock.__init__  s     	YY|:TUDN!DNS%(
%' 3'

 S%(
 3%''

r&   c                P   Ve   VP                  4       M/ pV P                  e&   V P                  \        P                  ! V4      4      pV P	                  V4      pV P
                  ! V3RV/VB pWQ,           pV P                  V4      pV P                  ! V3RV/VB pWQ,           pV# )Nr  )r  r  r  r  rB  r  rC  rl  )r:   r%  r  r  r  r  s   &&&&  r#   rN  SkipFFTransformerBlock.forward  s    BXBd!7!<!<!>jl>>%$(NN166:O3P$Q!!ZZ6jj
"7
 %
 $3!ZZ6jj
"7
 %
 $3r&   )r  rl  r  rB  rC  )r
  NFT)	rR   r`   ra   rb   r>  rN  rd   re   rR  rS  s   @@r#   r  r    s     (
 (
T r&   r  c                      a a ] tR tRt oRtRV3R lV 3R llltV3R lR ltRV3R lR lltRV3R	 lR
 lltRV3R lR llt	RV3R lR llt
RtVtV ;t# )FreeNoiseTransformerBlocki6  a  
A FreeNoise Transformer block.

Parameters:
    dim (`int`):
        The number of channels in the input and output.
    num_attention_heads (`int`):
        The number of heads to use for multi-head attention.
    attention_head_dim (`int`):
        The number of channels in each head.
    dropout (`float`, *optional*, defaults to 0.0):
        The dropout probability to use.
    cross_attention_dim (`int`, *optional*):
        The size of the encoder_hidden_states vector for cross attention.
    activation_fn (`str`, *optional*, defaults to `"geglu"`):
        Activation function to be used in feed-forward.
    num_embeds_ada_norm (`int`, *optional*):
        The number of diffusion steps used during training. See `Transformer2DModel`.
    attention_bias (`bool`, defaults to `False`):
        Configure if the attentions should contain a bias parameter.
    only_cross_attention (`bool`, defaults to `False`):
        Whether to use only cross-attention layers. In this case two cross attention layers are used.
    double_self_attention (`bool`, defaults to `False`):
        Whether to use two self-attention layers. In this case no cross attention layers are used.
    upcast_attention (`bool`, defaults to `False`):
        Whether to upcast the attention computation to float32. This is useful for mixed precision training.
    norm_elementwise_affine (`bool`, defaults to `True`):
        Whether to use learnable elementwise affine parameters for normalization.
    norm_type (`str`, defaults to `"layer_norm"`):
        The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
    final_dropout (`bool` defaults to `False`):
        Whether to apply a final dropout after the last feed-forward layer.
    attention_type (`str`, defaults to `"default"`):
        The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`.
    positional_embeddings (`str`, *optional*):
        The type of positional embeddings to apply to.
    num_positional_embeddings (`int`, *optional*, defaults to `None`):
        The maximum number of positional embeddings to apply.
    ff_inner_dim (`int`, *optional*):
        Hidden dimension of feed-forward MLP.
    ff_bias (`bool`, defaults to `True`):
        Whether or not to use bias in feed-forward MLP.
    attention_out_bias (`bool`, defaults to `True`):
        Whether or not to use bias in attention output project layer.
    context_length (`int`, defaults to `16`):
        The maximum number of frames that the FreeNoise block processes at once.
    context_stride (`int`, defaults to `4`):
        The number of frames to be skipped before starting to process a new batch of `context_length` frames.
    weighting_scheme (`str`, defaults to `"pyramid"`):
        The weighting scheme to use for weighting averaging of processed latent frames. As described in the
        Equation 9. of the [FreeNoise](https://huggingface.co/papers/2310.15169) paper, "pyramid" is the default
        setting used.
c          .         < V ^8  d   QhRS[ RS[ RS[ RS[RS[ R,          RS[RS[ R,          R	S[R
S[RS[RS[RS[RS[RS[RS[RS[R,          RS[ R,          RS[ R,          RS[RS[RS[ RS[ RS[/# )r   r   rW  rX  r  rf  Nr:  r  r  r  r  r   r  rd  r  r  r  r  r  r  r  context_lengthcontext_strideweighting_schemer   r   r    rr   )r!   r"   s   "r#   r$   &FreeNoiseTransformerBlock.__annotate__n  s    p pp !p  	p
 p !4Zp p !4Zp p #p  $p p "&p p p  !p"  #Tz#p$ $':%p& Dj'p( )p* !+p, -p. /p0 1pr&   c                  < \         SV `  4        Wn        W n        W0n        W@n        WPn        W`n        Wn        Wn	        Wn
        VV n        VV n        Wn        V P                  VVV4       VR J;'       d    VR8H  V n        VR J;'       d    VR8H  V n        VR8H  V n        VR8H  V n        VR8H  V n        VR9   d   Vf   \)        RV RV R24      hWn        Wpn        V'       d   Vf   \)        R	4      hVR
8X  d   \/        VVR7      V n        MR V n        \2        P4                  ! WVR7      V n        \9        TTTTTV	'       d   TMR VVR7      V n        Vf	   V
'       d?   \2        P4                  ! WV4      V n        \9        TV
'       g   TMR VVVVVVR7      V n        \A        VVVVVVR7      V n!        \2        P4                  ! WV4      V n"        R V n#        ^ V n$        R # )Nr_  r  r  ra  r  r  r  r3   r  r  r  rh  r  r  r  r  )%r=  r>  r   rW  rX  r  rf  r:  r  r  r  r  r  r  set_free_noise_propertiesr  r  r  r  r  rJ   rd  r  r   r  r-   r  rB  r   r  rC  rl  rA  r$  r  ro  rp  )r:   r   rW  rX  r  rf  r:  r  r  r  r  r   r  rd  r  r  r  r  r  r  r  r  r  r  rQ   s   &&&&&&&&&&&&&&&&&&&&&&&&r#   r>  "FreeNoiseTransformerBlock.__init__n  s   4 	#6 "4#6 *,%:"'>$%:")B&$8!&&~~GWX )<4(G'i'iYZiMi$#6d#B"_"_	U_H_)26G)G&'<7-6:O-O*55:M:U( 4KKT+UVX 
 ##6  &?&Gn  !L0:3OhiDN!DN \\#W_`
%'7K 3QU-'	

 *.Cc5LMDJ"?T$7Z^)+#!1+	DJ ''"
 \\#1HI
  r&   c                L   < V ^8  d   QhRS[ RS[S[S[ S[ 3,          ,          /# )r   r  r   )r   listr   )r!   r"   s   "r#   r$   r    s(      S T%S/5J r&   c                    . p\        ^ WP                  ,
          ^,           V P                  4       F3  pTp\        WV P                  ,           4      pVP	                  WE34       K5  	  V# r  )ranger  r  minappend)r:   r  frame_indicesiwindow_start
window_ends   &&    r#   _get_frame_indices,FreeNoiseTransformerBlock._get_frame_indices  sb    q*':'::Q>@S@STALZT-@-@)@AJ  ,!;< U r&   c                <   < V ^8  d   QhRS[ RS[RS[S[,          /# )r   r  r  r   )r   r    r  r   )r!   r"   s   "r#   r$   r    s)      S C X\]bXc r&   c                   VR 8X  d   R.V,          pV# VR8X  d   V^,          ^ 8X  d:   V^,          p\        \        ^V^,           4      4      pW3RRR1,          ,           pV# V^,           ^,          p\        \        ^V4      4      pW4.,           VRRR1,          ,           p V# VR8X  d   V^,          ^ 8X  dB   V^,          pR.V^,
          ,          V.,           pV\        \        V^ R4      4      ,           pV# V^,           ^,          pR.V,          pV\        \        V^ R4      4      ,           p V# \        RV 24      h)flatg      ?pyramidNdelayed_reverse_sawtoothg{Gz?z'Unsupported value for weighting_scheme=r   )r  r  rJ   )r:   r  r  weightsmids   &&&  r#   _get_frame_weights,FreeNoiseTransformerBlock._get_frame_weights  sL   v%ej(G8 5 *A~" AouQa01!DbDM1* % "A~!+uQ}-!E/GDbDM9   !;;A~" Ao&C!G,u4!DsAr):$;;  "A~!+&3,!DsAr):$;;  FGWFXYZZr&   c                0   < V ^8  d   QhRS[ RS[ RS[RR/# )r   r  r  r  r   N)r   r    )r!   r"   s   "r#   r$   r    s-     1 1!1361JM1	1r&   c                *    Wn         W n        W0n        R # r  )r  r  r  )r:   r  r  r  s   &&&&r#   r  3FreeNoiseTransformerBlock.set_free_noise_properties  s     -, 0r&   c                8   < V ^8  d   QhRS[ R,          RS[ RR/# )r   r'  Nr   r   r   )r!   r"   s   "r#   r$   r    s&      t # d r&   c                    Wn         W n        R # r  rw  rx  s   &&&r#   ry  0FreeNoiseTransformerBlock.set_chunk_feed_forward  r{  r&   c                   < V ^8  d   QhRS[ P                  RS[ P                  R,          RS[ P                  R,          RS[ P                  R,          RS[S[S[3,          RS[ P                  /# )r   r%  r   Nr  r  r  r   )r,   r   r   r    r   )r!   r"   s   "r#   r$   r    sz     { {||{ t+{  %||d2	{
 !&t 3{ !%S#X{ 
{r&   c                l	   Ve*   VP                  RR 4      e   \        P                  R4       Ve   VP                  4       M/ pVP                  pVP
                  p	VP                  ^4      p
V P                  V
4      pV P                  V P                  V P                  4      p\        P                  ! WV	R7      P                  ^ 4      P                  R4      pVR,          ^,          V
8H  pV'       gg   WP                  8  d   \        RV
: RV P                  : 24      hWR,          ^,          ,
          pVP                  WP                  ,
          V
34       \        P                   ! ^V
^3VR7      p\        P"                  ! V4      p\%        V4       EF  w  pw  pp\        P&                  ! VRVV13,          4      pVV,          pVRVV13,          pV P)                  V4      pV P*                  e   V P+                  V4      pV P,                  ! V3RV P.                  '       d   TMR R	V/VB pVV,           pVP0                  ^8X  d   VP3                  ^4      pV P4                  eb   V P7                  V4      pV P*                  e#   V P8                  R
8w  d   V P+                  V4      pV P4                  ! V3RVR	V/VB pVV,           pV\;        V4      ^,
          8X  di   V'       ga   VRX) R 13;;,          VRV) R 13,          VRV) R 13,          ,          ,          uu&   VRV) R 13;;,          VRV) 3,          ,          uu&   EK  VRVV13;;,          VV,          ,          uu&   VRVV13;;,          V,          uu&   EK  	  \        P<                  ! \?        VPA                  V P                  ^R7      VPA                  V P                  ^R7      4       UUu. uF(  w  pp\        PB                  ! V^ 8  VV,          V4      NK*  	  upp^R7      PE                  V	4      pV PG                  V4      pV PH                  e.   \K        V PL                  VV PN                  V PH                  4      pMV PM                  V4      pVV,           pVP0                  ^8X  d   VP3                  ^4      pV# u uppi )Nr   r  r   zExpected num_frames=z1 to be greater or equal than self.context_length=)r   rK  r  r   r  r   r   )(r  rk   r  r  r   r   sizer&  r/  r  r  r,   r   r  rJ   r!  r  
zeros_like	enumerate	ones_likerB  r  r  r  r   r  rl  rC  rd  rH   r   zipsplitwherer   r  ro  r-  r$  rp  )r:   r%  r   r  r  r  argsr  r   r   r  r"  frame_weightsis_last_frame_batch_completelast_frame_batch_lengthnum_times_accumulatedaccumulated_valuesr#  frame_start	frame_endr-  hidden_states_chunkr  r  accumulated_splitnum_times_splitr,  s   &&&&&&*,                   r#   rN  !FreeNoiseTransformerBlock.forward  s    "-%))'48DtuBXBd!7!<!<!>jl %%##"''*
//
;//0C0CTEZEZ[]OYYZ[\ffgij'4R'8';z'I$
 ,/// #8ZM9kW[WjWjVl!mnn&03DQ3G&G#  */B/B"BJ!OP %Q
A,>v N"--m<+4]+C'A'Y oo&;A{9?T<T&UVG}$G"/;y3H0H"I "&,?!@~~)%)^^4F%G"**"?C?X?X?X&;^b  . )	K #.0C"C"''1,&9&A&A!&D# zz%%)ZZ0C%D">>-$..DU2U)-8J)K&"jj&*? $: -	 '24G&G#C&**3O"1'>&>&?#?@',C+C+D(DEPQTkSkSlPlHmm@ &a*A)A)B&BCwqSjRjOjGkkC"1k)&;#;<@SV]@]]<%aY)>&>?7J?c ,D| 		 ;>&,,T-@-@a,H)//0C0C/K;;6% Oa/1B_1TVgh; 	
 "U) 	 "ZZ6'-dgg7I4??\`\l\lmI 23I!M1")11!4M-s   .R0
)rp  ro  r:  r  rX  r  rl  r  r  rf  r   r  r  r$  rB  rC  r  r  rd  rW  r  r  r  r  r  r  r  r  r  r  r  )r
  Nr8  NFFFFTra  r  FNNNTT      r+  )r+  r  )NNNN)rR   r`   ra   rb   rQ  r>  r&  r/  r  ry  rN  rd   re   rR  rS  s   @@r#   r  r  6  sS     4lp pd  @1 1 
{ { {r&   r  c                   X   a a ] tR tRt oRt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	# )	rA  i  a  
A feed-forward layer.

Parameters:
    dim (`int`): The number of channels in the input.
    dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
    mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
    dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
    activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
    final_dropout (`bool` *optional*, defaults to False): Apply a final dropout.
    bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
c                R   < V ^8  d   QhRS[ RS[ R,          RS[ RS[RS[RS[RS[/# )	r   r   rj  Nmultr  r:  r  r   r  )r!   r"   s   "r#   r$   FeedForward.__annotate__  sU     &1 &1&1 t&1 	&1
 &1 &1 &1 &1r&   c	                  < \         S
V `  4        Vf   \        W,          4      pVe   TMTpVR8X  d   \        WVR7      p	VR8X  d   \        WRVR7      p	MTVR8X  d   \	        WVR7      p	M?VR8X  d   \        WVR7      p	M*VR8X  d   \        WVR7      p	MVR	8X  d   \        WVR
R7      p	\        P                  ! . 4      V n
        V P                  P                  X	4       V P                  P                  \        P                  ! V4      4       V P                  P                  \        P                  ! WrVR7      4       V'       d2   V P                  P                  \        P                  ! V4      4       R # R # )Ngelur  ri  rL  )approximater   r8  zgeglu-approximateswigluzlinear-silur  )r   
activation)r=  r>  r   r   r   r   r   r   r-   
ModuleListnetr!  Dropoutr   )r:   r   rj  rP  r  r:  r  r  r   act_fnrQ   s   &&&&&&&&& r#   r>  FeedForward.__init__  s     	CJI$0'cF"#t4F..#f4HFg%35F11$S$?Fh&C6Fm+%c4FSF==$

7+,		)4@AHHOOBJJw/0 r&   c                N   < V ^8  d   QhRS[ P                  RS[ P                  /# )r   r%  r   r   )r!   r"   s   "r#   r$   rQ    s#      U\\ u|| r&   c                    \        V4      ^ 8  g   VP                  RR4      e   Rp\        RRV4       V P                   F  pV! V4      pK  	  V# )r   r   NzThe `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`.z1.0.0)rH   r  r   rX  )r:   r%  r@  r  deprecation_messager*   s   &&*,  r#   rN  FeedForward.forward  sQ    t9q=FJJw5A #Ugw(;<hhF"=1M r&   )rX  )NrM  r
  r8  FNTrP  rS  s   @@r#   rA  rA    s$     &1 &1P  r&   rA  )6typingr   r   r,   torch.nnr-   torch.nn.functional
functionalr  utilsr   r   utils.import_utilsr   r   r	   utils.torch_utilsr
   activationsr   r   r   r   r   r   attention_processorr   r   r   
embeddingsr   normalizationr   r   r   r   r   r   r   
get_loggerrR   rk   r   rT   r-  r.   r/  rU  r  r  r  r  r  rA  r_   r&   r#   <module>rl     s]   !     & f f 4 Y Y U U 5 q q D 
		H	%O, O,dN% N%b &bii & &R h4BII h4 h4V HBII H HV
.M		 .Mb ~BII ~ ~BERYY EP X		 X Xv
<")) <r&   