Ë
    óÿæib•  ã                   ó(  — d dl Z d dlZd dlmZmZmZmZ d dlZd dlmZ d dl	m
Z  G d„ dej                  j                  «      Z G d„ dej                  j                  «      Z G d	„ d
ej                  j                  «      Z G d„ dej                  j                  «      Z G d„ dej                  j                  «      Z G d„ dej                  j                  «      Z G d„ dej                  «      Z G d„ dej                  «      Zdej*                  dededej*                  fd„Zd„ Zd'dej*                  dedededej*                  f
d„Zd(d ej*                  ded!ededej*                  f
d"„Zd#ee   defd$„Zd#ee   defd%„Zd#ee   defd&„Zy))é    N)ÚAnyÚDictÚListÚOptional)Únn)Ú
functionalc            	       ó˜   ‡ — e Zd ZdZddedededefˆ fd„Zede	j                  fd„«       Zd	e	j                  de	j                  fd
„Zˆ xZS )Ú_ScaledEmbeddingaF  Make continuous embeddings and boost learning rate

    Args:
        num_embeddings (int): number of embeddings
        embedding_dim (int): embedding dimensions
        scale (float, optional): amount to scale learning rate (Default: 10.0)
        smooth (bool, optional): choose to apply smoothing (Default: ``False``)
    Únum_embeddingsÚembedding_dimÚscaleÚsmoothc                 óÎ  •— t         ‰| �  «        t        j                  ||«      | _        |r‰t        j                  | j                  j                  j                  d¬«      }|t        j                  d|dz   «      j                  «       d d …d f   z  }|| j                  j                  j                  d d  | j                  j                  xj                  |z  c_        || _        y )Nr   ©Údimé   )ÚsuperÚ__init__r   Ú	EmbeddingÚ	embeddingÚtorchÚcumsumÚweightÚdataÚarangeÚsqrtr   )Úselfr   r   r   r   r   Ú	__class__s         €úo/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchaudio/models/_hdemucs.pyr   z_ScaledEmbedding.__init__-   s«   ø€ Ü‰ÑÔÜŸ™ n°mÓDˆŒÙÜ—\‘\ $§.¡.×"7Ñ"7×"<Ñ"<À!ÔDˆFàœeŸl™l¨1¨n¸qÑ.@ÓA×FÑFÓHÊÈDÈÑQÑQˆFØ,2ˆD�N‰N×!Ñ!×&Ñ&¡qÐ)Ø�‰×Ñ×"Ò" eÑ+Õ"Øˆ�
ó    Úreturnc                 óH   — | j                   j                  | j                  z  S ©N)r   r   r   )r   s    r   r   z_ScaledEmbedding.weight8   s   € à�~‰~×$Ñ$ t§z¡zÑ1Ð1r    Úxc                 óB   — | j                  |«      | j                  z  }|S )zøForward pass for embedding with scale.
        Args:
            x (torch.Tensor): input tensor of shape `(num_embeddings)`

        Returns:
            (Tensor):
                Embedding output of shape `(num_embeddings, embedding_dim)`
        )r   r   )r   r$   Úouts      r   Úforwardz_ScaledEmbedding.forward<   s    € ð �n‰n˜QÓ $§*¡*Ñ,ˆØˆ
r    )g      $@F)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚintÚfloatÚboolr   Úpropertyr   ÚTensorr   r'   Ú__classcell__©r   s   @r   r
   r
   #   sd   ø„ ññ	 sð 	¸3ð 	Àuð 	Ð]aõ 	ð ð2˜Ÿ™ò 2ó ð2ð
˜Ÿ™ð 
¨%¯,©,÷ 
r    r
   c                   óÌ   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 ddededededededed	ed
edeeee	f      defˆ fd„Z
ddej                  deej                     dej                  fd„Zˆ xZS )Ú
_HEncLayerat  Encoder layer. This used both by the time and the frequency branch.
    Args:
        chin (int): number of input channels.
        chout (int): number of output channels.
        kernel_size (int, optional): Kernel size for encoder (Default: 8)
        stride (int, optional): Stride for encoder layer (Default: 4)
        norm_groups (int, optional): number of groups for group norm. (Default: 4)
        empty (bool, optional): used to make a layer with just the first conv. this is used
            before merging the time and freq. branches. (Default: ``False``)
        freq (bool, optional): boolean for whether conv layer is for frequency domain (Default: ``True``)
        norm_type (string, optional): Norm type, either ``group_norm `` or ``none`` (Default: ``group_norm``)
        context (int, optional): context size for the 1x1 conv. (Default: 0)
        dconv_kw (Dict[str, Any] or None, optional): dictionary of kwargs for the DConv class. (Default: ``None``)
        pad (bool, optional): true to pad the input. Padding is done so that the output size is
            always the input size / stride. (Default: ``True``)
    ÚchinÚchoutÚkernel_sizeÚstrideÚnorm_groupsÚemptyÚfreqÚ	norm_typeÚcontextÚdconv_kwÚpadc                 ó\  •‡— t         ‰| �  «        |
€i }
d„ }|dk(  rˆfd„}|r|dz  nd}t        j                  }|| _        || _        || _        || _        || _        |r|dg}|dg}|dg}t        j                  } ||||||«      | _
         ||«      | _        | j                  rLt        j                  «       | _        t        j                  «       | _        t        j                  «       | _        y  ||d|z  dd|	z  z   d|	«      | _         |d|z  «      | _        t!        |fi |
¤Ž| _        y )Nc                 ó*   — t        j                  «       S r#   ©r   ÚIdentity©Úds    r   Ú<lambda>z%_HEncLayer.__init__.<locals>.<lambda>m   ó
   € œBŸK™KœMr    Ú
group_normc                 ó0   •— t        j                  ‰| «      S r#   ©r   Ú	GroupNorm©rE   r9   s    €r   rF   z%_HEncLayer.__init__.<locals>.<lambda>o   ó   ø€ ¤§¡¨[¸!Ô <r    é   r   r   é   )r   r   r   ÚConv1dr;   r7   r8   r:   r?   ÚConv2dÚconvÚnorm1rC   ÚrewriteÚnorm2ÚdconvÚ_DConv)r   r5   r6   r7   r8   r9   r:   r;   r<   r=   r>   r?   Únorm_fnÚpad_valÚklassr   s        `         €r   r   z_HEncLayer.__init__\   s  ù€ ô 	‰ÑÔØÐØˆHÙ)ˆØ˜Ò$Û<ˆGÙ&)�+ Ò"¨qˆÜ—	‘	ˆØˆŒ	Ø&ˆÔØˆŒØˆŒ
ØˆŒÙØ&¨Ð*ˆKØ˜a�[ˆFØ �lˆGÜ—I‘IˆEÙ˜$  {°F¸GÓDˆŒ	Ù˜U“^ˆŒ
à�:Š:ÜŸ;™;›=ˆDŒLÜŸ™›ˆDŒJÜŸ™›ˆD�Já  ¨¨E©	°1°q¸7±{±?ÀAÀwÓOˆDŒLÙ   U¡Ó+ˆDŒJÜ Ñ2¨Ñ2ˆD�Jr    r$   Úinjectr!   c                 ó  — | j                   s7|j                  «       dk(  r$|j                  \  }}}}|j                  |d|«      }| j                   sS|j                  d   }|| j                  z  dk(  s2t        j                  |d| j                  || j                  z  z
  f«      }| j                  |«      }| j                  r|S |�a|j                  d   |j                  d   k7  rt        d«      ‚|j                  «       dk(  r|j                  «       dk(  r|dd…dd…df   }||z   }t        j                  | j                  |«      «      }| j                   rn|j                  \  }}}}|j                  dddd«      j                  d||«      }| j                  |«      }|j                  ||||«      j                  dddd«      }n| j                  |«      }| j                  | j!                  |«      «      }	t        j"                  |	d¬	«      }	|	S )
a]  Forward pass for encoding layer.

        Size depends on whether frequency or time

        Args:
            x (torch.Tensor): tensor input of shape `(B, C, F, T)` for frequency and shape
                `(B, C, T)` for time
            inject (torch.Tensor, optional): on last layer, combine frequency and time branches through inject param,
                same shape as x (default: ``None``)

        Returns:
            Tensor
                output tensor after encoder layer of shape `(B, C, F / stride, T)` for frequency
                    and shape `(B, C, ceil(T / stride))` for time
        rN   éÿÿÿÿr   NzInjection shapes do not aligné   rO   r   r   )r;   r   ÚshapeÚviewr8   ÚFr?   rR   r:   Ú
ValueErrorÚgelurS   ÚpermuteÚreshaperV   rU   rT   Úglu)
r   r$   r[   ÚBÚCÚFrÚTÚleÚyÚzs
             r   r'   z_HEncLayer.forwardˆ   s·  € ð" �yŠy˜QŸU™U›W¨š\ØŸ'™'‰KˆAˆq�"�aØ—‘�q˜"˜aÓ ˆAà�yŠyØ—‘˜‘ˆBØ˜Ÿ™Ñ# qÒ(Ü—E‘E˜!˜a §¡°°T·[±[Ñ0@Ñ!AÐBÓC�Ø�I‰I�a‹LˆØ�:Š:ØˆHØÐØ�|‰|˜BÑ 1§7¡7¨2¡;Ò.Ü Ð!@ÓAÐAØ�z‰z‹|˜qÒ  Q§U¡U£W°¢\Ø¢¢1 d 
Ñ+�Ø�F‘
ˆAÜ�F‰F�4—:‘:˜a“=Ó!ˆØ�9Š9ØŸ'™'‰KˆAˆq�"�aØ—	‘	˜!˜Q  1Ó%×-Ñ-¨b°!°QÓ7ˆAØ—
‘
˜1“ˆAØ—‘�q˜"˜a Ó#×+Ñ+¨A¨q°!°QÓ7‰Aà—
‘
˜1“ˆAØ�J‰J�t—|‘| A“Ó'ˆÜ�E‰E�!˜ŒOˆØˆr    )	é   rN   rN   FTrH   r   NTr#   ©r(   r)   r*   r+   r,   r.   Ústrr   r   r   r   r   r0   r'   r1   r2   s   @r   r4   r4   I   sÒ   ø„ ñð* ØØØØØ%ØØ-1Øñ*3àð*3ð ð*3ð ð	*3ð
 ð*3ð ð*3ð ð*3ð ð*3ð ð*3ð ð*3ð ˜4  S ™>Ñ*ð*3ð õ*3ñX,˜Ÿ™ð ,¨x¸¿¹Ñ/Eð ,ÐQV×Q]ÑQ]÷ ,r    r4   c                   ó¸   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 	 ddededededededed	ed
ededeeee	f      defˆ fd„Z
dej                  deej                     fd„Zˆ xZS )Ú
_HDecLayera»  Decoder layer. This used both by the time and the frequency branches.
    Args:
        chin (int): number of input channels.
        chout (int): number of output channels.
        last (bool, optional): whether current layer is final layer (Default: ``False``)
        kernel_size (int, optional): Kernel size for encoder (Default: 8)
        stride (int): Stride for encoder layer (Default: 4)
        norm_groups (int, optional): number of groups for group norm. (Default: 1)
        empty (bool, optional): used to make a layer with just the first conv. this is used
            before merging the time and freq. branches. (Default: ``False``)
        freq (bool, optional): boolean for whether conv layer is for frequency (Default: ``True``)
        norm_type (str, optional): Norm type, either ``group_norm `` or ``none`` (Default: ``group_norm``)
        context (int, optional): context size for the 1x1 conv. (Default: 1)
        dconv_kw (Dict[str, Any] or None, optional): dictionary of kwargs for the DConv class. (Default: ``None``)
        pad (bool, optional): true to pad the input. Padding is done so that the output size is
            always the input size / stride. (Default: ``True``)
    r5   r6   Úlastr7   r8   r9   r:   r;   r<   r=   r>   r?   c                 óŽ  •‡— t         ‰| �  «        |€i }d„ }|	dk(  rˆfd„}|r||z
  dz  dk7  rt        d«      ‚||z
  dz  }nd}|| _        || _        || _        || _        || _        || _        || _	        t        j                  }t        j                  }|r(|dg}|dg}t        j                  }t        j                  } |||||«      | _         ||«      | _        | j                  r3t        j"                  «       | _        t        j"                  «       | _        y  ||d|z  dd|
z  z   d|
«      | _         |d|z  «      | _        y )Nc                 ó*   — t        j                  «       S r#   rB   rD   s    r   rF   z%_HDecLayer.__init__.<locals>.<lambda>Ü   rG   r    rH   c                 ó0   •— t        j                  ‰| «      S r#   rJ   rL   s    €r   rF   z%_HDecLayer.__init__.<locals>.<lambda>Þ   rM   r    rO   r   z#Kernel size and stride do not alignr   )r   r   rb   r?   rs   r;   r5   r:   r8   r7   r   rP   ÚConvTranspose1drQ   ÚConvTranspose2dÚconv_trrU   rC   rT   rS   )r   r5   r6   rs   r7   r8   r9   r:   r;   r<   r=   r>   r?   rX   rZ   Úklass_trr   s         `         €r   r   z_HDecLayer.__init__Ê   s@  ù€ ô 	‰ÑÔØÐØˆHÙ)ˆØ˜Ò$Û<ˆGÙØ˜fÑ$¨Ñ)¨QÒ.Ü Ð!FÓGÐGØ Ñ'¨AÑ-‰CàˆCØˆŒØˆŒ	ØˆŒ	ØˆŒ	ØˆŒ
ØˆŒØ&ˆÔÜ—	‘	ˆÜ×%Ñ%ˆÙØ&¨Ð*ˆKØ˜a�[ˆFÜ—I‘IˆEÜ×)Ñ)ˆHÙ  e¨[¸&ÓAˆŒÙ˜U“^ˆŒ
Ø�:Š:ÜŸ;™;›=ˆDŒLÜŸ™›ˆD�Já   q¨4¡x°°Q¸±[±À!ÀWÓMˆDŒLÙ   T¡Ó*ˆD�Jr    r$   Úskipc                 óÀ  — | j                   rA|j                  «       dk(  r.|j                  \  }}}|j                  || j                  d|«      }| j
                  s;||z   }t        j                  | j                  | j                  |«      «      d¬«      }n|}|�t        d«      ‚| j                  | j                  |«      «      }| j                   r.| j                  r_|d| j                  | j                   …dd…f   }n=|d| j                  | j                  |z   …f   }|j                  d   |k7  rt        d«      ‚| j                  st        j                  |«      }||fS )	a,  Forward pass for decoding layer.

        Size depends on whether frequency or time

        Args:
            x (torch.Tensor): tensor input of shape `(B, C, F, T)` for frequency and shape
                `(B, C, T)` for time
            skip (torch.Tensor, optional): on first layer, separate frequency and time branches using param
                (default: ``None``)
            length (int): Size of tensor for output

        Returns:
            (Tensor, Tensor):
                Tensor
                    output tensor after decoder layer of shape `(B, C, F * stride, T)` for frequency domain except last
                        frequency layer shape is `(B, C, kernel_size, T)`. Shape is `(B, C, stride * T)`
                        for time domain.
                Tensor
                    contains the output just before final transposed convolution, which is used when the
                        freq. and time branch separate. Otherwise, does not matter. Shape is
                        `(B, C, F, T)` for frequency and `(B, C, T)` for time.
        r^   r]   r   r   Nz%Skip must be none when empty is true..z'Last index of z must be equal to length)r;   r   r_   r`   r5   r:   ra   rf   rS   rT   rb   rU   ry   r?   rs   rc   )	r   r$   r{   Úlengthrg   rh   rj   rl   rm   s	            r   r'   z_HDecLayer.forwardü   s$  € ð. �9Š9˜Ÿ™› AšØ—g‘g‰GˆAˆq�!Ø—‘�q˜$Ÿ)™) R¨Ó+ˆAà�zŠzØ�D‘ˆAÜ—‘�d—j‘j §¡¨a£Ó1°qÔ9‰AàˆAØÐÜ Ð!HÓIÐIà�J‰J�t—|‘| A“Ó'ˆØ�9Š9Ø�xŠxØ�c˜4Ÿ8™8 t§x¡x iÐ/²Ð2Ñ3‘à�#�t—x‘x $§(¡(¨VÑ"3Ð3Ð3Ñ4ˆAØ�w‰w�r‰{˜fÒ$Ü Ð!JÓKÐKØ�yŠyÜ—‘�q“	ˆAà�!ˆtˆr    )
Frn   rN   r   FTrH   r   NTro   r2   s   @r   rr   rr   ·   sÑ   ø„ ñð, ØØØØØØ%ØØ-1Øñ0+àð0+ð ð0+ð ð	0+ð
 ð0+ð ð0+ð ð0+ð ð0+ð ð0+ð ð0+ð ð0+ð ˜4  S ™>Ñ*ð0+ð õ0+ðd.˜Ÿ™ð .¨X°e·l±lÑ-C÷ .r    rr   c            +       ó  ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d$dee   dedededededed	ed
ededededededededededededef*ˆ fd„Z	d„ Z
d%d„Zd&dej                  dedededef
d„Zd „ Zd!„ Zd"ej                  fd#„Zˆ xZS )'ÚHDemucsa#
  Hybrid Demucs model from
    *Hybrid Spectrogram and Waveform Source Separation* :cite:`defossez2021hybrid`.

    See Also:
        * :class:`torchaudio.pipelines.SourceSeparationBundle`: Source separation pipeline with pre-trained models.

    Args:
        sources (List[str]): list of source names. List can contain the following source
            options: [``"bass"``, ``"drums"``, ``"other"``, ``"mixture"``, ``"vocals"``].
        audio_channels (int, optional): input/output audio channels. (Default: 2)
        channels (int, optional): initial number of hidden channels. (Default: 48)
        growth (int, optional): increase the number of hidden channels by this factor at each layer. (Default: 2)
        nfft (int, optional): number of fft bins. Note that changing this requires careful computation of
            various shape parameters and will not work out of the box for hybrid models. (Default: 4096)
        depth (int, optional): number of layers in encoder and decoder (Default: 6)
        freq_emb (float, optional): add frequency embedding after the first frequency layer if > 0,
            the actual value controls the weight of the embedding. (Default: 0.2)
        emb_scale (int, optional): equivalent to scaling the embedding learning rate (Default: 10)
        emb_smooth (bool, optional): initialize the embedding with a smooth one (with respect to frequencies).
            (Default: ``True``)
        kernel_size (int, optional): kernel_size for encoder and decoder layers. (Default: 8)
        time_stride (int, optional): stride for the final time layer, after the merge. (Default: 2)
        stride (int, optional): stride for encoder and decoder layers. (Default: 4)
        context (int, optional): context for 1x1 conv in the decoder. (Default: 4)
        context_enc (int, optional): context for 1x1 conv in the encoder. (Default: 0)
        norm_starts (int, optional): layer at which group norm starts being used.
            decoder layers are numbered in reverse order. (Default: 4)
        norm_groups (int, optional): number of groups for group norm. (Default: 4)
        dconv_depth (int, optional): depth of residual DConv branch. (Default: 2)
        dconv_comp (int, optional): compression of DConv branch. (Default: 4)
        dconv_attn (int, optional): adds attention layers in DConv branch starting at this layer. (Default: 4)
        dconv_lstm (int, optional): adds a LSTM layer in DConv branch starting at this layer. (Default: 4)
        dconv_init (float, optional): initial scale for the DConv branch LayerScale. (Default: 1e-4)
    ÚsourcesÚaudio_channelsÚchannelsÚgrowthÚnfftÚdepthÚfreq_embÚ	emb_scaleÚ
emb_smoothr7   Útime_strider8   r=   Úcontext_encÚnorm_startsr9   Údconv_depthÚ
dconv_compÚ
dconv_attnÚ
dconv_lstmÚ
dconv_initc                 óÊ  •— t         ‰+| �  «        || _        || _        || _        || _        |
| _        || _        || _        || _	        | j                  dz  | _
        d | _        t        j                  «       | _        t        j                  «       | _        t        j                  «       | _        t        j                  «       | _        |}|dz  }|}|}| j                  dz  }t%        | j                  «      D �]ß  }||k\  }||k\  }||k\  rdnd}|dkD  }|} |
}!|s|dk7  rt'        d«      ‚|dz  }!|} d}"d}#|r||
k  r|}!d}"d}#|!| ||"|||||||d	œd
œ}$t)        |$«      }%d|%d<   |
|%d<   ||%d<   d|%d<   t)        |$«      }&|#rt+        ||«      }|}t-        ||fd|i|$¤Ž}'|r>|#du r|dk(  r
d|%d<   d|%d<   t-        ||f||#dœ|%¤Ž}(| j                   j/                  |(«       | j                  j/                  |'«       |dk(  r'| j                  t1        | j
                  «      z  }|dz  }t3        ||f|dk(  |dœ|&¤Ž})|r0t3        ||f|#|dk(  |dœ|%¤Ž}*| j"                  j5                  d|*«       | j                  j5                  d|)«       |}|}t7        ||z  «      }t7        ||z  «      }|r||
k  rd}n||z  }|dk(  s�ŒÁ|s�ŒÅt9        |||	|¬«      | _        || _        �Œâ t=        | «       y )NrN   rO   rH   Únoner   z$When freq is false, freqs must be 1.TF)ÚlstmÚattnr…   ÚcompressÚinit)r7   r8   r;   r?   r<   r9   r>   r   r;   r7   r8   r?   r=   é   )r=   r:   )rs   r=   )r:   rs   r=   )r   r   )r   r   r…   r„   r�   r€   r7   r=   r8   r‚   Ú
hop_lengthr†   r   Ú
ModuleListÚfreq_encoderÚfreq_decoderÚtime_encoderÚtime_decoderÚrangerb   ÚdictÚmaxr4   ÚappendÚlenrr   Úinsertr,   r
   Úfreq_emb_scaleÚ_rescale_module),r   r€   r�   r‚   rƒ   r„   r…   r†   r‡   rˆ   r7   r‰   r8   r=   rŠ   r‹   r9   rŒ   r�   rŽ   r�   r�   r5   Úchin_zr6   Úchout_zÚfreqsÚindexr“   r”   r<   r;   ÚstriÚkerr?   Ú	last_freqÚkwÚkwtÚkw_decÚencÚtencÚdecÚtdecr   s,                                              €r   r   zHDemucs.__init__Q  sI  ø€ ô0 	‰ÑÔØˆŒ
ØˆŒ	Ø,ˆÔØˆŒØ&ˆÔØˆŒØˆŒØ ˆŒàŸ)™) q™.ˆŒØˆŒäŸM™M›OˆÔÜŸM™M›OˆÔäŸM™M›OˆÔÜŸM™M›OˆÔàˆØ˜‘ˆØˆØˆØ—	‘	˜Q‘ˆä˜4Ÿ:™:×&ˆEØ˜JÑ&ˆDØ˜JÑ&ˆDØ(-°Ò(<™À&ˆIØ˜1‘9ˆDØˆDØˆCÙØ˜A’:Ü$Ð%KÓLÐLØ! A‘o�Ø"�àˆCØˆIÙ˜ Ò,Ø�Ø�Ø �	ð  #ØØØØ&Ø*à Ø Ø(Ø *Ø&ññˆBô �r“(ˆCØˆC�‰KØ!,ˆC�ÑØ"ˆC�‰MØˆC�‰JÜ˜"“XˆFáÜ˜e WÓ-�Ø�ä˜V WÑH°kÐHÀRÑHˆCÙØ Ñ$¨°ªØ$%�C˜‘MØ)*�C˜Ñ&Ü! $¨Ð[°{È)Ñ[ÐWZÑ[�Ø×!Ñ!×(Ñ(¨Ô.à×Ñ×$Ñ$ SÔ)Ø˜ŠzØ×*Ñ*¬S°·±Ó->Ñ>�Ø ™�Ü˜W fÐY°5¸A±:ÀwÑYÐRXÑYˆCÙÜ! %¨Ðh°YÀUÈaÁZÐY`ÑhÐdgÑh�Ø×!Ñ!×(Ñ(¨¨DÔ1Ø×Ñ×$Ñ$ Q¨Ô,àˆDØˆFÜ˜ ™Ó'ˆEÜ˜& 7Ñ*Ó+ˆGÙØ˜KÒ'Ø‘Eà˜fÑ$�EØ˜Œz›hÜ 0°¸ÀzÐYbÔ c�”Ø&.�Ö#ðW 'ôZ 	˜Õr    c                 ó¨  — | j                   }| j                  }|}||dz  k7  rt        d«      ‚t        t	        j
                  |j                  d   |z  «      «      }|dz  dz  }| j                  |||||z  z   |j                  d   z
  d¬«      }t        |||«      dd d…d d …f   }|j                  d   |dz   k7  rt        d	«      ‚|ddd|z   …f   }|S )
NrN   zHop length must be nfft // 4r]   rO   r^   Úreflect)Úmode.zESpectrogram's last dimension must be 4 + input size divided by stride)	r˜   r„   rb   r,   ÚmathÚceilr_   Ú_pad1dÚ_spectro)r   r$   Úhlr„   Úx0rk   r?   rm   s           r   Ú_speczHDemucs._specÑ  sâ   € Ø�_‰_ˆØ�y‰yˆØˆð �˜‘Š?ÜÐ;Ó<Ð<Ü”—‘˜1Ÿ7™7 2™;¨Ñ+Ó,Ó-ˆØ�A‰g˜‰kˆØ�K‰K˜˜3  b¨2¡g¡°·±¸±Ñ ;À)ˆKÓLˆä�Q˜˜bÓ! # s¨ sªA +Ñ.ˆØ�7‰7�2‰;˜"˜q™&Ò ÜÐdÓeÐeØˆc�1�q˜2‘v�:ˆoÑˆØˆr    c                 ó  — | j                   }t        j                  |g d¢«      }t        j                  |ddg«      }|dz  dz  }|t        t	        j
                  ||z  «      «      z  d|z  z   }t        |||¬«      }|d|||z   …f   }|S )N)r   r   r   r   rO   r^   )r}   .)r˜   ra   r?   r,   r·   r¸   Ú	_ispectro)r   rm   r}   r»   r?   rk   r$   s          r   Ú_ispeczHDemucs._ispecé  sŒ   € Ø�_‰_ˆÜ�E‰E�!’\Ó"ˆÜ�E‰E�!�a˜�VÓˆØ�A‰g˜‰kˆØ”#”d—i‘i ¨¡Ó,Ó-Ñ-°°C±Ñ7ˆÜ�a˜ BÔ'ˆØˆc�3˜˜v™Ð%Ð%Ñ&ˆØˆr    r$   Úpadding_leftÚpadding_rightr¶   Úvaluec                 ó¼   — |j                   d   }|dk(  r/t        ||«      }||k  rt        j                  |d||z
  dz   f«      }t        j                  |||f||«      S )z¤Wrapper around F.pad, in order for reflect padding when num_frames is shorter than max_pad.
        Add extra zero padding around in order for padding to not break.r]   rµ   r   r   )r_   r    ra   r?   )r   r$   rÁ   rÂ   r¶   rÃ   r}   Úmax_pads           r   r¹   zHDemucs._pad1dó  sf   € ð —‘˜‘ˆØ�9ÒÜ˜,¨Ó6ˆGØ˜Ò Ü—E‘E˜!˜a ¨6Ñ!1°AÑ!5Ð6Ó7�Ü�u‰u�Q˜ }Ð5°t¸UÓCÐCr    c                 ó¦   — |j                   \  }}}}t        j                  |«      j                  ddddd«      }|j	                  ||dz  ||«      }|S )Nr   r   rN   rO   r^   )r_   r   Úview_as_realrd   re   )r   rm   rg   rh   ri   rj   Úms          r   Ú
_magnitudezHDemucs._magnitudeý  sS   € à—g‘g‰ˆˆ1ˆb�!Ü×Ñ˜qÓ!×)Ñ)¨!¨Q°°1°aÓ8ˆØ�I‰I�a˜˜Q™  AÓ&ˆØˆr    c                 óÄ   — |j                   \  }}}}}|j                  ||dd||«      j                  dddddd«      }t        j                  |j                  «       «      }|S )Nr]   rO   r   r   rN   é   r^   )r_   r`   rd   r   Úview_as_complexÚ
contiguous)r   rÈ   rg   ÚSrh   ri   rj   r&   s           r   Ú_maskzHDemucs._mask  s^   € àŸ™‰ˆˆ1ˆa��QØ�f‰f�Q˜˜2˜q " aÓ(×0Ñ0°°A°q¸!¸QÀÓBˆÜ×#Ñ# C§N¡NÓ$4Ó5ˆØˆ
r    Úinputc                 óÜ  — |j                   dk7  rt        d|j                  › �«      ‚|j                  d   | j                  k7  rt        d|j                  d   › d�«      ‚|}|j                  d   }| j	                  |«      }| j                  |«      }|}|j                  \  }}}}	|j                  dd¬	«      }
|j                  dd¬	«      }||
z
  d
|z   z  }|}|j                  dd¬	«      }|j                  dd¬	«      }||z
  d
|z   z  }g }g }g }g }t        | j                  «      D �]7  \  }}|j                  |j                  d   «       d}|t        | j                  «      k  rU|j                  |j                  d   «       | j                  |   } ||«      }|j                  s|j                  |«       n|} |||«      }|dk(  r…| j                  �yt        j                   |j                  d   |j"                  ¬«      }| j                  |«      j%                  «       ddd…dd…df   j'                  |«      }|| j(                  |z  z   }|j                  |«       �Œ: t        j*                  |«      }t        j*                  |«      }t        | j,                  «      D ]ë  \  }}|j/                  d«      } ||||j/                  d«      «      \  }}| j0                  t        | j2                  «      z
  }||k\  sŒ[| j2                  ||z
     }|j/                  d«      }|j                  rD|j                  d   dk7  rt        d|j                  › �«      ‚|dd…dd…df   } ||d|«      \  }}ŒÎ|j/                  d«      } ||||«      \  }}Œí t        |«      dk7  rt5        d«      ‚t        |«      dk7  rt5        d«      ‚t        |«      dk7  rt5        d«      ‚t        | j6                  «      } |j9                  || d||	«      }||dd…df   z  |
dd…df   z   }| j;                  |«      }!| j=                  |!|«      }|j9                  || d|«      }||dd…df   z  |dd…df   z   }||z   }|S )a  HDemucs forward call

        Args:
            input (torch.Tensor): input mixed tensor of shape `(batch_size, channel, num_frames)`

        Returns:
            Tensor
                output tensor split into sources of shape `(batch_size, num_sources, channel, num_frames)`
        r^   zDExpected 3D tensor with dimensions (batch, channel, frames). Found: r   zZThe channel dimension of input Tensor must match `audio_channels` of HDemucs model. Found:Ú.r]   )r   rO   r^   T)r   Úkeepdimgñhãˆµøä>)r   rO   Nr   éþÿÿÿ)ÚdevicerO   z0If tdec empty is True, pre shape does not match zsaved is not emptyzlengths_t is not emptyzsaved_t is not empty)Úndimrb   r_   r�   r½   rÉ   ÚmeanÚstdÚ	enumeraterš   r¡   r¢   rœ   r:   r†   r   r   rÕ   ÚtÚ	expand_asr¤   Ú
zeros_liker›   Úpopr…   r�   ÚAssertionErrorr€   r`   rÏ   rÀ   )"r   rÐ   r$   r}   rm   Úmagrg   rh   ÚFqrj   r×   rØ   ÚxtÚmeantÚstdtÚsavedÚsaved_tÚlengthsÚ	lengths_tÚidxÚencoder[   r±   ÚfrsÚembÚdecoder{   ÚpreÚoffsetr³   Úlength_tÚ_rÎ   Úzouts"                                     r   r'   zHDemucs.forward  sR  € ð �:‰:˜Š?ÜÐcÐdi×doÑdoÐcpÐqÓrÐrà�;‰;�q‰>˜T×0Ñ0Ò0ÜðØŸ™ Q™Ð(¨ð+óð ð
 ˆØ—‘˜‘ˆà�J‰J�uÓˆØ�o‰o˜aÓ ˆØˆà—g‘g‰ˆˆ1ˆb�!ð �v‰v˜)¨TˆvÓ2ˆØ�e‰e˜	¨4ˆeÓ0ˆØ�‰X˜$ ™*Ñ%ˆð ˆØ—‘˜F¨D�Ó1ˆØ�v‰v˜&¨$ˆvÓ/ˆØ�5‰j˜T D™[Ñ)ˆàˆØˆØˆØ!ˆ	ä$ T×%6Ñ%6×7‰KˆC�Ø�N‰N˜1Ÿ7™7 2™;Ô'ØˆFØ”S˜×*Ñ*Ó+Ò+à× Ñ  §¡¨"¡Ô.Ø×(Ñ(¨Ñ-�Ù˜"“X�Ø—z’zà—N‘N 2Õ&ð  �FÙ�q˜&Ó!ˆAØ�aŠx˜DŸM™MÐ5ô —l‘l 1§7¡7¨2¡;°q·x±xÔ@�Ø—m‘m CÓ(×*Ñ*Ó,¨T²1²a¸Ð-=Ñ>×HÑHÈÓK�Ø˜×+Ñ+¨cÑ1Ñ1�à�L‰L˜ŽOð/ 8ô2 ×Ñ˜QÓˆÜ×Ñ˜aÓ ˆô % T×%6Ñ%6Ö7‰KˆC�Ø—9‘9˜R“=ˆDÙ˜A˜t W§[¡[°£_Ó5‰FˆAˆsð —Z‘Z¤# d×&7Ñ&7Ó"8Ñ8ˆFØ�f‹}Ø×(Ñ(¨¨v©Ñ6�Ø$Ÿ=™=¨Ó,�Ø—:’:Ø—y‘y ‘| qÒ(Ü(Ð+[Ð\_×\eÑ\eÐ[fÐ)gÓhÐhØša¢ A˜g™,�CÙ   d¨HÓ5‘E�B™à"Ÿ;™; r›?�DÙ   T¨8Ó4‘E�B™ð! 8ô$ ˆu‹:˜Š?Ü Ð!5Ó6Ð6Üˆy‹>˜QÒÜ Ð!9Ó:Ð:Üˆw‹<˜1ÒÜ Ð!7Ó8Ð8ä�—‘ÓˆØ�F‰F�1�a˜˜R Ó#ˆØ�’A�t�G‘Ñ˜t¢A t G™}Ñ,ˆà�z‰z˜!‹}ˆØ�K‰K˜˜fÓ%ˆà�W‰W�Q˜˜2˜vÓ&ˆØ�$’q˜$�w‘-Ñ %ª¨4¨¡.Ñ0ˆØ�‰FˆØˆr    )rO   é0   rO   é   é   gš™™™™™É?é
   Trn   rO   rN   r   r   rN   rN   rO   rN   rN   rN   ç-Cëâ6?r#   )Úzerog        )r(   r)   r*   r+   r   rp   r,   r-   r.   r   r½   rÀ   r   r0   r¹   rÉ   rÏ   r'   r1   r2   s   @r   r   r   -  s‘  ø„ ñ!ðL  ØØØØØØØØØØØØØØØØØØØ ñ-~à�c‘ð~ð ð~ð ð	~ð
 ð~ð ð~ð ð~ð ð~ð ð~ð ð~ð ð~ð ð~ð ð~ð ð~ð ð~ð  ð!~ð" ð#~ð$ ð%~ð& ð'~ð( ð)~ð* ð+~ð, õ-~ò@ó0ñD˜Ÿ™ð D°Cð DÈð DÐSVð DÐhmó Dòòðo˜UŸ\™\÷ or    r   c                   óf   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 ddededededededed	ed
edefˆ fd„Zd„ Z	ˆ xZ
S )rW   a±  
    New residual branches in each encoder layer.
    This alternates dilated convolutions, potentially with LSTMs and attention.
    Also before entering each residual branch, dimension is projected on a smaller subspace,
    e.g. of dim `channels // compress`.

    Args:
        channels (int): input/output channels for residual branch.
        compress (float, optional): amount of channel compression inside the branch. (default: 4)
        depth (int, optional): number of layers in the residual branch. Each layer has its own
            projection, and potentially LSTM and attention.(default: 2)
        init (float, optional): initial scale for LayerNorm. (default: 1e-4)
        norm_type (bool, optional): Norm type, either ``group_norm `` or ``none`` (Default: ``group_norm``)
        attn (bool, optional): use LocalAttention. (Default: ``False``)
        heads (int, optional): number of heads for the LocalAttention.  (default: 4)
        ndecay (int, optional): number of decay controls in the LocalAttention. (default: 4)
        lstm (bool, optional): use LSTM. (Default: ``False``)
        kernel_size (int, optional): kernel size for the (dilated) convolutions. (default: 3)
    r‚   r•   r…   r–   r<   r”   ÚheadsÚndecayr“   r7   c                 ó&  •— t         ‰| �  «        |
dz  dk(  rt        d«      ‚|| _        || _        t        |«      | _        |dkD  }d„ }|dk(  rd„ }t        ||z  «      }t        j                  }t        j                  g «      | _        t        | j                  «      D ]ñ  }|rt        d|«      nd}||
dz  z  }t        j                  |||
||¬«       ||«       |«       t        j                  |d|z  d«       |d|z  «      t        j                  d«      t!        ||«      g}|r|j#                  d	t%        |||¬
«      «       |	r|j#                  d	t'        |dd¬«      «       t        j(                  |Ž }| j                  j+                  |«       Œó y )NrO   r   z(Kernel size should not be divisible by 2c                 ó*   — t        j                  «       S r#   rB   rD   s    r   rF   z!_DConv.__init__.<locals>.<lambda>©  rG   r    rH   c                 ó.   — t        j                  d| «      S )Nr   rJ   rD   s    r   rF   z!_DConv.__init__.<locals>.<lambda>«  s   € ¤§¡¨Q°Ô 2r    r   )ÚdilationÚpaddingr^   )rù   rú   T)Úlayersr{   )r   r   rb   r‚   r•   Úabsr…   r,   r   ÚGELUr™   r   rž   ÚpowrP   ÚGLUÚ_LayerScaler£   Ú_LocalStateÚ_BLSTMÚ
Sequentialr¡   )r   r‚   r•   r…   r–   r<   r”   rù   rú   r“   r7   ÚdilaterX   ÚhiddenÚactrE   rþ   rÿ   ÚmodsÚlayerr   s                       €r   r   z_DConv.__init__’  sj  ø€ ô 	‰ÑÔØ˜‰?˜aÒÜÐGÓHÐHØ ˆŒØ ˆŒÜ˜“ZˆŒ
Ø˜‘ˆñ *ˆØ˜Ò$Ù2ˆGä�X Ñ(Ó)ˆä�g‰gˆä—m‘m BÓ'ˆŒÜ�t—z‘zÖ"ˆAÙ$*”s˜1˜a”y°ˆHØ +°Ñ"2Ñ3ˆGä—	‘	˜( F¨KÀ(ÐT[Ô\Ù˜“Ù“Ü—	‘	˜& ! h¡,°Ó2Ù˜˜H™Ó%Ü—‘�q“	Ü˜H dÓ+ðˆDñ Ø—‘˜Aœ{¨6¸ÀvÔNÔOÙØ—‘˜Aœv f°Q¸TÔBÔCÜ—M‘M 4Ð(ˆEØ�K‰K×Ñ˜uÕ%ñ# #r    c                 ó>   — | j                   D ]  }| ||«      z   }Œ |S )zÁDConv forward call

        Args:
            x (torch.Tensor): input tensor for convolution

        Returns:
            Tensor
                Output after being run through layers.
        )r   )r   r$   r  s      r   r'   z_DConv.forwardÅ  s$   € ð —[”[ˆEØ‘E˜!“H‘‰Að !àˆr    )	rN   rO   rö   rH   FrN   rN   Fr^   )r(   r)   r*   r+   r,   r-   rp   r.   r   r'   r1   r2   s   @r   rW   rW   }  s’   ø„ ñð. ØØØ%ØØØØØñ1&àð1&ð ð1&ð ð	1&ð
 ð1&ð ð1&ð ð1&ð ð1&ð ð1&ð ð1&ð õ1&öfr    rW   c                   óf   ‡ — e Zd ZdZddedefˆ fd„Zdej                  dej                  fd„Z	ˆ xZ
S )	r  ae  
    BiLSTM with same hidden units as input dim.
    If `max_steps` is not None, input will be splitting in overlapping
    chunks and the LSTM applied separately on each chunk.
    Args:
        dim (int): dimensions at LSTM layer.
        layers (int, optional): number of LSTM layers. (default: 1)
        skip (bool, optional): (default: ``False``)
    r   r{   c                 ó¶   •— t         ‰| �  «        d| _        t        j                  d|||¬«      | _        t        j                  d|z  |«      | _        || _        y )NéÈ   T)ÚbidirectionalÚ
num_layersÚhidden_sizeÚ
input_sizerO   )	r   r   Ú	max_stepsr   ÚLSTMr“   ÚLinearÚlinearr{   )r   r   r   r{   r   s       €r   r   z_BLSTM.__init__ß  sI   ø€ Ü‰ÑÔØˆŒÜ—G‘G¨$¸6ÈsÐ_bÔcˆŒ	Ü—i‘i  C¡¨Ó-ˆŒØˆ�	r    r$   r!   c           	      óB  — |j                   \  }}}|}d}d}d}d}	| j                  �c|| j                  kD  rT| j                  }|dz  }t        |||«      }
|
j                   d   }	d}|
j                  dddd«      j	                  d||«      }|j                  ddd«      }| j                  |«      d   }| j                  |«      }|j                  ddd«      }|r·g }|j	                  |d||«      }
|dz  }t        |	«      D ]m  }|dk(  r |j                  |
dd…|dd…d| …f   «       Œ(||	dz
  k(  r|j                  |
dd…|dd…|d…f   «       ŒO|j                  |
dd…|dd…|| …f   «       Œo t        j                  |d«      }|d	d|…f   }|}| j                  r||z   }|S )
a  BLSTM forward call

        Args:
            x (torch.Tensor): input tensor for BLSTM shape is `(batch_size, dim, time_steps)`

        Returns:
            Tensor
                Output after being run through bidirectional LSTM. Shape is `(batch_size, dim, time_steps)`
        Fr   NrO   Tr   r^   r]   .)r_   r  Ú_unfoldrd   re   r“   r  rž   r¡   r   Úcatr{   )r   r$   rg   rh   rj   rl   ÚframedÚwidthr8   ÚnframesÚframesr&   ÚlimitÚks                 r   r'   z_BLSTM.forwardæ  s¶  € ð —'‘'‰ˆˆ1ˆaØˆØˆØˆØˆØˆØ�>‰>Ð%¨!¨d¯n©nÒ*<Ø—N‘NˆEØ˜a‘ZˆFÜ˜Q  vÓ.ˆFØ—l‘l 1‘oˆGØˆFØ—‘˜q ! Q¨Ó*×2Ñ2°2°q¸%Ó@ˆAà�I‰I�a˜˜AÓˆà�I‰I�a‹L˜‰OˆØ�K‰K˜‹NˆØ�I‰I�a˜˜AÓˆÙØˆCØ—Y‘Y˜q " a¨Ó/ˆFØ˜a‘KˆEÜ˜7–^�Ø˜’6Ø—J‘J˜v¢a¨ªA¨w°°¨wÐ&6Ñ7Õ8Ø˜' A™+Ò%Ø—J‘J˜v¢a¨ªA¨u©v oÑ6Õ7à—J‘J˜v¢a¨ªA¨u°e°V¨|Ð&;Ñ<Õ=ð $ô —)‘)˜C Ó$ˆCØ�c˜2˜A˜2�g‘,ˆCØˆAØ�9Š9Ø�A‘ˆAàˆr    )r   F)r(   r)   r*   r+   r,   r.   r   r   r0   r'   r1   r2   s   @r   r  r  Ô  s6   ø„ ññ Cð °4õ ð.˜Ÿ™ð .¨%¯,©,÷ .r    r  c                   ój   ‡ — e Zd ZdZd	dededefˆ fd„Zdej                  dej                  fd„Zˆ xZ	S )
r  a   Local state allows to have attention based only on data (no positional embedding),
    but while setting a constraint on the time window (e.g. decaying penalty term).
    Also a failed experiments with trying to provide some frequency based attention.
    r‚   rù   rú   c                 ót  •— t         t        | �  «        ||z  dk7  rt        d«      ‚|| _        || _        t        j                  ||d«      | _        t        j                  ||d«      | _	        t        j                  ||d«      | _
        t        j                  |||z  d«      | _        |rm| j                  j                  xj                  dz  c_        | j                  j                  €t        d«      ‚d| j                  j                  j                  dd t        j                  ||dz  z   |d«      | _        y)z¬
        Args:
            channels (int): Size of Conv1d layers.
            heads (int, optional):  (default: 4)
            ndecay (int, optional): (default: 4)
        r   z$Channels must be divisible by heads.r   g{®Gáz„?Nzbias must not be None.rÔ   )r   r  r   rb   rù   rú   r   rP   ÚcontentÚqueryÚkeyÚquery_decayr   r   ÚbiasÚproj)r   r‚   rù   rú   r   s       €r   r   z_LocalState.__init__  s  ø€ ô 	Œk˜4Ñ)Ô+Ø�eÑ˜qÒ ÜÐCÓDÐDØˆŒ
ØˆŒÜ—y‘y ¨8°QÓ7ˆŒÜ—Y‘Y˜x¨°1Ó5ˆŒ
Ü—9‘9˜X x°Ó3ˆŒäŸ9™9 X¨u°v©~¸qÓAˆÔÙà×Ñ×#Ñ#×(Ò(¨DÑ0Õ(Ø×Ñ×$Ñ$Ð,Ü Ð!9Ó:Ð:Ø,.ˆD×Ñ×!Ñ!×&Ñ&¡qÐ)Ü—I‘I˜h¨°©Ñ2°H¸aÓ@ˆ�	r    r$   r!   c                 óì  — |j                   \  }}}| j                  }t        j                  ||j                  |j
                  ¬«      }|dd…df   |ddd…f   z
  }| j                  |«      j                  ||d|«      }| j                  |«      j                  ||d|«      }	t        j                  d|	|«      }
|
t        j                  |	j                   d   «      z  }
| j                  rÔt        j                  d| j                  dz   |j                  |j
                  ¬«      }| j                  |«      j                  ||d|«      }t        j                  |«      dz  }|j                  ddd«       |j                  «       z  t        j                  | j                  «      z  }|
t        j                  d||«      z  }
|
j!                  t        j"                  ||
j                  t        j$                  ¬«      d«       t        j&                  |
d¬	«      }| j)                  |«      j                  ||d|«      }t        j                  d
||«      }|j+                  |d|«      }|| j-                  |«      z   S )zÏLocalState forward call

        Args:
            x (torch.Tensor): input tensor for LocalState

        Returns:
            Tensor
                Output after being run through LocalState layer.
        )rÕ   ÚdtypeNr]   zbhct,bhcs->bhtsrO   r   zfts,bhfs->bhtsiœÿÿÿr   zbhts,bhct->bhcs)r_   rù   r   r   rÕ   r,  r&  r`   r'  Úeinsumr·   r   rú   r(  Úsigmoidr  Úmasked_fill_Úeyer.   Úsoftmaxr%  re   r*  )r   r$   rg   rh   rj   rù   ÚindexesÚdeltaÚqueriesÚkeysÚdotsÚdecaysÚdecay_qÚdecay_kernelÚweightsr%  Úresults                    r   r'   z_LocalState.forward6  só  € ð —'‘'‰ˆˆ1ˆaØ—
‘
ˆÜ—,‘,˜q¨¯©¸¿¹ÔAˆàš˜4˜Ñ  7¨4²¨7Ñ#3Ñ3ˆà—*‘*˜Q“-×$Ñ$ Q¨¨r°1Ó5ˆØ�x‰x˜‹{×Ñ  5¨"¨aÓ0ˆä�|‰|Ð-¨t°WÓ=ˆØ”—	‘	˜$Ÿ*™* Q™-Ó(Ñ(ˆØ�;Š;Ü—\‘\ ! T§[¡[°1¡_¸Q¿X¹XÈQÏWÉWÔUˆFØ×&Ñ& qÓ)×.Ñ.¨q°%¸¸QÓ?ˆGÜ—m‘m GÓ,¨qÑ0ˆGØ"ŸK™K¨¨A¨qÓ1Ð1°E·I±I³KÑ?Ä$Ç)Á)ÈDÏKÉKÓBXÑXˆLØ”E—L‘LÐ!1°<ÀÓIÑIˆDð 	×Ñœ%Ÿ)™) A¨d¯k©kÄÇÁÔLÈdÔSÜ—-‘- ¨!Ô,ˆà—,‘,˜q“/×&Ñ& q¨%°°QÓ7ˆÜ—‘Ð/°¸'ÓBˆØ—‘  2 qÓ)ˆØ�4—9‘9˜VÓ$Ñ$Ð$r    )rN   rN   )
r(   r)   r*   r+   r,   r   r   r0   r'   r1   r2   s   @r   r  r    sA   ø„ ññ
A ð A¨Sð A¸cõ Að2#%˜Ÿ™ð #%¨%¯,©,÷ #%r    r  c                   óf   ‡ — e Zd ZdZddedefˆ fd„Zdej                  dej                  fd„Z	ˆ xZ
S )	r  z£Layer scale from [Touvron et al 2021] (https://arxiv.org/pdf/2103.17239.pdf).
    This rescales diagonally residual outputs close to 0 initially, then learnt.
    r‚   r–   c                 ó²   •— t         ‰| �  «        t        j                  t	        j
                  |d¬«      «      | _        || j                  j                  dd y)z‹
        Args:
            channels (int): Size of  rescaling
            init (float, optional): Scale to default to (default: 0)
        T)Úrequires_gradN)r   r   r   Ú	Parameterr   Úzerosr   r   )r   r‚   r–   r   s      €r   r   z_LayerScale.__init__a  s=   ø€ ô 	‰ÑÔÜ—\‘\¤%§+¡+¨hÀdÔ"KÓLˆŒ
Ø!ˆ�
‰
�‰™Ñr    r$   r!   c                 ó.   — | j                   dd…df   |z  S )z½LayerScale forward call

        Args:
            x (torch.Tensor): input tensor for LayerScale

        Returns:
            Tensor
                Output after rescaling tensor.
        N)r   )r   r$   s     r   r'   z_LayerScale.forwardk  s   € ð �z‰zš!˜T˜'Ñ" QÑ&Ð&r    )r   )r(   r)   r*   r+   r,   r-   r   r   r0   r'   r1   r2   s   @r   r  r  \  s6   ø„ ññ" ð "¨Eõ "ð
'˜Ÿ™ð 
'¨%¯,©,÷ 
'r    r  Úar7   r8   r!   c                 óö  — t        | j                  dd «      }t        | j                  d   «      }t        j                  ||z  «      }|dz
  |z  |z   }t        j                  | d||z
  g¬«      } t        | j                  «       «      D �cg c]  }| j                  |«      ‘Œ }}|d   dk7  rt        d«      ‚|dd |dgz   }|j                  |«       |j                  |«       | j                  ||«      S c c}w )zûGiven input of size [*OT, T], output Tensor of size [*OT, F, K]
    with K the kernel size, by extracting frames with the given stride.
    This will pad the input so that `F = ceil(T / K)`.
    see https://github.com/pytorch/pytorch/issues/60466
    Nr]   r   r   )rÐ   r?   zData should be contiguous.)Úlistr_   r,   r·   r¸   ra   r?   rž   r   r8   rb   r¡   Ú
as_strided)	rB  r7   r8   r_   r}   Ún_framesÚ
tgt_lengthr   Ústridess	            r   r  r  x  sé   € ô �—‘˜˜"�Ó€EÜ�—‘˜‘Ó€FÜ�y‰y˜ &™Ó)€HØ˜Q‘, &Ñ(¨;Ñ6€JÜ	�‰�A˜A˜z¨FÑ2Ð3Ô4€AÜ(-¨a¯e©e«g¬Ó7© ˆq�x‰x˜�}¨€GÐ7Øˆr�{�aÒÜÐ5Ó6Ð6Ø�c�rˆl˜f a˜[Ñ(€GØ	‡L�L�ÔØ	‡L�L�ÔØ�<‰<˜˜wÓ'Ð'ùò 8s   ÂC6c                 ó¶  — | j                  «       D ]Æ  }t        |t        j                  t        j                  t        j
                  t        j                  f«      sŒL|j                  j                  «       j                  «       }|dz  dz  }|j                  xj                  |z  c_
        |j                  €Œ¨|j                  xj                  |z  c_
        ŒÈ y)zI
    Rescales initial weight scale for all models within the module.
    gš™™™™™¹?g      à?N)ÚmodulesÚ
isinstancer   rP   rw   rQ   rx   r   rØ   Údetachr   r)  )ÚmoduleÚsubrØ   r   s       r   r¥   r¥   Œ  s‘   € ð �~‰~ÖˆÜ�cœBŸI™I¤r×'9Ñ'9¼2¿9¹9Äb×FXÑFXÐYÕZØ—*‘*—.‘.Ó"×)Ñ)Ó+ˆCØ˜3‘Y 3Ñ&ˆEØ�J‰J�OŠO˜uÑ$�OØ�x‰xÑ#Ø—‘—’ Ñ&–ñ  r    r$   Ún_fftr˜   r?   c                 óz  — t        | j                  d d «      }t        | j                  d   «      }| j                  d|«      } t	        j
                  | |d|z   z  |t	        j                  |«      j                  | «      |dddd¬«	      }|j                  \  }}}	|j                  ||	g«       |j                  |«      S )Nr]   r   Trµ   )ÚwindowÚ
win_lengthÚ
normalizedÚcenterÚreturn_complexÚpad_mode)
rD  r_   r,   re   r   ÚstftÚhann_windowÚtoÚextendr`   )
r$   rO  r˜   r?   Úotherr}   rm   rð   r¨   Úframes
             r   rº   rº   ™  s¯   € Ü�—‘˜˜"�Ó€EÜ�—‘˜‘Ó€FØ	�	‰	�"�fÓ€AÜ�
‰
Ø	Ø��S‘ÑØÜ× Ñ  Ó'×*Ñ*¨1Ó-ØØØØØô
	€Að —g‘g�O€A€uˆeØ	‡L�L�%˜�Ô Ø�6‰6�%‹=Ðr    rm   r}   c           
      óÌ  — t        | j                  d d «      }t        | j                  d   «      }t        | j                  d   «      }d|z  dz
  }| j                  d||«      } |d|z   z  }t	        j
                  | ||t	        j                  |«      j                  | j                  «      |d|d¬«      }	|	j                  \  }
}|j                  |«       |	j                  |«      S )NrÔ   r]   rO   r   T)rQ  rR  rS  r}   rT  )
rD  r_   r,   r`   r   ÚistftrX  rY  Úrealr¡   )rm   r˜   r}   r?   r[  r¨   r   rO  rR  r$   rð   s              r   r¿   r¿   ­  sÐ   € Ü�—‘˜˜"�Ó€EÜ�—‘˜‘Ó€EÜ�—‘˜‘Ó€Fà�‰I˜‰M€EØ	�‰ˆr�5˜&Ó!€AØ˜1˜s™7Ñ#€JÜ�‰Ø	ØØÜ× Ñ  Ó,×/Ñ/°·±Ó7ØØØØô		€Að —‘�I€A€vØ	‡L�L�ÔØ�6‰6�%‹=Ðr    r€   c                 ó   — t        | dd¬«      S )zÚBuilds low nfft (1024) version of :class:`HDemucs`, suitable for sample rates around 8 kHz.

    Args:
        sources (List[str]): See :py:func:`HDemucs`.

    Returns:
        HDemucs:
            HDemucs model.
    i   rË   ©r€   r„   r…   ©r   ©r€   s    r   Úhdemucs_lowrd  Ä  ó   € ô ˜7¨°QÔ7Ð7r    c                 ó   — t        | dd¬«      S )aÉ  Builds medium nfft (2048) version of :class:`HDemucs`, suitable for sample rates of 16-32 kHz.

    .. note::

        Medium HDemucs has not been tested against the original Hybrid Demucs as this nfft and depth configuration is
        not compatible with the original implementation in https://github.com/facebookresearch/demucs

    Args:
        sources (List[str]): See :py:func:`HDemucs`.

    Returns:
        HDemucs:
            HDemucs model.
    r—   rô   ra  rb  rc  s    r   Úhdemucs_mediumrg  Ò  s   € ô  ˜7¨°QÔ7Ð7r    c                 ó   — t        | dd¬«      S )zßBuilds medium nfft (4096) version of :class:`HDemucs`, suitable for sample rates of 44.1-48 kHz.

    Args:
        sources (List[str]): See :py:func:`HDemucs`.

    Returns:
        HDemucs:
            HDemucs model.
    ró   rô   ra  rb  rc  s    r   Úhdemucs_highri  å  re  r    )i   r   r   )r   r   r   )r·   ÚtypingÚtpr   r   r   r   r   r   Útorch.nnr   ra   ÚModuler
   r4   rr   r   rW   r  r  r  r0   r,   r  r¥   rº   r¿   rp   rd  rg  ri  © r    r   Ú<module>ro     s   ðó4 Û ß ,Ó ,ã Ý Ý $ô#�u—x‘x—‘ô #ôLk�—‘—‘ô kô\s�—‘—‘ô sôlMˆe�h‰h�o‰oô Mô`
TˆU�X‰X�_‰_ô Tôn@ˆU�X‰X�_‰_ô @ôFB%�"—)‘)ô B%ôJ'�"—)‘)ô 'ð8(ˆu�|‰|ð (¨#ð (°sð (¸u¿|¹|ó (ò(
'ñ�—‘ð  Sð ¸Cð È#ð ÐV[×VbÑVbó ñ(�—‘ð ¨3ð ¸Cð È#ð ÐV[×VbÑVbó ð.8˜˜c™ð 8 wó 8ð8˜D ™Ið 8¨'ó 8ð&8˜$˜s™)ð 8¨ô 8r    