Ë
    (täi   ã                   ó*  — d dl mZ d dlZd dlmZ d dlmc mZ ddlm	Z	 ddl
mZ ddlmZ ddlmZmZmZmZmZ dd	lmZ dd
lmZmZmZmZmZmZ  G d„ dej:                  «      Z G d„ dej:                  «      Zdej@                  dej@                  fd„Z! G d„ dej:                  «      Z" G d„ dej:                  «      Z# G d„ dej:                  «      Z$ G d„ dej:                  «      Z% G d„ dej:                  «      Z& G d„ dej:                  «      Z'y)é    )ÚpartialNé   )Ú	deprecateé   )Úget_activation)ÚSpatialNorm)ÚDownsample1DÚDownsample2DÚFirDownsample2DÚKDownsample2DÚdownsample_2d)ÚAdaGroupNorm)ÚFirUpsample2DÚKUpsample2DÚ
Upsample1DÚ
Upsample2DÚupfirdn2d_nativeÚupsample_2dc            "       óî   ‡ — e Zd ZdZddddddddd	d
ddddddœdededz  dedededededz  dedededededz  dededededz  f ˆ fd„Zde	j                  de	j                  d e	j                  fd!„Zˆ xZS )"Ú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`.
    NFç        é   é    ç�íµ ÷Æ°>ÚswishÚ	ada_groupç      ð?T)Úout_channelsÚconv_shortcutÚdropoutÚtemb_channelsÚgroupsÚ
groups_outÚepsÚnon_linearityÚtime_embedding_normÚoutput_scale_factorÚuse_in_shortcutÚupÚdownÚconv_shortcut_biasÚconv_2d_out_channelsÚin_channelsr   r   r    r!   r"   r#   r$   r%   r&   r'   r(   r)   r*   r+   r,   c                ó.  •— t         ‰| �  «        || _        |€|n|}|| _        || _        || _        || _        || _        |
| _        |€|}| j                  dk(  rt        ||||¬«      | _
        n9| j                  dk(  rt        ||«      | _
        nt        d| j                  › �«      ‚t        j                  ||ddd¬«      | _        | j                  dk(  rt        ||||¬«      | _        n9| j                  dk(  rt        ||«      | _        nt        d| j                  › �«      ‚t"        j                  j%                  |«      | _        |xs |}t        j                  ||ddd¬«      | _        t+        |	«      | _        d x| _        | _        | j
                  rt3        |d¬	«      | _        n | j                  rt5        |ddd
¬«      | _        |€| j                  |k7  n|| _        d | _        | j6                  r!t        j                  ||ddd|¬«      | _        y y )Nr   )r$   Úspatialz" unsupported time_embedding_norm: é   r   ©Úkernel_sizeÚstrideÚpaddingF©Úuse_convÚop©r6   r4   Únamer   ©r2   r3   r4   Ú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                    €úf/Volumes/fast/ai/experiments/MLX_z-image/.venv/lib/python3.12/site-packages/diffusers/models/resnet.pyr=   zResnetBlockCondNorm2D.__init__I   s   ø€ ô( 	‰ÑÔØ&ˆÔØ&2Ð&:‘{ÀˆØ(ˆÔØ!.ˆÔØˆŒØˆŒ	Ø#6ˆÔ Ø#6ˆÔ àÐØˆJà×#Ñ# {Ò2Ü% m°[À&ÈcÔRˆD�JØ×%Ñ%¨Ò2Ü$ [°-Ó@ˆD�JäÐAÀ$×BZÑBZÐA[Ð\Ó]Ð]ä—Y‘Y˜{¨LÀaÐPQÐ[\Ô]ˆŒ
à×#Ñ# {Ò2Ü% m°\À:ÐSVÔWˆD�JØ×%Ñ%¨Ò2Ü$ \°=ÓAˆD�JäÐAÀ$×BZÑBZÐA[Ð\Ó]Ð]ä—x‘x×'Ñ'¨Ó0ˆŒà3ÒC°|ÐÜ—Y‘Y˜|Ð-AÈqÐYZÐdeÔfˆŒ
ä*¨=Ó9ˆÔà*.Ð.ˆŒ˜œØ�7Š7Ü& {¸UÔCˆD�MØ�YŠYÜ*¨;ÀÐPQÐX\Ô]ˆDŒOàKZÐKb˜t×/Ñ/Ð3GÒGÐhwˆÔà!ˆÔØ×ÒÜ!#§¡ØØ$ØØØØ'ô"ˆDÕð  ó    Úinput_tensorÚtembÚreturnc                 óÖ  — t        |«      dkD  s|j                  dd «      �d}t        dd|«       |}| j                  ||«      }| j	                  |«      }| j
                  �U|j                  d   dk\  r |j                  «       }|j                  «       }| j                  |«      }| j                  |«      }n.| j                  �"| j                  |«      }| j                  |«      }| j                  |«      }| j                  ||«      }| j	                  |«      }| j                  |«      }| j                  |«      }| j                  �| j                  |«      }||z   | j                  z  }|S )Nr   ÚscaleúðThe `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   r?   rH   rI   ÚshapeÚ
contiguousrJ   rC   rD   r    rG   r   r'   )rK   rO   rP   ÚargsÚkwargsÚdeprecation_messageÚhidden_statesÚoutput_tensors           rM   ÚforwardzResnetBlockCondNorm2D.forward”   sT  € Üˆt‹9�qŠ=˜FŸJ™J w°Ó5ÐAð #UÐÜ�g˜wÐ(;Ô<à$ˆàŸ
™
 =°$Ó7ˆà×)Ñ)¨-Ó8ˆà�=‰=Ð$à×"Ñ" 1Ñ%¨Ò+Ø+×6Ñ6Ó8�Ø -× 8Ñ 8Ó :�ØŸ=™=¨Ó6ˆLØ ŸM™M¨-Ó8‰Mà�_‰_Ð(ØŸ?™?¨<Ó8ˆLØ ŸO™O¨MÓ:ˆMàŸ
™
 =Ó1ˆàŸ
™
 =°$Ó7ˆà×)Ñ)¨-Ó8ˆàŸ™ ]Ó3ˆØŸ
™
 =Ó1ˆà×ÑÐ)Ø×-Ñ-¨lÓ;ˆLà%¨Ñ5¸×9QÑ9QÑQˆàÐrN   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚintÚboolÚfloatÚstrr=   rE   ÚTensorr`   Ú__classcell__©rL   s   @rM   r   r   +   s(  ø„ ñðB $(Ø#ØØ ØØ!%ØØ$Ø#.Ø%(Ø'+ØØØ#'Ø+/ò%Ið ðIð ˜D‘jð	Ið
 ðIð ðIð ðIð ðIð ˜$‘JðIð ðIð ðIð !ðIð #ðIð  ™ðIð ðIð  ð!Ið" !ð#Ið$ " D™jõ%IðV% E§L¡Lð %¸¿¹ð %ÐZ_×ZfÑZf÷ %rN   r   c            (       ó  ‡ — e Zd ZdZddddddddd	dd
ddddddddœdededz  dedededededz  dedededededej                  dz  dededz  dededededz  f&ˆ fd „Z
d!ej                  d"ej                  d#ej                  fd$„Zˆ xZS )%ÚResnetBlock2Da9  
    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`.
    NFr   r   r   Tr   r   Údefaultr   )r   r   r    r!   r"   r#   Úpre_normr$   r%   Úskip_time_actr&   Úkernelr'   r(   r)   r*   r+   r,   r-   r   r   r    r!   r"   r#   ro   r$   r%   rp   r&   rq   r'   r(   r)   r*   r+   r,   c                ó’  •‡— t         ‰| �  «        |dk(  rt        d«      ‚|dk(  rt        d«      ‚d| _        || _        |€|n|}|| _        || _        || _        || _        || _	        || _
        || _        |€|}t        j                  j                  |||	d¬«      | _        t        j                   ||ddd¬	«      | _        |�r| j                  d
k(  rt        j$                  ||«      | _        nN| j                  dk(  rt        j$                  |d|z  «      | _        n t        d| j                  › d�«      ‚d | _        t        j                  j                  |||	d¬«      | _        t        j                  j+                  |«      | _        |xs |}t        j                   ||ddd¬	«      | _        t1        |
«      | _        d x| _        | _        | j                  rL|dk(  rdŠˆfd„| _        n“|dk(  r"t9        t:        j<                  dd¬«      | _        nlt?        |d¬«      | _        nY| j                  rM|dk(  rdŠˆfd„| _        n;|dk(  r"t9        t:        j@                  dd¬«      | _        ntC        |ddd¬«      | _        |€| j                  |k7  n|| _"        d | _#        | jD                  r!t        j                   ||ddd|¬«      | _#        y y )Nr   zkThis class cannot be used with `time_embedding_norm==ada_group`, please use `ResnetBlockCondNorm2D` insteadr/   ziThis class cannot be used with `time_embedding_norm==spatial`, please use `ResnetBlockCondNorm2D` insteadT©Ú
num_groupsÚnum_channelsr$   Úaffiner0   r   r1   rn   Úscale_shiftr   zunknown time_embedding_norm : Ú Úfir)r   r0   r0   r   c                 ó   •— t        | ‰¬«      S ©N)rq   )r   ©ÚxÚ
fir_kernels    €rM   Ú<lambda>z(ResnetBlock2D.__init__.<locals>.<lambda>$  s   ø€ ¬+°aÀ
Õ*KrN   Úsde_vpg       @Únearest)Úscale_factorÚmodeFr5   c                 ó   •— t        | ‰¬«      S r{   )r   r|   s    €rM   r   z(ResnetBlock2D.__init__.<locals>.<lambda>,  s   ø€ ¬M¸!ÀJÕ,OrN   )r2   r3   r7   r8   r   r:   )$r<   r=   r@   ro   r-   r   r>   r)   r*   r'   r&   rp   rE   rA   Ú	GroupNormr?   rB   rC   ÚLinearÚtime_emb_projrD   rF   r    rG   r   rH   rI   rJ   r   ÚFÚinterpolater   Ú
avg_pool2dr
   r(   r   )rK   r-   r   r   r    r!   r"   r#   ro   r$   r%   rp   r&   rq   r'   r(   r)   r*   r+   r,   r~   rL   s                       @€rM   r=   zResnetBlock2D.__init__Û   s¨  ù€ ô. 	‰ÑÔØ +Ò-ÜØ}óð ð  )Ò+ÜØ{óð ð ˆŒØ&ˆÔØ&2Ð&:‘{ÀˆØ(ˆÔØ!.ˆÔØˆŒØˆŒ	Ø#6ˆÔ Ø#6ˆÔ Ø*ˆÔàÐØˆJä—X‘X×'Ñ'°6ÈÐY\ÐeiÐ'ÓjˆŒ
ä—Y‘Y˜{¨LÀaÐPQÐ[\Ô]ˆŒ
àÐ$Ø×'Ñ'¨9Ò4Ü%'§Y¡Y¨}¸lÓ%K�Õ"Ø×)Ñ)¨]Ò:Ü%'§Y¡Y¨}¸aÀ,Ñ>NÓ%O�Õ"ä Ð#AÀ$×BZÑBZÐA[Ð[\Ð!]Ó^Ð^à!%ˆDÔä—X‘X×'Ñ'°:ÈLÐ^aÐjnÐ'ÓoˆŒ
ä—x‘x×'Ñ'¨Ó0ˆŒØ3ÒC°|ÐÜ—Y‘Y˜|Ð-AÈqÐYZÐdeÔfˆŒ
ä*¨=Ó9ˆÔà*.Ð.ˆŒ˜œØ�7Š7Ø˜ŠØ)�
Û K�•Ø˜8Ò#Ü '¬¯©ÀCÈiÔ X�•ä *¨;ÀÔ G�•Ø�YŠYØ˜ŠØ)�
Û"O�•Ø˜8Ò#Ü")¬!¯,©,ÀAÈaÔ"P�•ä".¨{ÀUÐTUÐ\`Ô"a�”àKZÐKb˜t×/Ñ/Ð3GÒGÐhwˆÔà!ˆÔØ×ÒÜ!#§¡ØØ$ØØØØ'ô"ˆDÕð  rN   rO   rP   rQ   c                 ó¦  — t        |«      dkD  s|j                  dd «      �d}t        dd|«       |}| j                  |«      }| j	                  |«      }| j
                  �U|j                  d   dk\  r |j                  «       }|j                  «       }| j                  |«      }| j                  |«      }n.| j                  �"| j                  |«      }| j                  |«      }| j                  |«      }| j                  �9| j                  s| j	                  |«      }| j                  |«      d d …d d …d d f   }| j                  dk(  r|�||z   }| j                  |«      }nr| j                  dk(  rR|€t        d| j                  › �«      ‚t        j                   |d	d
¬«      \  }}| j                  |«      }|d
|z   z  |z   }n| j                  |«      }| j	                  |«      }| j#                  |«      }| j%                  |«      }| j&                  �-| j(                  r|j                  «       }| j'                  |«      }||z   | j*                  z  }	|	S )Nr   rS   rT   rU   rV   rn   rw   z9 `temb` should not be None when `time_embedding_norm` is r   r   )Údim)rW   rX   r   r?   rH   rI   rY   rZ   rJ   rC   r‡   rp   r&   rD   r@   rE   Úchunkr    rG   r   Útrainingr'   )
rK   rO   rP   r[   r\   r]   r^   Ú
time_scaleÚ
time_shiftr_   s
             rM   r`   zResnetBlock2D.forward?  sB  € Üˆt‹9�qŠ=˜FŸJ™J w°Ó5ÐAð #UÐÜ�g˜wÐ(;Ô<à$ˆàŸ
™
 =Ó1ˆØ×)Ñ)¨-Ó8ˆà�=‰=Ð$à×"Ñ" 1Ñ%¨Ò+Ø+×6Ñ6Ó8�Ø -× 8Ñ 8Ó :�ØŸ=™=¨Ó6ˆLØ ŸM™M¨-Ó8‰MØ�_‰_Ð(ØŸ?™?¨<Ó8ˆLØ ŸO™O¨MÓ:ˆMàŸ
™
 =Ó1ˆà×ÑÐ)Ø×%Ò%Ø×(Ñ(¨Ó.�Ø×%Ñ% dÓ+ªAªq°$¸Ð,<Ñ=ˆDà×#Ñ# yÒ0ØÐØ -°Ñ 4�Ø ŸJ™J }Ó5‰MØ×%Ñ%¨Ò6Øˆ|Ü ØOÐPT×PhÑPhÐOiÐjóð ô &+§[¡[°°q¸aÔ%@Ñ"ˆJ˜
Ø ŸJ™J }Ó5ˆMØ)¨Q°©^Ñ<¸zÑI‰Mà ŸJ™J }Ó5ˆMà×)Ñ)¨-Ó8ˆàŸ™ ]Ó3ˆØŸ
™
 =Ó1ˆà×ÑÐ)ð �}Š}Ø+×6Ñ6Ó8�Ø×-Ñ-¨lÓ;ˆLà%¨Ñ5¸×9QÑ9QÑQˆàÐrN   )ra   rb   rc   rd   re   rf   rg   rh   rE   ri   r=   r`   rj   rk   s   @rM   rm   rm   ¼   s[  ø„ ñðD $(Ø#ØØ ØØ!%ØØØ$Ø#Ø#,Ø&*Ø%(Ø'+ØØØ#'Ø+/ò+bð ðbð ˜D‘jð	bð
 ðbð ðbð ðbð ðbð ˜$‘Jðbð ðbð ðbð ðbð ðbð !ðbð —‘˜tÑ#ðbð  #ð!bð"  ™ð#bð$ ð%bð& ð'bð( !ð)bð* " D™jõ+bðH: E§L¡Lð :¸¿¹ð :ÐZ_×ZfÑZf÷ :rN   rm   ÚtensorrQ   c                 ó  — t        | j                  «      dk(  r| d d …d d …d f   S t        | j                  «      dk(  r| d d …d d …d d d …f   S t        | j                  «      dk(  r| d d …d d …dd d …f   S t        dt        | «      › d�«      ‚)Nr   r0   é   r   z`len(tensor)`: z has to be 2, 3 or 4.)rW   rY   r@   )r‘   s    rM   Úrearrange_dimsr”   }  s…   € Ü
ˆ6�<‰<Ó˜AÒØ’aš˜D�jÑ!Ð!Ü
ˆ6�<‰<Ó˜AÒØ’aš˜D¢!�mÑ$Ð$Ü	ˆV�\‰\Ó	˜aÒ	Ø’aš˜Ašq�jÑ!Ð!ä˜?¬3¨v«;¨-Ð7LÐMÓNÐNrN   c                   ó†   ‡ — e Zd ZdZ	 	 ddededeeeef   z  dedef
ˆ fd„Zdej                  d	ej                  fd
„Z
ˆ xZS )ÚConv1dBlocka˜  
    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.
    Úinp_channelsr   r2   Ún_groupsÚ
activationc                 óº   •— t         ‰| �  «        t        j                  ||||dz  ¬«      | _        t        j
                  ||«      | _        t        |«      | _        y )Nr   ©r4   )	r<   r=   rA   ÚConv1dÚconv1dr…   Ú
group_normr   Úmish)rK   r—   r   r2   r˜   r™   rL   s         €rM   r=   zConv1dBlock.__init__”  sK   ø€ ô 	‰ÑÔä—i‘i ¨l¸KÐQ\Ð`aÑQaÔbˆŒÜŸ,™, x°Ó>ˆŒÜ" :Ó.ˆ�	rN   ÚinputsrQ   c                 ó˜   — | j                  |«      }t        |«      }| j                  |«      }t        |«      }| j                  |«      }|S ©N)r�   r”   rž   rŸ   )rK   r    Úintermediate_reprÚoutputs       rM   r`   zConv1dBlock.forward¢  sM   € Ø ŸK™K¨Ó/ÐÜ*Ð+<Ó=ÐØ ŸO™OÐ,=Ó>ÐÜ*Ð+<Ó=ÐØ—‘Ð,Ó-ˆØˆrN   )é   rŸ   ©ra   rb   rc   rd   re   Útuplerh   r=   rE   ri   r`   rj   rk   s   @rM   r–   r–   ˆ  sm   ø„ ñ	ð  Ø ñ/àð/ð ð/ð ˜5  c ™?Ñ*ð	/ð
 ð/ð õ/ð˜eŸl™lð ¨u¯|©|÷ rN   r–   c                   óž   ‡ — e Zd ZdZ	 	 ddedededeeeef   z  def
ˆ fd„Zdej                  d	ej                  d
ej                  fd„Z
ˆ xZS )ÚResidualTemporalBlock1Da•  
    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.
    r—   r   Ú	embed_dimr2   r™   c                 ó6  •— t         ‰| �  «        t        |||«      | _        t        |||«      | _        t        |«      | _        t        j                  ||«      | _	        ||k7  rt        j                  ||d«      | _        y t        j                  «       | _        y )Nr   )r<   r=   r–   Úconv_inÚconv_outr   Útime_emb_actrA   r†   Útime_embrœ   ÚIdentityÚresidual_conv)rK   r—   r   rª   r2   r™   rL   s         €rM   r=   z ResidualTemporalBlock1D.__init__¸  s…   ø€ ô 	‰ÑÔÜ" <°¸{ÓKˆŒÜ# L°,ÀÓLˆŒä*¨:Ó6ˆÔÜŸ	™	 )¨\Ó:ˆŒð 9EÈÒ8TŒB�I‰I�l L°!Ó4ð 	ÕÜZ\×ZeÑZeÓZgð 	ÕrN   r    ÚtrQ   c                 óÊ   — | j                  |«      }| j                  |«      }| j                  |«      t        |«      z   }| j	                  |«      }|| j                  |«      z   S )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±   )rK   r    r²   Úouts       rM   r`   zResidualTemporalBlock1D.forwardË  s^   € ð ×Ñ˜aÓ ˆØ�M‰M˜!ÓˆØ�l‰l˜6Ó"¤^°AÓ%6Ñ6ˆØ�m‰m˜CÓ ˆØ�T×'Ñ'¨Ó/Ñ/Ð/rN   )é   rŸ   r¦   rk   s   @rM   r©   r©   ¬  sx   ø„ ñ	ð  ./Ø ñ
àð
ð ð
ð ð	
ð
 ˜5  c ™?Ñ*ð
ð õ
ð&0˜eŸl™lð 0¨u¯|©|ð 0ÀÇÁ÷ 0rN   r©   c            	       ó€   ‡ — e Zd ZdZ	 	 	 ddededz  dedefˆ fd„Zddej                  d	ed
ej                  fd„Z	ˆ xZ
S )ÚTemporalConvLayeraà  
    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.
    NÚin_dimÚout_dimr    Únorm_num_groupsc                 ób  •— t         ‰| �  «        |xs |}|| _        || _        t	        j
                  t	        j                  ||«      t	        j                  «       t	        j                  ||dd¬«      «      | _	        t	        j
                  t	        j                  ||«      t	        j                  «       t	        j                  |«      t	        j                  ||dd¬«      «      | _        t	        j
                  t	        j                  ||«      t	        j                  «       t	        j                  |«      t	        j                  ||dd¬«      «      | _        t	        j
                  t	        j                  ||«      t	        j                  «       t	        j                  |«      t	        j                  ||dd¬«      «      | _        t        j                  j                  | j                  d   j                   «       t        j                  j                  | j                  d   j"                  «       y )N©r0   r   r   )r   r   r   r›   éÿÿÿÿ)r<   r=   r¸   r¹   rA   Ú
Sequentialr…   ÚSiLUÚConv3drC   rF   rG   Úconv3Úconv4ÚinitÚzeros_Úweightr;   )rK   r¸   r¹   r    rº   rL   s        €rM   r=   zTemporalConvLayer.__init__æ  sv  ø€ ô 	‰ÑÔØÒ#˜VˆØˆŒØˆŒô —]‘]Ü�L‰L˜¨&Ó1Ü�G‰G‹IÜ�I‰I�f˜g y¸)ÔDó
ˆŒ
ô
 —]‘]Ü�L‰L˜¨'Ó2Ü�G‰G‹IÜ�J‰J�wÓÜ�I‰I�g˜v y¸)ÔDó	
ˆŒ
ô —]‘]Ü�L‰L˜¨'Ó2Ü�G‰G‹IÜ�J‰J�wÓÜ�I‰I�g˜v y¸)ÔDó	
ˆŒ
ô —]‘]Ü�L‰L˜¨'Ó2Ü�G‰G‹IÜ�J‰J�wÓÜ�I‰I�g˜v y¸)ÔDó	
ˆŒ
ô 	�‰�‰�t—z‘z "‘~×,Ñ,Ô-Ü
�‰�‰�t—z‘z "‘~×*Ñ*Õ+rN   r^   Ú
num_framesrQ   c                 ó¼  — |d d d …f   j                  d|f|j                  dd  z   «      j                  ddddd«      }|}| j                  |«      }| j	                  |«      }| j                  |«      }| j                  |«      }||z   }|j                  ddddd«      j                  |j                  d   |j                  d   z  df|j                  dd  z   «      }|S )Nr½   r   r   r   r0   r“   )ÚreshaperY   ÚpermuterC   rG   rÁ   rÂ   )rK   r^   rÆ   Úidentitys       rM   r`   zTemporalConvLayer.forward  sú   € à˜$¢˜'Ñ"×*Ñ*¨B°
Ð+;¸m×>QÑ>QÐRSÐRTÐ>UÑ+UÓV×^Ñ^Ð_`ÐbcÐefÐhiÐklÓmð 	ð !ˆØŸ
™
 =Ó1ˆØŸ
™
 =Ó1ˆØŸ
™
 =Ó1ˆØŸ
™
 =Ó1ˆà  =Ñ0ˆà%×-Ñ-¨a°°A°q¸!Ó<×DÑDØ× Ñ  Ñ# m×&9Ñ&9¸!Ñ&<Ñ<¸bÐAÀM×DWÑDWÐXYÐXZÐD[Ñ[ó
ˆð ÐrN   )Nr   r   )r   ©ra   rb   rc   rd   re   rg   r=   rE   ri   r`   rj   rk   s   @rM   r·   r·   Û  se   ø„ ñð #ØØ!ñ',àð',ð �t‘ð',ð ð	',ð
 õ',ñR U§\¡\ð ¸sð È5Ï<É<÷ rN   r·   c            	       ó’   ‡ — e Zd ZdZ	 	 	 ddededz  dedefˆ fd„Zdej                  d	ej                  d
ej                  fd„Z	ˆ xZ
S )ÚTemporalResnetBlockaÞ  
    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.
    Nr-   r   r!   r$   c                 óØ  •— t         ‰| �  «        || _        |€|n|}|| _        d}|D �cg c]  }|dz  ‘Œ	 }}t        j
                  j                  d||d¬«      | _        t        j                  |||d|¬«      | _	        |�t        j                  ||«      | _        nd | _        t        j
                  j                  d||d¬«      | _        t        j
                  j                  d«      | _        t        j                  |||d|¬«      | _        t!        d	«      | _        | j                  |k7  | _        d | _        | j$                  r t        j                  ||ddd
¬«      | _        y y c c}w )Nr¼   r   r   Trs   r   r1   r   Úsilur   )r<   r=   r-   r   rE   rA   r…   r?   rÀ   rC   r†   r‡   rD   rF   r    rG   r   rH   r(   r   )	rK   r-   r   r!   r$   r2   Úkr4   rL   s	           €rM   r=   zTemporalResnetBlock.__init__.  s`  ø€ ô 	‰ÑÔØ&ˆÔØ&2Ð&:‘{ÀˆØ(ˆÔàˆÙ#.Ó/¡;˜a�1˜“6 ;ˆÐ/ä—X‘X×'Ñ'°2ÀKÐUXÐaeÐ'ÓfˆŒ
Ü—Y‘YØØØ#ØØô
ˆŒ
ð Ð$Ü!#§¡¨=¸,Ó!GˆDÕà!%ˆDÔä—X‘X×'Ñ'°2ÀLÐVYÐbfÐ'ÓgˆŒ
ä—x‘x×'Ñ'¨Ó,ˆŒÜ—Y‘YØØØ#ØØô
ˆŒ
ô +¨6Ó2ˆÔà#×/Ñ/°<Ñ?ˆÔà!ˆÔØ×ÒÜ!#§¡ØØØØØô"ˆDÕð  ùòA 0s   ªE'rO   rP   rQ   c                 óè  — |}| j                  |«      }| j                  |«      }| j                  |«      }| j                  �J| j                  |«      }| j                  |«      d d …d d …d d …d d f   }|j	                  ddddd«      }||z   }| j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }| j                  �| j                  |«      }||z   }|S )Nr   r   r   r0   r“   )	r?   rH   rC   r‡   rÉ   rD   r    rG   r   )rK   rO   rP   r^   r_   s        rM   r`   zTemporalResnetBlock.forwardd  sõ   € Ø$ˆàŸ
™
 =Ó1ˆØ×)Ñ)¨-Ó8ˆØŸ
™
 =Ó1ˆà×ÑÐ)Ø×$Ñ$ TÓ*ˆDØ×%Ñ% dÓ+ªAªq²!°T¸4Ð,?Ñ@ˆDØ—<‘<  1 a¨¨AÓ.ˆDØ)¨DÑ0ˆMàŸ
™
 =Ó1ˆØ×)Ñ)¨-Ó8ˆØŸ™ ]Ó3ˆØŸ
™
 =Ó1ˆà×ÑÐ)Ø×-Ñ-¨lÓ;ˆLà$ }Ñ4ˆàÐrN   )Nr   r   rË   rk   s   @rM   rÍ   rÍ   "  si   ø„ ñ	ð $(Ø Øñ4àð4ð ˜D‘jð4ð ð	4ð
 õ4ðl E§L¡Lð ¸¿¹ð ÈÏÉ÷ rN   rÍ   c                   ó¾   ‡ — e Zd ZdZ	 	 	 	 	 	 	 ddededz  dedededz  ded	efˆ fd
„Z	 	 ddej                  dej                  dz  dej                  dz  fd„Z
ˆ xZS )ÚSpatioTemporalResBlockaé  
    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.
    Nr-   r   r!   r$   Útemporal_epsÚmerge_factorÚswitch_spatial_to_temporal_mixc	                 ó°   •— t         ‰	| �  «        t        ||||¬«      | _        t	        |�|n||�|n|||�|n|¬«      | _        t        |||¬«      | _        y )N)r-   r   r!   r$   )ÚalphaÚmerge_strategyrÖ   )r<   r=   rm   Úspatial_res_blockrÍ   Útemporal_res_blockÚAlphaBlenderÚ
time_mixer)
rK   r-   r   r!   r$   rÔ   rÕ   rÙ   rÖ   rL   s
            €rM   r=   zSpatioTemporalResBlock.__init__‘  sp   ø€ ô 	‰ÑÔä!.Ø#Ø%Ø'Øô	"
ˆÔô #6Ø(4Ð(@™ÀkØ)5Ð)A™À{Ø'Ø ,Ð 8‘¸cô	#
ˆÔô 'ØØ)Ø+Iô
ˆ�rN   r^   rP   Úimage_only_indicatorc                 óô  — |j                   d   }| j                  ||«      }|j                   \  }}}}||z  }	|d d d …f   j                  |	||||«      j                  ddddd«      }
|d d d …f   j                  |	||||«      j                  ddddd«      }|�|j                  |	|d«      }| j	                  ||«      }| j                  |
||¬«      }|j                  ddddd«      j                  ||||«      }|S )Nr½   r   r   r   r0   r“   )Ú	x_spatialÚ
x_temporalrÞ   )rY   rÚ   rÈ   rÉ   rÛ   rÝ   )rK   r^   rP   rÞ   rÆ   Úbatch_framesÚchannelsÚheightÚwidthÚ
batch_sizeÚhidden_states_mixs              rM   r`   zSpatioTemporalResBlock.forward²  s>  € ð *×/Ñ/°Ñ3ˆ
Ø×.Ñ.¨}¸dÓCˆà0=×0CÑ0CÑ-ˆ�h ¨Ø! ZÑ/ˆ
ð ˜$¢˜'Ñ"×*Ñ*¨:°zÀ8ÈVÐUZÓ[×cÑcÐdeÐghÐjkÐmnÐpqÓrð 	ð ˜$¢˜'Ñ"×*Ñ*¨:°zÀ8ÈVÐUZÓ[×cÑcÐdeÐghÐjkÐmnÐpqÓrð 	ð ÐØ—<‘< 
¨J¸Ó;ˆDà×/Ñ/°¸tÓDˆØŸ™Ø'Ø$Ø!5ð (ó 
ˆð &×-Ñ-¨a°°A°q¸!Ó<×DÑDÀ\ÐS[Ð]cÐejÓkˆØÐrN   )Nr   r   Ng      à?Úlearned_with_imagesF)NN)ra   rb   rc   rd   re   rg   rf   r=   rE   ri   r`   rj   rk   s   @rM   rÓ   rÓ     s°   ø„ ñð( $(Ø ØØ%)Ø!Ø,Ø/4ñ
àð
ð ˜D‘jð
ð ð	
ð
 ð
ð ˜d‘lð
ð ð
ð )-õ
ðH %)Ø48ñ	à—|‘|ðð �l‰l˜TÑ!ðð $Ÿl™l¨TÑ1÷	rN   rÓ   c            	       óì   ‡ — e Zd ZdZg d¢Z	 	 ddededefˆ fd„Zde	j                  ded	e	j                  fd
„Z	 dde	j                  de	j                  de	j                  dz  d	e	j                  fd„Zˆ xZS )rÜ   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.
    )ÚlearnedÚfixedrè   rØ   rÙ   rÖ   c                 óè  •— t         ‰| �  «        || _        || _        || j                  vrt        d| j                  › �«      ‚| j                  dk(  r'| j                  dt        j                  |g«      «       y | j                  dk(  s| j                  dk(  rD| j                  dt        j                  j                  t        j                  |g«      «      «       y t        d| j                  › �«      ‚)Nzmerge_strategy needs to be in rë   Ú
mix_factorrê   rè   zUnknown merge strategy )r<   r=   rÙ   rÖ   Ú
strategiesr@   Úregister_bufferrE   ri   Úregister_parameterrA   Ú	Parameter)rK   rØ   rÙ   rÖ   rL   s       €rM   r=   zAlphaBlender.__init__á  sÏ   ø€ ô 	‰ÑÔØ,ˆÔØ.LˆÔ+à §¡Ñ0ÜÐ=¸d¿o¹oÐ=NÐOÓPÐPà×Ñ 'Ò)Ø× Ñ  ¬u¯|©|¸U¸GÓ/DÕEØ× Ñ  IÒ-°×1DÑ1DÐH]Ò1]Ø×#Ñ# L´%·(±(×2DÑ2DÄUÇ\Á\ÐSXÐRYÓEZÓ2[Õ\äÐ6°t×7JÑ7JÐ6KÐLÓMÐMrN   rÞ   ÚndimsrQ   c                 ó2  — | j                   dk(  r| j                  }|S | j                   dk(  r!t        j                  | j                  «      }|S | j                   dk(  r¶|€t	        d«      ‚t        j
                  |j                  «       t        j                  dd|j                  ¬«      t        j                  | j                  «      d   «      }|dk(  r|d d …d d d …d d f   }|S |d	k(  r|j                  d
«      d d …d d f   }|S t	        d|› d�«      ‚t        ‚)Nrë   rê   rè   zMPlease provide image_only_indicator to use learned_with_images merge strategyr   )Údevice).Nrµ   r0   r½   zUnexpected ndims z. Dimensions should be 3 or 5)rÙ   rí   rE   Úsigmoidr@   Úwhererf   Úonesrô   rÈ   ÚNotImplementedError)rK   rÞ   rò   rØ   s       rM   Ú	get_alphazAlphaBlender.get_alphaõ  s   € Ø×Ñ 'Ò)Ø—O‘OˆEð6 ˆð3 × Ñ  IÒ-Ü—M‘M $§/¡/Ó2ˆEð0 ˆð- × Ñ Ð$9Ò9Ø#Ð+Ü Ð!pÓqÐqä—K‘KØ$×)Ñ)Ó+Ü—
‘
˜1˜aÐ(<×(CÑ(CÔDÜ—‘˜dŸo™oÓ.¨yÑ9óˆEð ˜ŠzØša ¢q¨$°Ð4Ñ5�ð ˆð ˜!’ØŸ™ bÓ)ª!¨T°4¨-Ñ8�ð ˆô !Ð#4°U°GÐ;XÐ!YÓZÐZô &Ð%rN   Nrà   rá   c                 ó²   — | j                  ||j                  «      }|j                  |j                  «      }| j                  rd|z
  }||z  d|z
  |z  z   }|S )Nr   )rù   ÚndimÚtoÚdtyperÖ   )rK   rà   rá   rÞ   rØ   r}   s         rM   r`   zAlphaBlender.forward  sZ   € ð —‘Ð3°Y·^±^ÓDˆØ—‘˜Ÿ™Ó)ˆà×.Ò.Ø˜%‘KˆEà�IÑ  u¡°
Ñ :Ñ:ˆØˆrN   )rè   Fr¢   )ra   rb   rc   rd   rî   rg   rh   rf   r=   rE   ri   re   rù   r`   rj   rk   s   @rM   rÜ   rÜ   Ó  s¤   ø„ ñ	ò =€Jð
 4Ø/4ñ	NàðNð ðNð )-õ	Nð(¨e¯l©lð À3ð È5Ï<É<ó ðF 59ñ	à—<‘<ðð —L‘Lðð $Ÿl™l¨TÑ1ð	ð
 
�‰÷rN   rÜ   )(Ú	functoolsr   rE   Útorch.nnrA   Útorch.nn.functionalÚ
functionalrˆ   Úutilsr   Úactivationsr   Úattention_processorr   Údownsamplingr	   r
   r   r   r   Únormalizationr   Ú
upsamplingr   r   r   r   r   r   ÚModuler   rm   ri   r”   r–   r©   r·   rÍ   rÓ   rÜ   © rN   rM   Ú<module>r
     sé   ðõ  ã Ý ß Ð å Ý 'Ý ,÷õ õ (÷÷ ôN˜BŸI™Iô Nôb}�B—I‘Iô }ðBO˜5Ÿ<™<ð O¨E¯L©Ló Oô �"—)‘)ô  ôH,0˜bŸi™iô ,0ô^D˜Ÿ	™	ô DôNY˜"Ÿ)™)ô YôzQ˜RŸY™Yô QôhN�2—9‘9õ NrN   