a
    dB0                    @   s0  d Z ddlZddlZddlmZmZmZ ddlZddl	Z	ddl
m  mZ ddlZ	ddl	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mZ dd
lmZ ddlmZm Z m!Z!m"Z" ddl#m$Z$ e"%e&Z'dZ(dZ)dZ*g dZ+dZ,dZ-dZ.ddgZ/dZ0dZ1g dZ2dTee3e3f e4e3ee	j5 e3ej6dddZ7G dd dej8Z9G dd dej8Z:G dd  d ej8Z;G d!d" d"ej8Z<G d#d$ d$ej8Z=G d%d& d&ej8Z>G d'd( d(e>Z?G d)d* d*ej8Z@G d+d, d,ej8ZAG d-d. d.ej8ZBG d/d0 d0ej8ZCG d1d2 d2ej8ZDG d3d4 d4ej8ZEG d5d6 d6ej8ZFG d7d8 d8ej8ZGG d9d: d:ej8ZHG d;d< d<ej8ZIG d=d> d>eZJd?ZKd@ZLe dAeKG dBdC dCeJZMe dDeKG dEdF dFeJZNe dGeKG dHdI dIeJZOe dJeKG dKdL dLeJZPG dMdN dNej8ZQG dOdP dPej8ZRe dQeKG dRdS dSeJZSdS )Uz PyTorch WavLM model.    N)OptionalTupleUnion)nn)CrossEntropyLoss   )ACT2FN)is_deepspeed_zero3_enabled)BaseModelOutputCausalLMOutputSequenceClassifierOutputTokenClassifierOutputWav2Vec2BaseModelOutputXVectorOutput)PreTrainedModel)add_code_sample_docstringsadd_start_docstrings%add_start_docstrings_to_model_forwardlogging   )WavLMConfig   r   z1patrickvonplaten/wavlm-libri-clean-100h-base-plus)r   i$  i   zZ'mister quilter is the aposle of the middle classes and we are glad to welcome his gospel'gQ)@zmicrosoft/wavlm-base-plus-sdzmicrosoft/wavlm-base-plus-svg
ףp=
?)zmicrosoft/wavlm-basezmicrosoft/wavlm-base-pluszmicrosoft/wavlm-large)shape	mask_probmask_lengthattention_mask	min_masksreturnc                    s  | \}dk rt dkr6t d d dtjd   fdd}|durt|d	  nfd
dt|D }tj	|ft
d}g }	|}
|
dkr|S |D ]v}||}tjjt|d  |dd}t|dkrd }n|d }t|tj|
| tjd| g}|	| qt|	}	t|	dddddf ||
f}	|	||
 }	tddddf }t|||
f||
 }|	| }	|	 d kr҈d |	|	d k< t||	dd	 |S )af  
    Computes random mask spans for a given shape. Used to implement [SpecAugment: A Simple Data Augmentation Method for
    ASR](https://arxiv.org/abs/1904.08779). Note that this method is not optimized to run on TPU and should be run on
    CPU as part of the preprocessing during training.

    Args:
        shape: The shape for which to compute masks. This should be of a tuple of size 2 where
               the first element is the batch size and the second element is the length of the axis to span.
        mask_prob:  The percentage of the whole axis (between 0 and 1) which will be masked. The number of
                    independently generated mask spans of length `mask_length` is computed by
                    `mask_prob*shape[1]/mask_length`. Note that due to overlaps, `mask_prob` is an upper bound and the
                    actual percentage will be smaller.
        mask_length: size of the mask
        min_masks: minimum number of masked spans
        attention_mask: A (right-padded) attention mask which independently shortens the feature axis of
                        each batch dimension.
    r   z&`mask_length` has to be bigger than 0.zO`mask_length` has to be smaller than `sequence_length`, but got `mask_length`: z and `sequence_length`: `c                    sX   t |     }t|}| kr2 }| d  |k rTt| d  d}|S )z;Given input length, compute how many spans should be maskedr   r   )intmax)input_lengthnum_masked_spanepsilonr   r   r   sequence_length q/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/transformers/models/wavlm/modeling_wavlm.pycompute_num_masked_spanr   s    
z6_compute_mask_indices.<locals>.compute_num_masked_spanNc                    s   g | ]} qS r&   r&   .0_)r%   r&   r'   
<listcomp>       z)_compute_mask_indices.<locals>.<listcomp>dtyper   F)replace)
ValueErrornprandomZranditemsumdetachtolistrangezerosboolchoicearangelenZconcatenateonesint32appendarraybroadcast_toZreshaper    Zput_along_axis)r   r   r   r   r   
batch_sizer(   input_lengthsZspec_aug_maskZspec_aug_mask_idxsZmax_num_masked_spanr!   r"   Zspec_aug_mask_idxZdummy_mask_idxoffsetsr&   r#   r'   _compute_mask_indicesL   s\    

rG   c                       s&   e Zd Zd fdd	Zdd Z  ZS )WavLMNoLayerNormConvLayerr   c                    sj   t    |dkr |j|d  nd| _|j| | _tj| j| j|j| |j| |j	d| _
t|j | _d S )Nr   r   kernel_sizestridebias)super__init__conv_dimin_conv_dimout_conv_dimr   Conv1dconv_kernelconv_stride	conv_biasconvr   feat_extract_activation
activationselfconfiglayer_id	__class__r&   r'   rN      s    
z"WavLMNoLayerNormConvLayer.__init__c                 C   s   |  |}| |}|S N)rV   rX   rZ   hidden_statesr&   r&   r'   forward   s    

z!WavLMNoLayerNormConvLayer.forward)r   __name__
__module____qualname__rN   rb   __classcell__r&   r&   r]   r'   rH      s   rH   c                       s&   e Zd Zd fdd	Zdd Z  ZS )WavLMLayerNormConvLayerr   c                    s|   t    |dkr |j|d  nd| _|j| | _tj| j| j|j| |j| |j	d| _
tj| jdd| _t|j | _d S )Nr   r   rI   T)Zelementwise_affine)rM   rN   rO   rP   rQ   r   rR   rS   rT   rU   rV   	LayerNorm
layer_normr   rW   rX   rY   r]   r&   r'   rN      s    
z WavLMLayerNormConvLayer.__init__c                 C   s:   |  |}|dd}| |}|dd}| |}|S )Nr)   )rV   	transposerj   rX   r`   r&   r&   r'   rb      s    


zWavLMLayerNormConvLayer.forward)r   rc   r&   r&   r]   r'   rh      s   rh   c                       s&   e Zd Zd fdd	Zdd Z  ZS )WavLMGroupNormConvLayerr   c                    s   t    |dkr |j|d  nd| _|j| | _tj| j| j|j| |j| |j	d| _
t|j | _tj| j| jdd| _d S )Nr   r   rI   T)
num_groupsZnum_channelsZaffine)rM   rN   rO   rP   rQ   r   rR   rS   rT   rU   rV   r   rW   rX   	GroupNormrj   rY   r]   r&   r'   rN      s    
z WavLMGroupNormConvLayer.__init__c                 C   s"   |  |}| |}| |}|S r_   )rV   rj   rX   r`   r&   r&   r'   rb     s    


zWavLMGroupNormConvLayer.forward)r   rc   r&   r&   r]   r'   rm      s   rm   c                       s$   e Zd Z fddZdd Z  ZS )WavLMPositionalConvEmbeddingc                    s   t    tj|j|j|j|jd |jd| _tjj	}t
tjjdrNtjjj	}t rdd l}|jj| jjdd" || jddd| _W d    n1 s0    Y  |j| | jj |j| | jj n|| jddd| _t|j| _t|j | _d S )Nr   )rJ   paddinggroupsweight_normr   )Zmodifier_rankweight)namedim)rM   rN   r   rR   hidden_sizenum_conv_pos_embeddingsZnum_conv_pos_embedding_groupsrV   utilsrs   hasattrZparametrizationsr	   	deepspeedZzeroZGatheredParametersrt   Zregister_external_parameterZweight_vZweight_gWavLMSamePadLayerrq   r   rW   rX   )rZ   r[   rs   r{   r]   r&   r'   rN     s(    

0z%WavLMPositionalConvEmbedding.__init__c                 C   s:   | dd}| |}| |}| |}| dd}|S Nr   r   )rl   rV   rq   rX   r`   r&   r&   r'   rb   +  s    


z$WavLMPositionalConvEmbedding.forwardrc   r&   r&   r]   r'   rp     s   rp   c                       s$   e Zd Z fddZdd Z  ZS )r|   c                    s$   t    |d dkrdnd| _d S Nr   r   r   )rM   rN   num_pad_remove)rZ   rx   r]   r&   r'   rN   8  s    
zWavLMSamePadLayer.__init__c                 C   s,   | j dkr(|d d d d d | j  f }|S )Nr   )r   r`   r&   r&   r'   rb   <  s    
zWavLMSamePadLayer.forwardrc   r&   r&   r]   r'   r|   7  s   r|   c                       s0   e Zd ZdZ fddZdd Zdd Z  ZS )WavLMFeatureEncoderz.Construct the features from raw audio waveformc                    s   t     jdkr@t ddg fddt jd D  }n6 jdkrd fddt jD }ntd	 j d
t|| _	d| _
d| _d S )Ngroupr   r\   c                    s   g | ]}t  |d  dqS )r   r   )rH   r+   ir[   r&   r'   r-   J  s   z0WavLMFeatureEncoder.__init__.<locals>.<listcomp>r   layerc                    s   g | ]}t  |d qS )r   )rh   r   r   r&   r'   r-   N  r.   z`config.feat_extract_norm` is z), but has to be one of ['group', 'layer']FT)rM   rN   Zfeat_extract_normrm   r9   Znum_feat_extract_layersr2   r   
ModuleListconv_layersgradient_checkpointing_requires_grad)rZ   r[   r   r]   r   r'   rN   F  s    



zWavLMFeatureEncoder.__init__c                 C   s   |   D ]
}d|_qd| _d S )NF)
parametersrequires_gradr   rZ   paramr&   r&   r'   _freeze_parametersW  s    z&WavLMFeatureEncoder._freeze_parametersc                 C   sj   |d d d f }| j r"| jr"d|_| jD ]<}| j r\| jr\| jr\dd }tjj|||}q(||}q(|S )NTc                    s    fdd}|S )Nc                     s    |  S r_   r&   inputsmoduler&   r'   custom_forwardg  s    zRWavLMFeatureEncoder.forward.<locals>.create_custom_forward.<locals>.custom_forwardr&   r   r   r&   r   r'   create_custom_forwardf  s    z:WavLMFeatureEncoder.forward.<locals>.create_custom_forward)r   trainingr   r   r   torchry   
checkpoint)rZ   input_valuesra   Z
conv_layerr   r&   r&   r'   rb   \  s    

zWavLMFeatureEncoder.forward)rd   re   rf   __doc__rN   r   rb   rg   r&   r&   r]   r'   r   C  s   r   c                       s   e Zd Z fddZ  ZS )WavLMFeatureExtractorc                    s8   t  | td| jj d| jjd j dt d S )NzThe class `zD` has been depreciated and will be removed in Transformers v5. Use `r   z
` instead.)rM   rN   warningswarnr^   rd   	__bases__FutureWarningrZ   r[   r]   r&   r'   rN   w  s    zWavLMFeatureExtractor.__init__)rd   re   rf   rN   rg   r&   r&   r]   r'   r   v  s   r   c                       s$   e Zd Z fddZdd Z  ZS )WavLMFeatureProjectionc                    sJ   t    tj|jd |jd| _t|jd |j| _	t
|j| _d S )Nr)   Zeps)rM   rN   r   ri   rO   layer_norm_epsrj   Linearrw   
projectionDropoutZfeat_proj_dropoutdropoutr   r]   r&   r'   rN     s    
zWavLMFeatureProjection.__init__c                 C   s&   |  |}| |}| |}||fS r_   )rj   r   r   )rZ   ra   Znorm_hidden_statesr&   r&   r'   rb     s    


zWavLMFeatureProjection.forwardrc   r&   r&   r]   r'   r     s   r   c                       s   e Zd ZdZdeeeeeed fddZdej	e
ej	 e
ej	 eeej	e
ej	 e
eej	  f dddZejeejejf ejeejejfdddZeeejdddZejejdddZ  ZS )WavLMAttentionz=Multi-headed attention from 'Attention Is All You Need' paper        @     T	embed_dim	num_headsr   num_bucketsmax_distancehas_relative_position_biasc                    s   t    || _|| _|| _|| | _| j| | jkrNtd| j d| d| jd | _t	||| _
t	||| _t	||| _t	||| _|| _|| _ttd| jdd| _t	| jd| _|rt| j| j| _d S )Nz;embed_dim must be divisible by num_heads (got `embed_dim`: z and `num_heads`: z).g      r      )rM   rN   r   r   r   Zhead_dimr2   Zscalingr   r   k_projv_projq_projout_projr   r   	Parameterr   r?   gru_rel_pos_constgru_rel_pos_linearZ	Embeddingrel_attn_embed)rZ   r   r   r   r   r   r   r]   r&   r'   rN     s,    	


zWavLMAttention.__init__NFr   )ra   r   position_biasoutput_attentionsr   c                 C   s  |  \}}}|du rH| ||}|d|ddd|| j ||}||jdd | jdf }	|	dddd}	| |	}
|
|	jdd d 	d}
t
|
jddd\}}||| j d	  d
 }||| j dd| }|d||f}| ||||\}}|||fS )z'Attention layer with relative attentionNr   r   r)   r   r   )r      rv         ?g       @)sizecompute_bias	unsqueezerepeatviewr   r   permuter   r6   r   Zsigmoidchunkr   torch_multi_head_self_attention)rZ   ra   r   r   r   indexZbszZtgt_lenr,   Zgated_hidden_statesZrelative_position_projZgate_aZgate_bZgate_outputgated_position_biasattn_outputattn_weightsr&   r&   r'   rb     s"    	$
zWavLMAttention.forward)ra   r   r   r   r   c                 C   s   | dd } }}|dur&|dnd}d }	}
d}tj|||| j| jtdgt| j	j
| jj
| jj
f|	|
|| j| jj| jj
| j|||d| j	j| jj| jjd\}}| dd}|dur|dddf |jdd | jf |jdd  }||fS )zCsimple wrapper around torch's multi_head_attention_forward functionr   r   NFT)Zuse_separate_proj_weightZq_proj_weightZk_proj_weightZv_proj_weight)rl   neFZmulti_head_attention_forwardr   r   r   emptycatr   rL   r   r   r   r   rt   r   rC   r   )rZ   ra   r   r   r   querykeyvalueZkey_padding_maskZbias_kZbias_vZadd_zero_attnr   r   r&   r&   r'   r     sB    	

"z.WavLMAttention.torch_multi_head_self_attention)query_length
key_lengthr   c                 C   sv   t j|t jdd d d f }t j|t jdd d d f }|| }| |}|| jjj}| |}|g d}|S )Nr/   )r   r   r   )	r   r=   long_relative_positions_buckettor   rt   devicer   )rZ   r   r   Zcontext_positionZmemory_positionZrelative_positionZrelative_position_bucketvaluesr&   r&   r'   r     s    

zWavLMAttention.compute_bias)relative_positionsr   c                 C   s   | j d }|dktj| }t|}|d }||k }t| | }|t| j|  }|||  }|| tj}t	|t
||d }|t|||7 }|S r~   )r   r   r   r   abslogfloatmathr   minZ	full_likewhere)rZ   r   r   Zrelative_bucketsZ	max_exactZis_smallZrelative_positions_if_largeZrelative_position_if_larger&   r&   r'   r   "  s    

z)WavLMAttention._relative_positions_bucket)r   r   r   T)NNFr   )rd   re   rf   r   r   r   r;   rN   r   Tensorr   r   rb   FloatTensorr   
LongTensorZ
BoolTensorr   r   r   rg   r&   r&   r]   r'   r     s@       '    +
7
r   c                       s$   e Zd Z fddZdd Z  ZS )WavLMFeedForwardc                    sp   t    t|j| _t|j|j| _	t
|jtrDt|j | _n|j| _t|j|j| _t|j| _d S r_   )rM   rN   r   r   Zactivation_dropoutintermediate_dropoutr   rw   Zintermediate_sizeintermediate_dense
isinstanceZ
hidden_actstrr   intermediate_act_fnoutput_densehidden_dropoutoutput_dropoutr   r]   r&   r'   rN   9  s    
zWavLMFeedForward.__init__c                 C   s6   |  |}| |}| |}| |}| |}|S r_   )r   r   r   r   r   r`   r&   r&   r'   rb   F  s    




zWavLMFeedForward.forwardrc   r&   r&   r]   r'   r   8  s   r   c                       s0   e Zd Zd
eed fddZddd	Z  ZS )WavLMEncoderLayerTr[   r   c                    sn   t    t|j|j|j|j|j|d| _t	
|j| _t	j|j|jd| _t|| _t	j|j|jd| _d S Nr   r   rM   rN   r   rw   Znum_attention_headsZattention_dropoutr   Zmax_bucket_distance	attentionr   r   r   r   ri   r   rj   r   feed_forwardfinal_layer_normrZ   r[   r   r]   r&   r'   rN   Q  s    

zWavLMEncoderLayer.__init__NFr   c           	      C   sl   |}| j |||||d\}}}| |}|| }| |}|| | }| |}||f}|rh||f7 }|S )Nr   r   r   r   )r   r   rj   r   r   )	rZ   ra   r   r   r   r   attn_residualr   outputsr&   r&   r'   rb   `  s"    



zWavLMEncoderLayer.forward)T)NNFr   rd   re   rf   r   r;   rN   rb   rg   r&   r&   r]   r'   r   P  s   r   c                       s0   e Zd Zd	eed fddZd
ddZ  ZS ) WavLMEncoderLayerStableLayerNormTr   c                    sn   t    t|j|j|j|j|j|d| _t	
|j| _t	j|j|jd| _t|| _t	j|j|jd| _d S r   r   r   r]   r&   r'   rN   z  s    

z)WavLMEncoderLayerStableLayerNorm.__init__NFc                 C   sf   |}|  |}| j||||d\}}}| |}|| }|| | | }||f}|rb||f7 }|S )N)r   r   r   )rj   r   r   r   r   )rZ   ra   r   r   r   r   r   r   r&   r&   r'   rb     s    


z(WavLMEncoderLayerStableLayerNorm.forward)T)NNFr   r&   r&   r]   r'   r   y  s   r   c                       s&   e Zd Z fddZdddZ  ZS )	WavLMEncoderc                    sf   t     | _t | _tj j jd| _	t
 j| _t fddt jD | _d| _d S )Nr   c                    s   g | ]}t  |d kdqS r   )r   )r   r   r   r&   r'   r-     r.   z)WavLMEncoder.__init__.<locals>.<listcomp>FrM   rN   r[   rp   pos_conv_embedr   ri   rw   r   rj   r   r   r   r   r9   num_hidden_layerslayersr   r   r]   r   r'   rN     s    

zWavLMEncoder.__init__NFTc                    sX  |rdnd } rdnd }|d ur*d|| < |  |}|| }| |}| |}t }	d }
t| jD ]\}}|rz||f }tjdd}| j	o|dko|| j
jk }|r|	r| jr| j	r܇ fdd}tjj|||||
}n||||
 |d}|d d \}}
|rd	} rd||d f }qd|r,||f }|sJtd
d |||fD S t|||dS )Nr&   r   r   r   c                    s    fdd}|S )Nc                     s    g | R  S r_   r&   r   r   r   r&   r'   r     s    zKWavLMEncoder.forward.<locals>.create_custom_forward.<locals>.custom_forwardr&   r   r   r   r'   r     s    z3WavLMEncoder.forward.<locals>.create_custom_forwardr   r   NNc                 s   s   | ]}|d ur|V  qd S r_   r&   r+   vr&   r&   r'   	<genexpr>  r.   z'WavLMEncoder.forward.<locals>.<genexpr>last_hidden_statera   
attentions)r   rj   r   r	   	enumerater   r3   r4   uniformr   r[   	layerdropr   r   ry   r   tupler
   rZ   ra   r   r   output_hidden_statesreturn_dictZall_hidden_statesZall_self_attentionsZposition_embeddingsZdeepspeed_zero3_is_enabledr   r   r   Zdropout_probabilityZskip_the_layerr   Zlayer_outputsr&   r  r'   rb     sZ    





zWavLMEncoder.forward)NFFTrc   r&   r&   r]   r'   r     s       r   c                       s&   e Zd Z fddZdddZ  ZS )	WavLMEncoderStableLayerNormc                    sf   t     | _t | _tj j jd| _	t
 j| _t fddt jD | _d| _d S )Nr   c                    s   g | ]}t  |d kdqS r   )r   r   r   r&   r'   r-     s   z8WavLMEncoderStableLayerNorm.__init__.<locals>.<listcomp>Fr   r   r]   r   r'   rN     s    


z$WavLMEncoderStableLayerNorm.__init__NFTc                    sT  |rdnd } rdnd }|d ur*d|| < |  |}|| }| |}t }	d }
t| jD ]\}}|rp||f }tjdd}| jo|dko|| j	j
k }|r|	r| jr| jr҇ fdd}tjj|||||
}n||| |
d}|d d \}}
|rd} rZ||d f }qZ| |}|r(||f }|sFtd	d
 |||fD S t|||dS )Nr&   r   r   c                    s    fdd}|S )Nc                     s    g | R  S r_   r&   r   r  r&   r'   r   )  s    zZWavLMEncoderStableLayerNorm.forward.<locals>.create_custom_forward.<locals>.custom_forwardr&   r   r  r   r'   r   (  s    zBWavLMEncoderStableLayerNorm.forward.<locals>.create_custom_forward)r   r   r   r   r  c                 s   s   | ]}|d ur|V  qd S r_   r&   r  r&   r&   r'   r  I  r.   z6WavLMEncoderStableLayerNorm.forward.<locals>.<genexpr>r  )r   r   r	   r
  r   r3   r4   r  r   r[   r  r   r   ry   r   rj   r  r
   r  r&   r  r'   rb     sT    





z#WavLMEncoderStableLayerNorm.forward)NFFTrc   r&   r&   r]   r'   r    s       r  c                       s4   e Zd ZdZ fddZedd Zdd Z  ZS )WavLMGumbelVectorQuantizerz
    Vector quantization using gumbel softmax. See [CATEGORICAL REPARAMETERIZATION WITH
    GUMBEL-SOFTMAX](https://arxiv.org/pdf/1611.01144.pdf) for more information.
    c                    s   t    |j| _|j| _|j| j dkrDtd|j d| j dt	t
d| j| j |j| j | _t|jd | j| j | _d| _d S )Nr   z`config.codevector_dim z5 must be divisible by `config.num_codevector_groups` z for concatenation.r   r)   r   )rM   rN   Znum_codevector_groupsrn   Znum_codevectors_per_groupnum_varsZcodevector_dimr2   r   r   r   r   codevectorsr   rO   weight_projtemperaturer   r]   r&   r'   rN   U  s    

z#WavLMGumbelVectorQuantizer.__init__c                 C   s8   | j dd}ttj|t|d  dd  }|S )Nr   r   gHz>r)   )meanr   expr6   r   )ZprobsZmarginal_probs
perplexityr&   r&   r'   _compute_perplexityj  s    (z.WavLMGumbelVectorQuantizer._compute_perplexityc                 C   s  |j \}}}| |}||| | j d}| jrtjj| | j	dd}|
|}tj||| | jd dd}| |}nH|jdd}|j|j  d|ddd}||| | jd}| |}||| d}|d| j }	|	|| | j| jd}
|
d||d}
|
|fS )Nr)   T)tauhardr   r   r   rk   )r   r  r   rn   r   r   
functionalZgumbel_softmaxr   r  Ztype_asr   softmaxr  argmaxZ	new_zerosZscatter_r   r  r  r6   )rZ   ra   rD   r%   rw   Zcodevector_probsZcodevector_soft_distr  Zcodevector_idxZcodevectors_per_groupr  r&   r&   r'   rb   p  s*    


z"WavLMGumbelVectorQuantizer.forward)	rd   re   rf   r   rN   staticmethodr  rb   rg   r&   r&   r]   r'   r  O  s
   
r  c                       s$   e Zd Z fddZdd Z  ZS )WavLMAdapterc                    sp   t     j jkr8t j j| _t j| _nd  | _| _t	 fddt
 jD | _ j| _d S )Nc                 3   s   | ]}t  V  qd S r_   )WavLMAdapterLayerr*   r   r&   r'   r    r.   z(WavLMAdapter.__init__.<locals>.<genexpr>)rM   rN   output_hidden_sizerw   r   r   projri   proj_layer_normr   r9   num_adapter_layersr   r  r   r]   r   r'   rN     s    
 zWavLMAdapter.__init__c                 C   sr   | j d ur(| jd ur(|  |}| |}|dd}| jD ]&}tj }| jrX|| jkr:||}q:|dd}|S r}   )r$  r%  rl   r   r3   r4   r   r  )rZ   ra   r   Zlayerdrop_probr&   r&   r'   rb     s    




zWavLMAdapter.forwardrc   r&   r&   r]   r'   r!    s   r!  c                       s$   e Zd Z fddZdd Z  ZS )r"  c                    s0   t    tj|jd|j |j|jdd| _d S )Nr   r   )rK   rq   )rM   rN   r   rR   r#  Zadapter_kernel_sizeadapter_striderV   r   r]   r&   r'   rN     s    
zWavLMAdapterLayer.__init__c                 C   s   |  |}tjj|dd}|S )Nr   r   )rV   r   r  Zglur`   r&   r&   r'   rb     s    
zWavLMAdapterLayer.forwardrc   r&   r&   r]   r'   r"    s   
r"  c                   @   sl   e Zd ZdZeZdZdZdgZdZ	dd Z
deejef ee d	d
dZdeejdddZdddZdS )WavLMPreTrainedModelz
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    wavlmr   Zposition_idsTc              	   C   s  t |tr>|jjjjddd |jjj  tj	
|j njt |trtj	j|jjddtd|jjd |jj   d tj	|jjd nt |trtd|jj }tj	j
|jj| |d tj	j
|jj| |d nt |tjr|jjjd| jjd |jdur|jj  nt |tjtjfrN|jj  |jjd nZt |tjrtj	|j |jdurt|j|j|jd   }tj	j
|j| |d dS )	zInitialize the weightsr   r   )r  stdr   r   )abNr   )r   r  r  rt   dataZnormal_rL   Zzero_r   inituniform_r  rp   rV   r   sqrtrJ   Zin_channelsZ	constant_r   r   Zin_featuresr   r[   Zinitializer_rangeri   ro   Zfill_rR   Zkaiming_normal_rr   )rZ   r   kr&   r&   r'   _init_weights  s6    

 
z"WavLMPreTrainedModel._init_weightsN)rE   add_adapterc                 C   sn   |du r| j jn|}dd }t| j j| j jD ]\}}||||}q.|rjt| j jD ]}||d| j j}qT|S )zH
        Computes the output length of the convolutional layers
        Nc                 S   s   t j| | |ddd S )Nfloor)Zrounding_moder   )r   divr!   rJ   rK   r&   r&   r'   _conv_out_length  s    zOWavLMPreTrainedModel._get_feat_extract_output_lengths.<locals>._conv_out_lengthr   )r[   r3  ziprS   rT   r9   r&  r'  )rZ   rE   r3  r7  rJ   rK   r,   r&   r&   r'    _get_feat_extract_output_lengths  s    z5WavLMPreTrainedModel._get_feat_extract_output_lengths)feature_vector_lengthr   c                 C   s   |j ddd d df }| j||d}|tj}|jd }tj||f|j|jd}d|tj	|jd |jd|d f< |
dg d
dg }|S )Nr)   r   r3  r   )r0   r   r   )r   )Zcumsumr9  r   r   r   r   r:   r0   r   r=   Zflipr;   )rZ   r:  r   r3  Znon_padded_lengthsZoutput_lengthsrD   r&   r&   r'   "_get_feature_vector_attention_mask  s    
"z7WavLMPreTrainedModel._get_feature_vector_attention_maskFc                 C   s   t |tttfr||_d S r_   )r   r   r  r   r   )rZ   r   r   r&   r&   r'   _set_gradient_checkpointing   s    z0WavLMPreTrainedModel._set_gradient_checkpointing)N)N)F)rd   re   rf   r   r   config_classZbase_model_prefixZmain_input_nameZ_keys_to_ignore_on_load_missingZsupports_gradient_checkpointingr2  r   r   r   r   r   r;   r9  r<  r=  r&   r&   r&   r'   r(    s    "  r(  a  
    WavLM was proposed in [WavLM: Unified Speech Representation Learning with Labeled and Unlabeled
    Data](https://arxiv.org/abs/2101.07597) by Chengyi Wang, Yu Wu, Yao Qian, Kenichi Kumatani, Shujie Liu, Furu Wei,
    Michael Zeng, Xuedong Huang.

    This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the
    library implements for all its model (such as downloading or saving etc.).

    This model is a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) sub-class. Use
    it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage and
    behavior.

    Parameters:
        config ([`WavLMConfig`]): Model configuration class with all the parameters of the model.
            Initializing with a config file does not load the weights associated with the model, only the
            configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.
aI  
    Args:
        input_values (`torch.FloatTensor` of shape `(batch_size, sequence_length)`):
            Float values of input raw speech waveform. Values can be obtained by loading a `.flac` or `.wav` audio file
            into an array of type `List[float]` or a `numpy.ndarray`, *e.g.* via the soundfile library (`pip install
            soundfile`). To prepare the array into `input_values`, the [`AutoProcessor`] should be used for padding and
            conversion into a tensor of type `torch.FloatTensor`. See [`Wav2Vec2Processor.__call__`] for details.
        attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Mask to avoid performing convolution and attention on padding token indices. Mask values selected in `[0,
            1]`:

            - 1 for tokens that are **not masked**,
            - 0 for tokens that are **masked**.

            [What are attention masks?](../glossary#attention-mask)

            <Tip warning={true}>

            `attention_mask` should only be passed if the corresponding processor has `config.return_attention_mask ==
            True`. For all models whose processor has `config.return_attention_mask == False`, `attention_mask` should
            **not** be passed to avoid degraded performance when doing batched inference. For such models
            `input_values` should simply be padded with 0 and passed without `attention_mask`. Be aware that these
            models also yield slightly different results depending on whether `input_values` is padded or not.

            </Tip>

        output_attentions (`bool`, *optional*):
            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
            tensors for more detail.
        output_hidden_states (`bool`, *optional*):
            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
            more detail.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
z_The bare WavLM Model transformer outputting raw hidden-states without any specific head on top.c                       s   e Zd Zed fddZdd Zdd Zdeje	ej e	ej
 d	d
dZeeeeeededde	ej e	ej e	ej e	e e	e e	e eeef dddZ  ZS )
WavLMModelr   c                    s   t  | || _t|| _t|| _|jdks:|jdkrRt	
t|j | _|jrdt|| _n
t|| _|jr|t|nd | _|   d S )Nr   )rM   rN   r[   r   feature_extractorr   feature_projectionmask_time_probmask_feature_probr   r   r   r   rw   r/  masked_spec_embedZdo_stable_layer_normr  encoderr   r3  r!  adapter	post_initr   r]   r&   r'   rN   c  s    


zWavLMModel.__init__c                 C   s   t dt |   dS z
        Calling this function will disable the gradient computation for the feature encoder so that its parameters will
        not be updated during training.
        The method `freeze_feature_extractor` is deprecated and will be removed in Transformers v5.Please use the equivalent `freeze_feature_encoder` method instead.Nr   r   r   freeze_feature_encoderrZ   r&   r&   r'   freeze_feature_extractorw  s
    z#WavLMModel.freeze_feature_extractorc                 C   s   | j   dS 
        Calling this function will disable the gradient computation for the feature encoder so that its parameter will
        not be updated during training.
        N)r@  r   rL  r&   r&   r'   rK    s    z!WavLMModel.freeze_feature_encoderN)ra   mask_time_indicesr   c                 C   s  t | jdds|S | \}}}|dur<| j|j||< nZ| jjdkr| jrt||f| jj| jj	|| jj
d}tj||jtjd}| j|j||< | jjdkr| jrt||f| jj| jj| jjd}tj||jtjd}|dddf d|d}d||< |S )	z
        Masks extracted features along time axis and/or along feature axis according to
        [SpecAugment](https://arxiv.org/abs/1904.08779).
        Zapply_spec_augmentTNr   )r   r   r   r   )r   r0   )r   r   r   r)   )getattrr[   r   rD  r   r0   rB  r   rG   Zmask_time_lengthZmask_time_min_masksr   Ztensorr   r;   rC  Zmask_feature_lengthZmask_feature_min_masksexpand)rZ   ra   rP  r   rD   r%   rw   Zmask_feature_indicesr&   r&   r'   _mask_hidden_states  s4    zWavLMModel._mask_hidden_statesaudior   output_typer>  modalityexpected_output)r   r   rP  r   r  r  r   c           
      C   s   |d ur|n| j j}|d ur |n| j j}|d ur4|n| j j}| |}|dd}|d urp| j|jd |dd}| |\}}| j	|||d}| j
|||||d}	|	d }| jd ur| |}|s||f|	dd   S t|||	j|	jdS )	Nr   r   Fr;  )rP  r   r   r   r  r  r   )r  extract_featuresra   r	  )r[   r   r  use_return_dictr@  rl   r<  r   rA  rS  rE  rF  r   ra   r	  )
rZ   r   r   rP  r   r  r  rZ  ra   Zencoder_outputsr&   r&   r'   rb     s@    


zWavLMModel.forward)NN)NNNNN)rd   re   rf   r   rN   rM  rK  r   r   r   r   rS  r   WAVLM_INPUTS_DOCSTRINGr   _CHECKPOINT_FOR_DOCr   _CONFIG_FOR_DOC_EXPECTED_OUTPUT_SHAPEr   r;   r   r   rb   rg   r&   r&   r]   r'   r?  ]  s@   
  .
     
r?  zcWavLM Model with a `language modeling` head on top for Connectionist Temporal Classification (CTC).c                       s   e Zd Zd fdd	Zdd Zdd Zeeee	e
eeeddeej eej ee ee ee eej eee
f d	d
dZ  ZS )WavLMForCTCNc                    s   t  | t|| _t|j| _|jd u r@t	d| j
 dt|drV|jrV|jn|j}t||j| _|d urt| jdd d u rt	d| dn8|d u rt| jdd d urtd n|d ur| | |   d S )NzYou are trying to instantiate z with a configuration that does not define the vocabulary size of the language model head. Please instantiate the model as follows: `WavLMForCTC.from_pretrained(..., vocab_size=vocab_size)`. or define `vocab_size` of your model's configuration.r3  Zadapter_attn_dimzCannot pass `target_lang`: z- if `config.adapter_attn_dim` is not defined.z)By default `target_lang` is set to 'eng'.)rM   rN   r?  r)  r   r   Zfinal_dropoutr   
vocab_sizer2   r^   rz   r3  r#  rw   r   lm_headrQ  r[   loggerinfoZload_adapterrG  )rZ   r[   Ztarget_langr#  r]   r&   r'   rN     s"    


zWavLMForCTC.__init__c                 C   s   t dt |   dS rO  rI  NrJ  rL  r&   r&   r'   rM    s
    z$WavLMForCTC.freeze_feature_extractorc                 C   s   | j j  dS rN  r)  r@  r   rL  r&   r&   r'   rK  #  s    z"WavLMForCTC.freeze_feature_encoder)r   rV  r>  rX  Zexpected_lossr   r   r   r  r  labelsr   c              
   C   s|  |dur|n| j j}| j|||||d}|d }| |}| |}	d}
|dur8| | j jkrttd| j j |dur|ntj	|tj
d}| |dtj
}|dk}|d}||}tjj|	dtjddd}tjjjd	d
6 tjj||||| j j| j j| j jd}
W d   n1 s.0    Y  |sh|	f|td  }|
durd|
f| S |S t|
|	|j|jdS )a  
        labels (`torch.LongTensor` of shape `(batch_size, target_length)`, *optional*):
            Labels for connectionist temporal classification. Note that `target_length` has to be smaller or equal to
            the sequence length of the output logits. Indices are selected in `[-100, 0, ..., config.vocab_size - 1]`.
            All labels set to `-100` are ignored (masked), the loss is only computed for labels in `[0, ...,
            config.vocab_size - 1]`.
        NrY  r   z$Label values must be <= vocab_size: r/   r)   )rv   r0   r   F)Zenabled)ZblankZ	reductionZzero_infinitylosslogitsra   r	  )r[   r[  r)  r   rb  r    ra  r2   r   Z	ones_liker   r9  r6   r   Zmasked_selectr   r  Zlog_softmaxfloat32rl   backendsZcudnnflagsZctc_lossZpad_token_idZctc_loss_reductionZctc_zero_infinity_HIDDEN_STATES_START_POSITIONr   ra   r	  )rZ   r   r   r   r  r  rh  r   ra   rk  rj  rE   Zlabels_maskZtarget_lengthsZflattened_targetsZ	log_probsoutputr&   r&   r'   rb   *  sL    




&
zWavLMForCTC.forward)N)NNNNN)rd   re   rf   rN   rM  rK  r   r\  r   r]  r   r^  _CTC_EXPECTED_OUTPUT_CTC_EXPECTED_LOSSr   r   r   r;   r   r   rb   rg   r&   r&   r]   r'   r`    s2   
     
r`  z
    WavLM Model with a sequence classification head on top (a linear layer over the pooled output) for tasks like
    SUPERB Keyword Spotting.
    c                       s   e Zd Z fddZdd Zdd Zdd Zeee	e
eed	d
deej eej ee ee ee eej eeef dddZ  ZS )WavLMForSequenceClassificationc                    s   t  | t|dr$|jr$tdt|| _|jd }|jrTt	
t|| | _t	|j|j| _t	|j|j| _|   d S )Nr3  z\Sequence classification does not support the use of WavLM adapters (config.add_adapter=True)r   )rM   rN   rz   r3  r2   r?  r)  r   use_weighted_layer_sumr   r   r   r?   layer_weightsr   rw   Zclassifier_proj_size	projector
num_labels
classifierrG  rZ   r[   
num_layersr]   r&   r'   rN     s    

z'WavLMForSequenceClassification.__init__c                 C   s   t dt |   dS rH  rJ  rL  r&   r&   r'   rM    s
    z7WavLMForSequenceClassification.freeze_feature_extractorc                 C   s   | j j  dS rN  rf  rL  r&   r&   r'   rK    s    z5WavLMForSequenceClassification.freeze_feature_encoderc                 C   s   | j  D ]
}d|_q
dS z
        Calling this function will disable the gradient computation for the base model so that its parameters will not
        be updated during training. Only the classification head will be updated.
        FNr)  r   r   r   r&   r&   r'   freeze_base_model  s    z0WavLMForSequenceClassification.freeze_base_modelrT  )r   rV  r>  rW  Nrg  c                 C   sf  |dur|n| j j}| j jr dn|}| j|||||d}| j jr|t }tj|dd}tjj	| j
dd}	||	ddd jdd}n|d }| |}|du r|jdd}
n<| |jd |}d|| < |jdd|jdddd }
| |
}d}|dur"t }||d| j j|d}|sR|f|td  }|durN|f| S |S t|||j|jd	S )
  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
        NTrY  r   r   r)   r   r   ri  )r[   r[  rt  r)  ro  r   stackr   r  r  ru  r   r6   rv  r  r<  r   rx  r   rw  r   ra   r	  )rZ   r   r   r   r  r  rh  r   ra   norm_weightsZpooled_outputZpadding_maskrk  rj  loss_fctrp  r&   r&   r'   rb     sF    

 

z&WavLMForSequenceClassification.forward)NNNNN)rd   re   rf   rN   rM  rK  r}  r   r\  r   r]  r   r^  r   r   r   r;   r   r   rb   rg   r&   r&   r]   r'   rs  z  s2   
     
rs  za
    WavLM Model with a frame classification head on top for tasks like Speaker Diarization.
    c                       s   e Zd Z fddZdd Zdd Zdd Zeee	e
eed	ed
deej eej eej ee ee ee eeef dddZ  ZS ) WavLMForAudioFrameClassificationc                    sz   t  | t|dr$|jr$tdt|| _|jd }|jrTt	
t|| | _t	|j|j| _|j| _|   d S )Nr3  z_Audio frame classification does not support the use of WavLM adapters (config.add_adapter=True)r   )rM   rN   rz   r3  r2   r?  r)  r   rt  r   r   r   r?   ru  r   rw   rw  rx  init_weightsry  r]   r&   r'   rN     s    

z)WavLMForAudioFrameClassification.__init__c                 C   s   t dt |   dS re  rJ  rL  r&   r&   r'   rM    s
    z9WavLMForAudioFrameClassification.freeze_feature_extractorc                 C   s   | j j  dS rN  rf  rL  r&   r&   r'   rK    s    z7WavLMForAudioFrameClassification.freeze_feature_encoderc                 C   s   | j  D ]
}d|_q
dS r{  r|  r   r&   r&   r'   r}  "  s    z2WavLMForAudioFrameClassification.freeze_base_modelrT  rU  N)r   r   rh  r   r  r  r   c                 C   s   |dur|n| j j}| j jr dn|}| j|||||d}| j jr|t }tj|dd}tjj	| j
dd}	||	ddd jdd}n|d }| |}
d}|durt }||
d| jtj|d| jdd}|s|
f|td  }|S t||
|j|jd	S )
r~  NTrY  r   r   r)   r   )Zaxisri  )r[   r[  rt  r)  ro  r   r  r   r  r  ru  r   r6   rx  r   rw  r  r   ra   r	  )rZ   r   r   rh  r   r  r  r   ra   r  rk  rj  r  rp  r&   r&   r'   rb   *  s:    
(z(WavLMForAudioFrameClassification.forward)NNNNN)rd   re   rf   rN   rM  rK  r}  r   r\  r   _FRAME_CLASS_CHECKPOINTr   r^  _FRAME_EXPECTED_OUTPUTr   r   r   r;   r   r   rb   rg   r&   r&   r]   r'   r    s4   
     
r  c                       s&   e Zd Zd fdd	Zdd Z  ZS )AMSoftmaxLoss      >@皙?c                    sF   t t|   || _|| _|| _tjt	||dd| _
t | _d S )NT)r   )rM   r  rN   scalemarginrw  r   r   r   Zrandnrt   r   rj  )rZ   Z	input_dimrw  r  r  r]   r&   r'   rN   j  s    zAMSoftmaxLoss.__init__c           	      C   sx   |  }tjj| jdd}tjj|dd}t||}|| j }tj|| j	}| j
t| || }| ||}|S )Nr   r   r   )flattenr   r  	normalizert   r   mmr  Zone_hotrw  r  r   r;   rj  )	rZ   ra   rh  rt   Z	cos_thetapsiZonehotrk  rj  r&   r&   r'   rb   r  s    
zAMSoftmaxLoss.forward)r  r  rc   r&   r&   r]   r'   r  i  s   r  c                       s&   e Zd Zd fdd	Zdd Z  ZS )	TDNNLayerr   c                    sv   t    |dkr |j|d  n|j| | _|j| | _|j| | _|j| | _t	
| j| j | j| _t	 | _d S )Nr   r   )rM   rN   tdnn_dimrP   rQ   tdnn_kernelrJ   Ztdnn_dilationdilationr   r   kernelZReLUrX   rY   r]   r&   r'   rN     s    
"zTDNNLayer.__init__c                 C   sV   | d}tjj|| j| jfd| jf| jdfd}|dd}| |}| 	|}|S )Nr   )rK   r  r   )
r   r   r  ZunfoldrJ   rP   r  rl   r  rX   r`   r&   r&   r'   rb     s    



zTDNNLayer.forward)r   rc   r&   r&   r]   r'   r    s   
r  zi
    WavLM Model with an XVector feature extraction head on top for tasks like Speaker Verification.
    c                       s   e Zd Z fddZdd Zdd Zdd Zeej	e
f d	d
dZeeeeeededdeej eej ee ee ee eej eeef dddZ  ZS )WavLMForXVectorc                    s   t    t | _ jd } jr<tt	|| | _
t j jd | _ fddtt jD }t|| _t jd d  j| _t j j| _t j j| _|   d S )Nr   r   c                    s   g | ]}t  |qS r&   )r  r   r   r&   r'   r-     r.   z,WavLMForXVector.__init__.<locals>.<listcomp>r)   r   )rM   rN   r?  r)  r   rt  r   r   r   r?   ru  r   rw   r  rv  r9   r>   r   tdnnZxvector_output_dimr@  rx  r  rw  	objectiver  )rZ   r[   rz  Ztdnn_layersr]   r   r'   rN     s    

zWavLMForXVector.__init__c                 C   s   t dt |   dS re  rJ  rL  r&   r&   r'   rM    s
    z(WavLMForXVector.freeze_feature_extractorc                 C   s   | j j  dS rN  rf  rL  r&   r&   r'   rK    s    z&WavLMForXVector.freeze_feature_encoderc                 C   s   | j  D ]
}d|_q
dS r{  r|  r   r&   r&   r'   r}    s    z!WavLMForXVector.freeze_base_model)rE   c                 C   s&   dd }| j jD ]}|||d}q|S )z?
        Computes the output length of the TDNN layers
        c                 S   s   | | | d S )Nr   r&   r6  r&   r&   r'   r7    s    zBWavLMForXVector._get_tdnn_output_lengths.<locals>._conv_out_lengthr   )r[   r  )rZ   rE   r7  rJ   r&   r&   r'   _get_tdnn_output_lengths  s    z(WavLMForXVector._get_tdnn_output_lengthsrT  rU  Nrg  c                 C   s  |dur|n| j j}| j jr dn|}| j|||||d}| j jr|t }tj|dd}tjj	| j
dd}	||	ddd jdd}n|d }| |}| jD ]}
|
|}q|du r|jdd}|jdd}n| |jdd}| |}g }g }t|D ]D\}}|||d|f jdd |||d|f jdd qt|}t|}tj||gdd}| |}| |}d}|dur| ||}|s||f|td  }|dur|f| S |S t||||j|jdS )	r~  NTrY  r   r   r)   r   )rj  rk  Z
embeddingsra   r	  )r[   r[  rt  r)  ro  r   r  r   r  r  ru  r   r6   rv  r  r  r*  r9  r  r
  rA   r   r@  rx  r  r   ra   r	  )rZ   r   r   r   r  r  rh  r   ra   r  Z
tdnn_layerZmean_featuresZstd_featuresZfeat_extract_output_lengthsZtdnn_output_lengthsr   lengthZstatistic_poolingZoutput_embeddingsrk  rj  rp  r&   r&   r'   rb     s\    



 




zWavLMForXVector.forward)NNNNN)rd   re   rf   rN   rM  rK  r}  r   r   r   r   r  r   r\  r   _XVECTOR_CHECKPOINTr   r^  _XVECTOR_EXPECTED_OUTPUTr   r   r;   r   rb   rg   r&   r&   r]   r'   r    s6   
     
r  )Nr   )Tr   r   r   typingr   r   r   numpyr3   r   Ztorch.nn.functionalr   r  r   Ztorch.utils.checkpointZtorch.nnr   Zactivationsr   r{   r	   Zmodeling_outputsr
   r   r   r   r   r   Zmodeling_utilsr   ry   r   r   r   r   Zconfiguration_wavlmr   Z
get_loggerrd   rc  ro  r^  r]  r_  rq  rr  r  r  r  r  Z#WAVLM_PRETRAINED_MODEL_ARCHIVE_LISTr   r   r   ZndarrayrG   ModulerH   rh   rm   rp   r|   r   r   r   r   r   r   r   r   r  r  r!  r"  r(  ZWAVLM_START_DOCSTRINGr\  r?  r`  rs  r  r  r  r  r&   r&   r&   r'   <module>   s    
  
x(3 ()%XYG ]%  vk