
    (ti                         d dl Z d dlZd dlmZ d dlmc mZ ddlmZm	Z	 ddl
mZ  G d dej                        Z G d dee      Zy)	    N   )ConfigMixinregister_to_config)
ModelMixinc                        e Zd Z	 	 	 	 	 ddedededeedf   dedef fdZd	ej                  d
ej                  fdZ
 xZS )ResBlockchannelskernel_sizestride	dilations.leaky_relu_negative_slopepadding_modec                 z   t         	|           || _        || _        t	        j
                  |D cg c]  }t	        j                  ||||||       c}      | _        t	        j
                  t        t        |            D cg c]  }t	        j                  ||||d|       c}      | _
        y c c}w c c}w )N)r   dilationpadding   )super__init__r   negative_slopenn
ModuleListConv1dconvs1rangelenconvs2)
selfr	   r
   r   r   r   r   r   _	__class__s
            o/Volumes/fast/ai/experiments/MLX_z-image/.venv/lib/python3.12/site-packages/diffusers/pipelines/ltx2/vocoder.pyr   zResBlock.__init__   s     	"7mm !* )H 		(Hk&S[eqr )
 mm s9~..A 		(Hk&ST^jk.
s   "B3"B8xreturnc                    t        | j                  | j                        D ]\  \  }}t        j                  || j
                        } ||      }t        j                  || j
                        } ||      }||z   }^ |S )Nr   )zipr   r   F
leaky_relur   )r   r!   conv1conv2xts        r    forwardzResBlock.forward'   sm    T[[9LE5a0C0CDBrBb1D1DEBrBBA :     )r   r   r   r      皙?same)__name__
__module____qualname__inttuplefloatstrr   torchTensorr+   __classcell__r   s   @r    r   r      sv     %.+."

 
 	

 c?
 $)
 
6 %,, r,   r   c                        e Zd ZdZedddg dg dg dg dg dg dgd	d
f	dedededee   dee   dee   deee      dedef fd       Zdde	j                  dede	j                  fdZ xZS )LTX2Vocoderz\
    LTX 2.0 vocoder for converting generated mel spectrograms back to audio waveforms.
       i      )            rC   )   r.   r?   r?   r?   )r         r-   r/   i]  in_channelshidden_channelsout_channelsupsample_kernel_sizesupsample_factorsresnet_kernel_sizesresnet_dilationsr   output_sampling_ratec
                    t         |           t        |      | _        t        |      | _        || _        t        j                  |      | _        || _	        | j                  t        |      k7  r%t        d| j                   dt        |       d      | j                  t        |      k7  r.t        dt        | j                         dt        |       d      t        j                  ||ddd      | _        t        j                         | _        t        j                         | _        |}
t#        t%        ||            D ]  \  }\  }}|
d	z  }| j                  j'                  t        j(                  |
|||||z
  d	z  
             t%        ||      D ]-  \  }}| j                   j'                  t+        ||||             / |}
 t        j                  |ddd
      | _        y )Nza`upsample_kernel_sizes` and `upsample_factors` should be lists of the same length but are length z and z, respectively.z_`resnet_kernel_sizes` and `resnet_dilations` should be lists of the same length but are length rE   r   r   )r
   r   r   r?   )r   r   )r   r   )r   r   r   num_upsample_layersresnets_per_upsamplerI   mathprodtotal_upsample_factorr   
ValueErrorr   r   conv_inr   
upsamplersresnets	enumerater%   appendConvTranspose1dr   conv_out)r   rG   rH   rI   rJ   rK   rL   rM   r   rN   input_channelsir   r
   output_channelsr   r   s                   r    r   zLTX2Vocoder.__init__6   s    	#&'<#= $'(;$<!(%)YY/?%@"7##s+;'<<,,-U37G3H2IZ 
 $$,<(==11235=M9N8O` 
 yyo1UV`ab--/}}((1#6FH]2^(_$A$,1OOO"""""#!(61a7 +..ACS*T&Y##'#"+2K	 +U -N+ )`. 		/<1VWXr,   hidden_states	time_lastr"   c           	         |s|j                  dd      }|j                  dd      }| j                  |      }t        | j                        D ]  }t        j                  || j                        } | j                  |   |      }|| j                  z  }|dz   | j                  z  }t        j                  t        ||      D cg c]  } | j                  |   |       c}d      }t        j                  |d      } t        j                  |d      }| j                  |      }t        j                  |      }|S c c}w )a  
        Forward pass of the vocoder.

        Args:
            hidden_states (`torch.Tensor`):
                Input Mel spectrogram tensor of shape `(batch_size, num_channels, time, num_mel_bins)` if `time_last`
                is `False` (the default) or shape `(batch_size, num_channels, num_mel_bins, time)` if `time_last` is
                `True`.
            time_last (`bool`, *optional*, defaults to `False`):
                Whether the last dimension of the input is the time/frame dimension or the Mel bins dimension.

        Returns:
            `torch.Tensor`:
                Audio waveform tensor of shape (batch_size, out_channels, audio_length)
        r?   r   r   r$   r   )dimg{Gz?)	transposeflattenrV   r   rP   r&   r'   r   rW   rQ   r8   stackrX   meanr\   tanh)r   r`   ra   r^   startendjresnet_outputss           r    r+   zLTX2Vocoder.forwardt   s$   $ )33Aq9M%--a3]3t//0ALLtGZGZ[M.DOOA.}=M 111Eq5D555C"[[RWX]_bRc)dRcQ/$,,q/-*HRc)djklN!JJ~1=M 1 ]4Hm4

=1 *es   E
)F)r1   r2   r3   __doc__r   r4   listr6   r   r8   r9   boolr+   r:   r;   s   @r    r=   r=   1   s      #+<&5)3-6	9,M+.$);Y;Y ;Y 	;Y
  $Cy;Y s);Y "#Y;Y tCy/;Y $);Y ";Y ;Yz*U\\ *d *u|| *r,   r=   )rR   r8   torch.nnr   torch.nn.functional
functionalr&   configuration_utilsr   r   models.modeling_utilsr   Moduler   r=    r,   r    <module>rw      s;         B /#ryy #Lm*k mr,   