+
    Lj                          ^ RI Ht ^ RIt^ RIHt ^ RIHu Ht ^RIH	t	 ^RI
Ht ^RIHt ^RIHtHt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	 R
]P:                  4      t ! R R]P:                  4      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# )    )partialN)	deprecate)get_activation)SpatialNorm)Downsample1DDownsample2DFirDownsample2DKDownsample2Ddownsample_2d)AdaGroupNorm)FirUpsample2DKUpsample2D
Upsample1D
Upsample2Dupfirdn2d_nativeupsample_2dc                      a a ] tR t^+t oRtRRRRRRRR	R
^ RRRRRRRRRRRRRRRRRRR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	# )ResnetBlockCondNorm2Da  
A Resnet block that use normalization layer that incorporate conditioning information.

Parameters:
    in_channels (`int`): The number of channels in the input.
    out_channels (`int`, *optional*, default to be `None`):
        The number of output channels for the first conv2d layer. If None, same as `in_channels`.
    dropout (`float`, *optional*, defaults to `0.0`): The dropout probability to use.
    temb_channels (`int`, *optional*, default to `512`): the number of channels in timestep embedding.
    groups (`int`, *optional*, default to `32`): The number of groups to use for the first normalization layer.
    groups_out (`int`, *optional*, default to None):
        The number of groups to use for the second normalization layer. if set to None, same as `groups`.
    eps (`float`, *optional*, defaults to `1e-6`): The epsilon to use for the normalization.
    non_linearity (`str`, *optional*, default to `"swish"`): the activation function to use.
    time_embedding_norm (`str`, *optional*, default to `"ada_group"` ):
        The normalization layer for time embedding `temb`. Currently only support "ada_group" or "spatial".
    kernel (`torch.Tensor`, optional, default to None): FIR filter, see
        [`~models.resnet.FirUpsample2D`] and [`~models.resnet.FirDownsample2D`].
    output_scale_factor (`float`, *optional*, default to be `1.0`): the scale factor to use for the output.
    use_in_shortcut (`bool`, *optional*, default to `True`):
        If `True`, add a 1x1 nn.conv2d layer for skip-connection.
    up (`bool`, *optional*, default to `False`): If `True`, add an upsample layer.
    down (`bool`, *optional*, default to `False`): If `True`, add a downsample layer.
    conv_shortcut_bias (`bool`, *optional*, default to `True`):  If `True`, adds a learnable bias to the
        `conv_shortcut` output.
    conv_2d_out_channels (`int`, *optional*, default to `None`): the number of channels in the output.
        If None, same as `out_channels`.
out_channelsNconv_shortcutFdropout        temb_channels   groups
groups_outepsư>non_linearityswishtime_embedding_norm	ada_groupoutput_scale_factor      ?use_in_shortcutupdownconv_shortcut_biasTconv_2d_out_channelsc          !         < V ^8  d   QhRS[ RS[ R,          RS[RS[RS[ RS[ RS[ R,          R	S[R
S[RS[RS[RS[R,          RS[RS[RS[RS[ R,          /# )   in_channelsr   Nr   r   r   r   r   r   r   r!   r#   r%   r&   r'   r(   r)   )intboolfloatstr)format__classdict__s   "C/app/.local/lib/python3.14/site-packages/diffusers/models/resnet.py__annotate__"ResnetBlockCondNorm2D.__annotate__I   s     I I I Dj	I
 I I I I $JI I I !I #I I I  !I" !#I$ "Dj%I    c          	     h  < \         SV `  4        Wn        Vf   TMTpW n        W0n        Wn        Wn        Wn        Wn        Vf   TpV P                  R8X  d   \        WQWhR7      V n
        M:V P                  R8X  d   \        W4      V n
        M\        RV P                   24      h\        P                  ! W^^^R7      V n        V P                  R8X  d   \        WRWxR7      V n        M:V P                  R8X  d   \        W%4      V n        M\        RV P                   24      h\"        P                  P%                  V4      V n        T;'       g    Tp\        P                  ! VV^^^R7      V n        \+        V	4      V n        R ;V n        V n        V P
                  '       d   \3        VRR7      V n        M&V P                  '       d   \5        VR^RR	7      V n        Vf   V P                  V8g  MTV n        R V n        V P6                  '       d$   \        P                  ! VV^^^ VR
7      V n        R # R # )Nr"   )r   spatialz" unsupported time_embedding_norm: kernel_sizestridepaddingFuse_convopr>   r<   namer:   r;   r<   bias)super__init__r,   r   use_conv_shortcutr&   r'   r#   r!   r   norm1r   
ValueErrornnConv2dconv1norm2torchDropoutr   conv2r   nonlinearityupsample
downsampler   r   r%   r   )selfr,   r   r   r   r   r   r   r   r   r!   r#   r%   r&   r'   r(   r)   	__class__s   &$$$$$$$$$$$$$$$$r3   rE   ResnetBlockCondNorm2D.__init__I   s   ( 	&&2&:{(!.	#6 #6 J##{2%m&RDJ%%2$[@DJA$BZBZA[\]]YY{aPQ[\]
##{2%m:WDJ%%2$\ADJA$BZBZA[\]]xx''03CC|YY|-AqYZdef
*=9*..777&{UCDMYYY*;PQX\]DOKZKbt//3GGhw!!#$'"D  r6   c                h   < V ^8  d   QhRS[ P                  RS[ P                  RS[ P                  /# r+   input_tensortembreturnrM   Tensor)r1   r2   s   "r3   r4   r5      s1     % %ELL % %Z_ZfZf %r6   c                    \        V4      ^ 8  g   VP                  RR4      e   Rp\        RRV4       TpV P                  Wb4      pV P	                  V4      pV P
                  e\   VP                  ^ ,          ^@8  d!   VP                  4       pVP                  4       pV P                  V4      pV P                  V4      pM0V P                  e#   V P                  V4      pV P                  V4      pV P                  V4      pV P                  Wb4      pV P	                  V4      pV P                  V4      pV P                  V4      pV P                  e   V P                  V4      pW,           V P                  ,          pV# )r   scaleNThe `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`.1.0.0)lengetr   rG   rP   rQ   shape
contiguousrR   rK   rL   r   rO   r   r#   )rS   rX   rY   argskwargsdeprecation_messagehidden_statesoutput_tensors   &&&*,   r3   forwardResnetBlockCondNorm2D.forward   sN   t9q=FJJw5A #Ugw(;<$

=7))-8==$""1%++668 - 8 8 :==6L MM-8M__(??<8L OOM:M

=1

=7))-8]3

=1)--l;L%59Q9QQr6   )rK   rO   r   r'   rR   r   r,   rP   rG   rL   r   r#   r!   r&   rQ   rF   r%   
__name__
__module____qualname____firstlineno____doc__rE   rj   __static_attributes____classdictcell____classcell__rT   r2   s   @@r3   r   r   +   s     :I $(	I
 $I I !I I "&I I %I $/I &)I (,I I  !I" $(#I$ ,0%I IV% % %r6   r   c            $          a a ] tR t^t oRtRRRRRRRR	R
^ RRRRRRRRRRRRRRRRRRRRRRRRR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	# )"ResnetBlock2Da  
A Resnet block.

Parameters:
    in_channels (`int`): The number of channels in the input.
    out_channels (`int`, *optional*, default to be `None`):
        The number of output channels for the first conv2d layer. If None, same as `in_channels`.
    dropout (`float`, *optional*, defaults to `0.0`): The dropout probability to use.
    temb_channels (`int`, *optional*, default to `512`): the number of channels in timestep embedding.
    groups (`int`, *optional*, default to `32`): The number of groups to use for the first normalization layer.
    groups_out (`int`, *optional*, default to None):
        The number of groups to use for the second normalization layer. if set to None, same as `groups`.
    eps (`float`, *optional*, defaults to `1e-6`): The epsilon to use for the normalization.
    non_linearity (`str`, *optional*, default to `"swish"`): the activation function to use.
    time_embedding_norm (`str`, *optional*, default to `"default"` ): Time scale shift config.
        By default, apply timestep embedding conditioning with a simple shift mechanism. Choose "scale_shift" for a
        stronger conditioning with scale and shift.
    kernel (`torch.Tensor`, optional, default to None): FIR filter, see
        [`~models.resnet.FirUpsample2D`] and [`~models.resnet.FirDownsample2D`].
    output_scale_factor (`float`, *optional*, default to be `1.0`): the scale factor to use for the output.
    use_in_shortcut (`bool`, *optional*, default to `True`):
        If `True`, add a 1x1 nn.conv2d layer for skip-connection.
    up (`bool`, *optional*, default to `False`): If `True`, add an upsample layer.
    down (`bool`, *optional*, default to `False`): If `True`, add a downsample layer.
    conv_shortcut_bias (`bool`, *optional*, default to `True`):  If `True`, adds a learnable bias to the
        `conv_shortcut` output.
    conv_2d_out_channels (`int`, *optional*, default to `None`): the number of channels in the output.
        If None, same as `out_channels`.
r   Nr   Fr   r   r   r   r   r   pre_normTr   r   r   r    skip_time_actr!   defaultkernelr#   r$   r%   r&   r'   r(   r)   c          '         < V ^8  d   QhRS[ RS[ R,          RS[RS[RS[ RS[ RS[ R,          R	S[R
S[RS[RS[RS[RS[P
                  R,          RS[RS[R,          RS[RS[RS[RS[ R,          /# )r+   r,   r   Nr   r   r   r   r   rx   r   r   ry   r!   r{   r#   r%   r&   r'   r(   r)   )r-   r.   r/   r0   rM   r\   )r1   r2   s   "r3   r4   ResnetBlock2D.__annotate__   s     b b b Dj	b
 b b b b $Jb b b b b !b t#b  #!b" #b$ %b& 'b( !)b* "Dj+br6   c          	       <a \         SV `  4        VR 8X  d   \        R4      hVR8X  d   \        R4      hRV n        Wn        Vf   TMTpW n        W0n        VV n        VV n        Wn	        Wn
        Wn        Vf   Tp\        P                  P                  WaV	RR7      V n        \        P                   ! W^^^R7      V n        Ve|   V P                  R8X  d   \        P$                  ! WR4      V n        MUV P                  R	8X  d%   \        P$                  ! V^V,          4      V n        M \        R
V P                   R24      hRV n        \        P                  P                  WrV	RR7      V n        \        P                  P+                  V4      V n        T;'       g    Tp\        P                   ! VV^^^R7      V n        \1        V
4      V n        R;V n        V n        V P                  '       dR   VR8X  d   RoV3R lV n        MVR8X  d#   \9        \:        P<                  RRR7      V n        Mw\?        VRR7      V n        MdV P                  '       dS   VR8X  d   RoV3R lV n        M=VR8X  d#   \9        \:        P@                  ^^R7      V n        M\C        VR^RR7      V n        Vf   V P                  V8g  MTV n"        RV n#        V PD                  '       d$   \        P                   ! VV^^^ VR7      V n#        R# R# )r"   zkThis class cannot be used with `time_embedding_norm==ada_group`, please use `ResnetBlockCondNorm2D` insteadr8   ziThis class cannot be used with `time_embedding_norm==spatial`, please use `ResnetBlockCondNorm2D` insteadTN
num_groupsnum_channelsr   affiner9   rz   scale_shiftzunknown time_embedding_norm :  firc                    < \        V SR 7      # )r{   )r   x
fir_kernels   &r3   <lambda>(ResnetBlock2D.__init__.<locals>.<lambda>$  s    +a
*Kr6   sde_vpg       @nearest)scale_factormodeFr=   c                    < \        V SR 7      # r   )r   r   s   &r3   r   r   ,  s    M!J,Or6   )r:   r;   r?   r@   rB   )      r   r   )$rD   rE   rH   rx   r,   r   rF   r&   r'   r#   r!   ry   rM   rI   	GroupNormrG   rJ   rK   Lineartime_emb_projrL   rN   r   rO   r   rP   rQ   rR   r   Finterpolater   
avg_pool2dr   r%   r   )rS   r,   r   r   r   r   r   r   rx   r   r   ry   r!   r{   r#   r%   r&   r'   r(   r)   r   rT   s   &$$$$$$$$$$$$$$$$$$$@r3   rE   ResnetBlock2D.__init__   s   . 	+-}  )+{  &&2&:{(!.	#6 #6 *JXX''6Y\ei'j
YY{aPQ[\]
$''94%'YY}%K"))]:%'YY}a,>N%O" #A$BZBZA[[\!]^^!%DXX'':^ajn'o
xx''03CC|YY|-AqYZdef
*=9*..777)
 K8# 'Ci X *; GYYY)
"O8#")!,,Aa"P".{UTU\`"aKZKbt//3GGhw!!#$'"D  r6   c                h   < V ^8  d   QhRS[ P                  RS[ P                  RS[ P                  /# rW   r[   )r1   r2   s   "r3   r4   r}   ?  s1     : :ELL : :Z_ZfZf :r6   c                   \        V4      ^ 8  g   VP                  RR4      e   Rp\        RRV4       TpV P                  V4      pV P	                  V4      pV P
                  e\   VP                  ^ ,          ^@8  d!   VP                  4       pVP                  4       pV P                  V4      pV P                  V4      pM0V P                  e#   V P                  V4      pV P                  V4      pV P                  V4      pV P                  e<   V P                  '       g   V P	                  V4      pV P                  V4      R	,          pV P                  R8X  d   Ve	   Wb,           pV P                  V4      pMV P                  R8X  da   Vf   \        RV P                   24      h\        P                   ! V^^R7      w  rxV P                  V4      pV^V,           ,          V,           pMV P                  V4      pV P	                  V4      pV P#                  V4      pV P%                  V4      pV P&                  e4   V P(                  '       d   VP                  4       pV P'                  V4      pW,           V P*                  ,          p	V	# )
r   r^   Nr_   r`   rz   r   z9 `temb` should not be None when `time_embedding_norm` is )dim)NNNr   NN)ra   rb   r   rG   rP   rQ   rc   rd   rR   rK   r   ry   r!   rL   rH   rM   chunkr   rO   r   trainingr#   )
rS   rX   rY   re   rf   rg   rh   
time_scale
time_shiftri   s
   &&&*,     r3   rj   ResnetBlock2D.forward?  s9   t9q=FJJw5A #Ugw(;<$

=1))-8==$""1%++668 - 8 8 :==6L MM-8M__(??<8L OOM:M

=1)%%%((.%%d+,<=D##y0 - 4 JJ}5M%%6| OPTPhPhOij  &+[[qa%@"J JJ}5M)Q^<zIM JJ}5M))-8]3

=1) }}}+668--l;L%59Q9QQr6   )rK   rO   r   r'   rR   r   r,   rP   rG   rL   r   r#   rx   ry   r   r!   r&   rQ   rF   r%   rl   ru   s   @@r3   rw   rw      s     <b $(	b
 $b b !b b "&b b b %b $b $-b '+b  &)!b" (,#b$ %b& 'b( $()b* ,0+b bH: : :r6   rw   c                X    V ^8  d   QhR\         P                  R\         P                  /# )r+   tensorrZ   r[   )r1   s   "r3   r4   r4   }  s&     O O5<< OELL Or6   c                    \        V P                  4      ^8X  d
   V R,          # \        V P                  4      ^8X  d
   V R,          # \        V P                  4      ^8X  d
   V R,          # \        R\        V 4       R24      h)r+   z`len(tensor)`: z has to be 2, 3 or 4.)r   r   N)r   r   Nr   )r   r   r   r   )ra   rc   rH   )r   s   &r3   rearrange_dimsr   }  so    
6<<Aj!!
6<<Am$$	V\\	a	j!!?3v;-7LMNNr6   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	# )	Conv1dBlocki  ax  
Conv1d --> GroupNorm --> Mish

Parameters:
    inp_channels (`int`): Number of input channels.
    out_channels (`int`): Number of output channels.
    kernel_size (`int` or `tuple`): Size of the convolving kernel.
    n_groups (`int`, default `8`): Number of groups to separate the channels into.
    activation (`str`, defaults to `mish`): Name of the activation function.
c          
      ^   < V ^8  d   QhRS[ RS[ RS[ S[S[ S[ 3,          ,          RS[ RS[/# )r+   inp_channelsr   r:   n_groups
activationr-   tupler0   )r1   r2   s   "r3   r4   Conv1dBlock.__annotate__  sJ     / // / 5c?*	/
 / /r6   c                   < \         SV `  4        \        P                  ! WW3^,          R7      V n        \        P
                  ! WB4      V n        \        V4      V n        R# )r+   r<   N)	rD   rE   rI   Conv1dconv1dr   
group_normr   mish)rS   r   r   r:   r   r   rT   s   &&&&&&r3   rE   Conv1dBlock.__init__  sD     	iiK`aQab,,x>":.	r6   c                N   < V ^8  d   QhRS[ P                  RS[ P                  /# )r+   inputsrZ   r[   )r1   r2   s   "r3   r4   r     s#      ell u|| r6   c                    V P                  V4      p\        V4      pV P                  V4      p\        V4      pV P                  V4      pV# N)r   r   r   r   )rS   r   intermediate_reproutputs   &&  r3   rj   Conv1dBlock.forward  sM     KK/*+<= OO,=>*+<=,-r6   )r   r   r   )   r   rl   ru   s   @@r3   r   r     s#     	/ /  r6   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	# )	ResidualTemporalBlock1Di  au  
Residual 1D block with temporal convolutions.

Parameters:
    inp_channels (`int`): Number of input channels.
    out_channels (`int`): Number of output channels.
    embed_dim (`int`): Embedding dimension.
    kernel_size (`int` or `tuple`): Size of the convolving kernel.
    activation (`str`, defaults `mish`): It is possible to choose the right activation function.
c                ^   < V ^8  d   QhRS[ RS[ RS[ RS[ S[S[ S[ 3,          ,          RS[/# )r+   r   r   	embed_dimr:   r   r   )r1   r2   s   "r3   r4   $ResidualTemporalBlock1D.__annotate__  sJ     
 

 
 	

 5c?*
 
r6   c                :  < \         SV `  4        \        WV4      V n        \        W"V4      V n        \        V4      V n        \        P                  ! W24      V n	        W8w  d   \        P                  ! W^4      V n        R# \        P                  ! 4       V n        R# )r   N)rD   rE   r   conv_inconv_outr   time_emb_actrI   r   time_embr   Identityresidual_conv)rS   r   r   r   r:   r   rT   s   &&&&&&r3   rE    ResidualTemporalBlock1D.__init__  s{     	"<{K#LL*:6		): 9E8TBIIl!4 	Z\ZeZeZg 	r6   c                h   < V ^8  d   QhRS[ P                  RS[ P                  RS[ P                  /# )r+   r   trZ   r[   )r1   r2   s   "r3   r4   r     s.     0 0ell 0u|| 0 0r6   c                    V P                  V4      pV P                  V4      pV P                  V4      \        V4      ,           pV P	                  V4      pW0P                  V4      ,           # )z
Args:
    inputs : [ batch_size x inp_channels x horizon ]
    t : [ batch_size x embed_dim ]

returns:
    out : [ batch_size x out_channels x horizon ]
)r   r   r   r   r   r   )rS   r   r   outs   &&& r3   rj   ResidualTemporalBlock1D.forward  s\     a MM!ll6"^A%66mmC ''///r6   )r   r   r   r   r   )   r   rl   ru   s   @@r3   r   r     s#     	
 
&0 0 0r6   r   c                   \   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tVtV ;t	# )
TemporalConvLayeri  a  
Temporal convolutional layer that can be used for video (sequence of images) input Code mostly copied from:
https://github.com/modelscope/modelscope/blob/1509fdb973e5871f37148a4b5e5964cafd43e64d/modelscope/models/multi_modal/video_synthesis/unet_sd.py#L1016

Parameters:
    in_dim (`int`): Number of input channels.
    out_dim (`int`): Number of output channels.
    dropout (`float`, *optional*, defaults to `0.0`): The dropout probability to use.
c                @   < V ^8  d   QhRS[ RS[ R,          RS[RS[ /# )r+   in_dimout_dimNr   norm_num_groupsr-   r/   )r1   r2   s   "r3   r4   TemporalConvLayer.__annotate__  s7     ', ',', t', 	',
 ',r6   c                  < \         SV `  4        T;'       g    TpWn        W n        \        P
                  ! \        P                  ! WA4      \        P                  ! 4       \        P                  ! WRRR7      4      V n	        \        P
                  ! \        P                  ! WB4      \        P                  ! 4       \        P                  ! V4      \        P                  ! W!RRR7      4      V n        \        P
                  ! \        P                  ! WB4      \        P                  ! 4       \        P                  ! V4      \        P                  ! W!RRR7      4      V n        \        P
                  ! \        P                  ! WB4      \        P                  ! 4       \        P                  ! V4      \        P                  ! W!RRR7      4      V n        \        P                  P                  V P                  R,          P                   4       \        P                  P                  V P                  R,          P"                  4       R# )r   r   Nr   r   r   )r   r   r   )rD   rE   r   r   rI   
Sequentialr   SiLUConv3drK   rN   rO   conv3conv4initzeros_weightrC   )rS   r   r   r   r   rT   s   &&&&&r3   rE   TemporalConvLayer.__init__  se    	##V ]]LL1GGIIIfy)D


 ]]LL2GGIJJwIIgy)D	

 ]]LL2GGIJJwIIgy)D	

 ]]LL2GGIJJwIIgy)D	

 	tzz"~,,-
tzz"~**+r6   c                T   < V ^8  d   QhRS[ P                  RS[RS[ P                  /# )r+   rh   
num_framesrZ   rM   r\   r-   )r1   r2   s   "r3   r4   r     s*      U\\ s 5<< r6   c                   VR,          P                  RV3VP                  R,          ,           4      P                  ^ ^^^^4      pTpV P                  V4      pV P	                  V4      pV P                  V4      pV P                  V4      pW1,           pVP                  ^ ^^^^4      P                  VP                  ^ ,          VP                  ^,          ,          R3VP                  R,          ,           4      pV# )N:r   NN:r   NNNr   r   )reshaperc   permuterK   rO   r   r   )rS   rh   r   identitys   &&& r3   rj   TemporalConvLayer.forward  s    '"**B
+;m>Q>QRT>U+UV^^_`bcefhiklm 	 !

=1

=1

=1

=1 0%--aAq!<DD  #m&9&9!&<<bAMDWDWXZD[[
 r6   )rK   rO   r   r   r   r   )Nr       )r   rl   ru   s   @@r3   r   r     s$     ', ',R  r6   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	# )	TemporalResnetBlocki"  a  
A Resnet block.

Parameters:
    in_channels (`int`): The number of channels in the input.
    out_channels (`int`, *optional*, default to be `None`):
        The number of output channels for the first conv2d layer. If None, same as `in_channels`.
    temb_channels (`int`, *optional*, default to `512`): the number of channels in timestep embedding.
    eps (`float`, *optional*, defaults to `1e-6`): The epsilon to use for the normalization.
c                @   < V ^8  d   QhRS[ RS[ R,          RS[ RS[/# )r+   r,   r   Nr   r   r   )r1   r2   s   "r3   r4    TemporalResnetBlock.__annotate__.  s7     4 44 Dj4 	4
 4r6   c                  < \         SV `  4        Wn        Vf   TMTpW n        RpV Uu. uF  qf^,          NK  	  pp\        P
                  P                  ^ WRR7      V n        \
        P                  ! VVV^VR7      V n	        Ve   \
        P                  ! W24      V n        MR V n        \        P
                  P                  ^ W$RR7      V n        \        P
                  P                  R4      V n        \
        P                  ! VVV^VR7      V n        \!        R4      V n        V P                  V8g  V n        R V n        V P$                  '       d#   \
        P                  ! VV^^^ R7      V n        R # R # u upi )NTr   r9   r   silur   )rD   rE   r,   r   rM   rI   r   rG   r   rK   r   r   rL   rN   r   rO   r   rP   r%   r   )	rS   r,   r   r   r   r:   kr<   rT   s	   &&&&&   r3   rE   TemporalResnetBlock.__init__.  sW    	&&2&:{(#./;a66;/XX''2Kae'f
YY#

 $!#=!GD!%DXX''2Lbf'g
xx'',YY#

 +62#//<?!!#"D  A 0s   E7c                h   < V ^8  d   QhRS[ P                  RS[ P                  RS[ P                  /# rW   r[   )r1   r2   s   "r3   r4   r   d  s.      ELL   r6   c                   TpV P                  V4      pV P                  V4      pV P                  V4      pV P                  eG   V P                  V4      pV P                  V4      R,          pVP	                  ^ ^^^^4      pW2,           pV P                  V4      pV P                  V4      pV P                  V4      pV P                  V4      pV P                  e   V P                  V4      pW,           pV# )N)r   r   r   NN)	rG   rP   rK   r   r   rL   r   rO   r   )rS   rX   rY   rh   ri   s   &&&  r3   rj   TemporalResnetBlock.forwardd  s    $

=1))-8

=1)$$T*D%%d+,?@D<<1aA.D)0M

=1))-8]3

=1)--l;L$4r6   )rK   rO   r   r   r,   rP   rG   rL   r   r   r%   )Nr   r   rl   ru   s   @@r3   r   r   "  s$     	4 4l  r6   r   c                   \   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tVtV ;t	# )
SpatioTemporalResBlocki  a  
A SpatioTemporal Resnet block.

Parameters:
    in_channels (`int`): The number of channels in the input.
    out_channels (`int`, *optional*, default to be `None`):
        The number of output channels for the first conv2d layer. If None, same as `in_channels`.
    temb_channels (`int`, *optional*, default to `512`): the number of channels in timestep embedding.
    eps (`float`, *optional*, defaults to `1e-6`): The epsilon to use for the spatial resenet.
    temporal_eps (`float`, *optional*, defaults to `eps`): The epsilon to use for the temporal resnet.
    merge_factor (`float`, *optional*, defaults to `0.5`): The merge factor to use for the temporal mixing.
    merge_strategy (`str`, *optional*, defaults to `learned_with_images`):
        The merge strategy to use for the temporal mixing.
    switch_spatial_to_temporal_mix (`bool`, *optional*, defaults to `False`):
        If `True`, switch the spatial and temporal mixing.
c                `   < V ^8  d   QhRS[ RS[ R,          RS[ RS[RS[R,          RS[RS[/# )	r+   r,   r   Nr   r   temporal_epsmerge_factorswitch_spatial_to_temporal_mix)r-   r/   r.   )r1   r2   s   "r3   r4   #SpatioTemporalResBlock.__annotate__  sY     
 

 Dj
 	

 
 dl
 
 )-
r6   c	                   < \         S	V `  4        \        VVVVR 7      V n        \	        Ve   TMTVe   TMTTVe   TMTR 7      V n        \        VVVR7      V n        R# ))r,   r   r   r   N)alphamerge_strategyr  )rD   rE   rw   spatial_res_blockr   temporal_res_blockAlphaBlender
time_mixer)
rS   r,   r   r   r   r   r   r  r  rT   s
   &&&&&&&&&r3   rE   SpatioTemporalResBlock.__init__  sp     	!.#%'	"
 #6(4(@k)5)A{' , 8c	#
 ')+I
r6   c                   < V ^8  d   QhRS[ P                  RS[ P                  R,          RS[ P                  R,          /# )r+   rh   rY   Nimage_only_indicatorr[   )r1   r2   s   "r3   r4   r    s?      || llT! $llT1	r6   c                   VP                   R,          pV P                  W4      pVP                   w  rVrxWT,          p	VR,          P                  WWgV4      P                  ^ ^^^^4      p
VR,          P                  WWgV4      P                  ^ ^^^^4      pVe   VP                  WR4      pV P	                  W4      pV P                  V
VVR7      pVP                  ^ ^^^^4      P                  WVWx4      pV# )r   )	x_spatial
x_temporalr  r   r   )rc   r  r   r   r  r	  )rS   rh   rY   r  r   batch_frameschannelsheightwidth
batch_sizehidden_states_mixs   &&&&       r3   rj   SpatioTemporalResBlock.forward  s    *//3
..}C0=0C0C-!/
 '"**:8UZ[ccdeghjkmnpqr 	 '"**:8UZ[ccdeghjkmnpqr 	 <<
;D//D'$!5 ( 
 &--aAq!<DD\]ckr6   )r  r  r	  )Nr   r   Ng      ?learned_with_imagesF)NNrl   ru   s   @@r3   r   r     s$     "
 
B  r6   r   c                   v   a a ] tR tRt oRt. R
O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# )r  i  a  
A module to blend spatial and temporal features.

Parameters:
    alpha (`float`): The initial value of the blending factor.
    merge_strategy (`str`, *optional*, defaults to `learned_with_images`):
        The merge strategy to use for the temporal mixing.
    switch_spatial_to_temporal_mix (`bool`, *optional*, defaults to `False`):
        If `True`, switch the spatial and temporal mixing.
c                ,   < V ^8  d   QhRS[ RS[RS[/# )r+   r  r  r  )r/   r0   r.   )r1   r2   s   "r3   r4   AlphaBlender.__annotate__  s.     N NN N )-	Nr6   c                  < \         SV `  4        W n        W0n        W P                  9  d   \        R V P                   24      hV P                  R8X  d*   V P                  R\        P                  ! V.4      4       R# V P                  R8X  g   V P                  R8X  dG   V P                  R\        P                  P                  \        P                  ! V.4      4      4       R# \        RV P                   24      h)zmerge_strategy needs to be in fixed
mix_factorlearnedr  zUnknown merge strategy N)rD   rE   r  r  
strategiesrH   register_bufferrM   r\   register_parameterrI   	Parameter)rS   r  r  r  rT   s   &&&&r3   rE   AlphaBlender.__init__  s     	,.L+0=doo=NOPP')  u||UG/DE  I-1D1DH]1]##L%((2D2DU\\SXRYEZ2[\6t7J7J6KLMMr6   c                T   < V ^8  d   QhRS[ P                  RS[RS[ P                  /# )r+   r  ndimsrZ   r   )r1   r2   s   "r3   r4   r    s*      ell 3 5<< r6   c           	     N   V P                   R 8X  d   V P                  pV# V P                   R8X  d#   \        P                  ! V P                  4      pV# V P                   R8X  d   Vf   \	        R4      h\        P
                  ! VP                  4       \        P                  ! ^^VP                  R7      \        P                  ! V P                  4      R,          4      pV^8X  d   VR,          pV# V^8X  d   VP                  R	4      R
,          pV# \	        RV R24      h\        h)r  r  r  zMPlease provide image_only_indicator to use learned_with_images merge strategy)devicezUnexpected ndims z. Dimensions should be 3 or 5).N)r   Nr   NNr   )r   NN)r  r  rM   sigmoidrH   wherer.   onesr'  r   NotImplementedError)rS   r  r%  r  s   &&& r3   	get_alphaAlphaBlender.get_alpha  s   ')OOE6 3   I-MM$//2E0 -   $99#+ !pqqKK$))+

1a(<(C(CDdoo.y9E z45  !b)-8  !#4UG;X!YZZ &%r6   c                   < V ^8  d   QhRS[ P                  RS[ P                  RS[ P                  R,          RS[ P                  /# )r+   r  r  r  NrZ   r[   )r1   r2   s   "r3   r4   r    sH      << LL $llT1	
 
r6   c                    V P                  W1P                  4      pVP                  VP                  4      pV P                  '       d
   R V,
          pWA,          R V,
          V,          ,           pV# )r$   )r,  ndimtodtyper  )rS   r  r  r  r  r   s   &&&&  r3   rj   AlphaBlender.forward  sY     3^^D)...%KEu
 ::r6   )r  r  )r  r  r  )r  Fr   )rm   rn   ro   rp   rq   r  rE   r,  rj   rr   rs   rt   ru   s   @@r3   r  r    s6     	 =JN N( >  r6   r  )'	functoolsr   rM   torch.nnrI   torch.nn.functional
functionalr   utilsr   activationsr   attention_processorr   downsamplingr   r   r	   r
   r   normalizationr   
upsamplingr   r   r   r   r   r   Moduler   rw   r   r   r   r   r   r   r   r6   r3   <module>r@     s           ' ,  ( NBII Nb}BII }BO "))  H,0bii ,0^D		 DNY")) YzQRYY QhN299 Nr6   