§
    ‚ŠtjmÇ  ã                   óö  — d dl mZ d dlmZ d dl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mZ dd	lmZ dd
lmZmZmZ ddlmZ ddlmZ ddlmZmZ ddlmZm Z  ddl!m"Z"m#Z# ddl$m%Z% ddl&m'Z'm(Z(m)Z)m*Z*m+Z+ ddl,m-Z-m.Z. ddl/m0Z0 ddl1m2Z2 ddl3m4Z4 ddl5m6Z6m7Z7 ddl8m9Z9  e+j:        e;¦  «        Z< e)d¬¦  «        e G d„ de'¦  «        ¦   «         ¦   «         Z= ed¦  «         G d„ d ej>        ¦  «        ¦   «         Z? G d!„ d"ej>        ¦  «        Z@ G d#„ d$ej>        ¦  «        ZAd%„ ZB ed&¦  «        dLd'„¦   «         ZCd(ejD        d)eEd*ejD        fd+„ZF	 dMd-ej>        d.ejD        d/ejD        d0ejD        d1ejD        dz  d2eGd3eGd4e%e(         fd5„ZH eeC¦  «         G d6„ d7ej>        ¦  «        ¦   «         ZI G d8„ d9e¦  «        ZJ e)d:¬¦  «        e) G d;„ d<e#¦  «        ¦   «         ¦   «         ZKe) G d=„ d>eK¦  «        ¦   «         ZL G d?„ d@ej>        ¦  «        ZM e)dA¬¦  «         G dB„ dCeKe¦  «        ¦   «         ZN G dD„ dEej>        ¦  «        ZOe) G dF„ dGeK¦  «        ¦   «         ZP e)dH¬¦  «         G dI„ dJeKe9¦  «        ¦   «         ZQg dK¢ZRdS )Né    )ÚCallable)Ú	dataclass)ÚOptionalNé   )Úinitialization)ÚACT2FN)ÚCacheÚDynamicCache)ÚGenerationMixin)Úuse_kernel_forward_from_hubÚuse_kernel_func_from_hubÚuse_kernelized_func)Úcreate_causal_mask)ÚGradientCheckpointingLayer)ÚBaseModelOutputWithPastÚCausalLMOutputWithPast)ÚROPE_INIT_FUNCTIONSÚdynamic_rope_update)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)ÚModelOutputÚTransformersKwargsÚauto_docstringÚcan_return_tupleÚlogging)Úmaybe_autocastÚmerge_with_config_defaults)Úis_torchdynamo_compiling)Úcapture_outputsé   )Ú	AutoModelé   )Ú	CsmConfigÚCsmDepthDecoderConfig)ÚCsmGenerationMixinz:
    Base class for the model autoregressive outputs.
    )Úcustom_introc                   óŠ  — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	e
dz  ed<   dZeej        df         dz  ed<   dZeej        df         dz  ed<   dZej        dz  ed	<   dZej        dz  ed
<   dZe
dz  ed<   dZeej        df         dz  ed<   dZeej        df         dz  ed<   dZej        dz  ed<   dS )ÚCsmOutputWithPasta�	  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Language modeling loss (for next-token prediction).
    logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
        Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
    past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
        It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).

        Contains pre-computed hidden-states (key and values in the self-attention blocks) that can be used (see
        `past_key_values` input) to speed up sequential decoding.
    depth_decoder_loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Language modeling loss (for next-token prediction) of the depth decoder model.
    depth_decoder_logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
        Prediction scores of the depth decoder (scores for each vocabulary token before SoftMax).
    depth_decoder_past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
        It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).
    depth_decoder_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
        one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.

        Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
    depth_decoder_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`.
    backbone_loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Language modeling loss (for next-token prediction) of the backbone model.
    NÚlossÚlogitsÚpast_key_values.Úhidden_statesÚ
attentionsÚdepth_decoder_lossÚdepth_decoder_logitsÚdepth_decoder_past_key_valuesÚdepth_decoder_hidden_statesÚdepth_decoder_attentionsÚbackbone_loss)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r*   ÚtorchÚFloatTensorÚ__annotations__r+   r,   r	   r-   Útupler.   r/   r0   r1   r2   r3   r4   © ó    úb/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/csm/modeling_csm.pyr)   r)   3   sK  € € € € € € ðð ð8 &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø'+€FˆEÔ Ñ$Ð+Ð+Ñ+Ø$(€O�U˜T‘\Ð(Ð(Ñ(Ø:>€M�5˜Ô*¨CÐ/Ô0°4Ñ7Ð>Ð>Ñ>Ø7;€J��eÔ'¨Ð,Ô-°Ñ4Ð;Ð;Ñ;Ø37Ð˜Ô)¨DÑ0Ð7Ð7Ñ7Ø59Ð˜%Ô+¨dÑ2Ð9Ð9Ñ9Ø26Ð! 5¨4¡<Ð6Ð6Ñ6ØHLÐ  uÔ'8¸#Ð'=Ô!>ÀÑ!EÐLÐLÑLØEIÐ˜e EÔ$5°sÐ$:Ô;¸dÑBÐIÐIÑIØ.2€M�5Ô$ tÑ+Ð2Ð2Ñ2Ð2Ð2r>   r)   ÚRMSNormc                   óT   ‡ — e Zd Zd	deddfˆ fd„Zdej        dej        fd„Zd„ Zˆ xZ	S )
Ú
CsmRMSNormç�íµ ÷Æ°>ÚepsÚreturnNc                 ó¬   •— t          ¦   «                              ¦   «          t          j        t	          j        |¦  «        ¦  «        | _        || _        dS )z9
        CsmRMSNorm is equivalent to T5LayerNorm
        N)ÚsuperÚ__init__ÚnnÚ	Parameterr9   ÚonesÚweightÚvariance_epsilon)ÚselfÚhidden_sizerD   Ú	__class__s      €r?   rH   zCsmRMSNorm.__init__e   sD   ø€ õ 	‰Œ×ÒÑÔÐÝ”l¥5¤:¨kÑ#:Ô#:Ñ;Ô;ˆŒØ #ˆÔÐÐr>   r-   c                 ó  — |j         }|                     t          j        ¦  «        }|                     d¦  «                             dd¬¦  «        }|t          j        || j        z   ¦  «        z  }| j        |                     |¦  «        z  S )Nr!   éÿÿÿÿT)Úkeepdim)	ÚdtypeÚtor9   Úfloat32ÚpowÚmeanÚrsqrtrM   rL   )rN   r-   Úinput_dtypeÚvariances       r?   ÚforwardzCsmRMSNorm.forwardm   s|   € Ø#Ô)ˆØ%×(Ò(­¬Ñ7Ô7ˆØ ×$Ò$ QÑ'Ô'×,Ò,¨R¸Ð,Ñ>Ô>ˆØ%­¬°H¸tÔ?TÑ4TÑ(UÔ(UÑUˆØŒ{˜]×-Ò-¨kÑ:Ô:Ñ:Ð:r>   c                 óH   — t          | j        j        ¦  «        › d| j        › �S )Nz, eps=)r<   rL   ÚshaperM   ©rN   s    r?   Ú
extra_reprzCsmRMSNorm.extra_reprt   s&   € Ý˜œÔ)Ñ*Ô*ÐIÐI°$Ô2GÐIÐIÐIr>   )rC   )
r5   r6   r7   ÚfloatrH   r9   ÚTensorr\   r`   Ú__classcell__©rP   s   @r?   rB   rB   c   sŒ   ø€ € € € € ð$ð $¨ð $¸$ð $ð $ð $ð $ð $ð $ð; U¤\ð ;°e´lð ;ð ;ð ;ð ;ðJð Jð Jð Jð Jð Jð Jr>   rB   c                   óÔ   ‡ — e Zd ZU ej        ed<   ddefˆ fd„Ze	 	 	 ddedz  de	d         de
dz  ded	ef         fd
„¦   «         Z ej        ¦   «         ed„ ¦   «         ¦   «         Zˆ xZS )ÚCsmRotaryEmbeddingÚinv_freqNÚconfigc                 ó²  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        || _        | j        j        d         | _        | j        }| j        dk    rt          | j                 } || j        |¦  «        \  }| _
        |                      d|d¬¦  «         |                      d|                     ¦   «         d¬¦  «         d S )NÚ	rope_typeÚdefaultrg   F©Ú
persistentÚoriginal_inv_freq)rG   rH   Úmax_position_embeddingsÚmax_seq_len_cachedÚoriginal_max_seq_lenrh   Úrope_parametersrj   Úcompute_default_rope_parametersr   Úattention_scalingÚregister_bufferÚclone)rN   rh   ÚdeviceÚrope_init_fnrg   rP   s        €r?   rH   zCsmRotaryEmbedding.__init__{   sÊ   ø€ Ý‰Œ×ÒÑÔÐØ"(Ô"@ˆÔØ$*Ô$BˆÔ!àˆŒàœÔ4°[ÔAˆŒØ!%Ô!EˆØŒ>˜YÒ&Ð&Ý.¨t¬~Ô>ˆLØ+7¨<¸¼ÀVÑ+LÔ+LÑ(ˆ�$Ô(à×Ò˜Z¨¸eÐÑDÔDÐDØ×ÒÐ0°(·.².Ñ2BÔ2BÈuÐÑUÔUÐUÐUÐUr>   rw   ztorch.deviceÚseq_lenrE   ztorch.Tensorc                 óü   — | j         d         }t          | dd¦  «        p| j        | j        z  }d}d|t	          j        d|dt          j        ¬¦  «                             |t          j        ¬¦  «        |z  z  z  }||fS )	a¨  
        Computes the inverse frequencies according to the original RoPE implementation
        Args:
            config ([`~transformers.PreTrainedConfig`]):
                The model configuration.
            device (`torch.device`):
                The device to use for initialization of the inverse frequencies.
            seq_len (`int`, *optional*):
                The current sequence length. Unused for this type of RoPE.
        Returns:
            Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the
            post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE).
        Ú
rope_thetaÚhead_dimNg      ð?r   r!   ©rT   ©rw   rT   )	rr   ÚgetattrrO   Únum_attention_headsr9   ÚarangeÚint64rU   ra   )rh   rw   ry   ÚbaseÚdimÚattention_factorrg   s          r?   rs   z2CsmRotaryEmbedding.compute_default_rope_parameters‹   sŒ   € ð& Ô% lÔ3ˆÝ�f˜j¨$Ñ/Ô/Ðc°6Ô3EÈÔIcÑ3cˆàÐð Ø•U”\ ! S¨!µ5´;Ð?Ñ?Ô?×BÒBÈ&ÕX]ÔXcÐBÑdÔdÐgjÑjÑkñ
ˆð Ð)Ð)Ð)r>   c                 óN  — | j         d d d …d f                              ¦   «                              |j        d         dd¦  «                             |j        ¦  «        }|d d …d d d …f                              ¦   «         }t          |j        j        t          ¦  «        r|j        j        dk    r|j        j        nd}t          |d¬¦  «        5  |                     ¦   «         |                     ¦   «         z   
                    dd¦  «        }t          j        ||fd¬	¦  «        }|                     ¦   «         | j        z  }|                     ¦   «         | j        z  }	d d d ¦  «         n# 1 swxY w Y   |                     |j        ¬
¦  «        |	                     |j        ¬
¦  «        fS )Nr   rR   r#   ÚmpsÚcpuF)Údevice_typeÚenabledr!   ©r„   r}   )rg   ra   Úexpandr^   rU   rw   Ú
isinstanceÚtypeÚstrr   Ú	transposer9   ÚcatÚcosrt   ÚsinrT   )
rN   ÚxÚposition_idsÚinv_freq_expandedÚposition_ids_expandedr‰   ÚfreqsÚembr’   r“   s
             r?   r\   zCsmRotaryEmbedding.forward©   s·  € ð !œM¨$°°°°4¨-Ô8×>Ò>Ñ@Ô@×GÒGÈÔHZÐ[\ÔH]Ð_aÐcdÑeÔe×hÒhÐijÔiqÑrÔrÐØ ,¨Q¨Q¨Q°°a°a°a¨ZÔ 8× >Ò >Ñ @Ô @Ðå'1°!´(´-ÅÑ'EÔ'EÐkÈ!Ì(Ì-Ð[`ÒJ`ÐJ`�a”h”m�mÐfkˆÝ¨¸UÐCÑCÔCð 	5ð 	5Ø&×,Ò,Ñ.Ô.Ð1F×1LÒ1LÑ1NÔ1NÑN×YÒYÐZ[Ð]^Ñ_Ô_ˆEÝ”)˜U E˜N°Ð3Ñ3Ô3ˆCØ—'’'‘)”)˜dÔ4Ñ4ˆCØ—'’'‘)”)˜dÔ4Ñ4ˆCð		5ð 	5ð 	5ñ 	5ô 	5ð 	5ð 	5ð 	5ð 	5ð 	5ð 	5øøøð 	5ð 	5ð 	5ð 	5ð �vŠv˜AœGˆvÑ$Ô$ c§f¢f°1´7 fÑ&;Ô&;Ð;Ð;s   ÃBE&Å&E*Å-E*©N©NNN)r5   r6   r7   r9   rb   r;   r$   rH   Ústaticmethodr   Úintr<   ra   rs   Úno_gradr   r\   rc   rd   s   @r?   rf   rf   x   sù   ø€ € € € € € ØŒlÐÐÑðVð V˜yð Vð Vð Vð Vð Vð Vð  à#'Ø+/Ø"ð*ð *Ø˜DÑ ð*à˜Ô(ð*ð �t‘ð*ð 
ˆ~˜uÐ$Ô	%ð	*ð *ð *ñ „\ð*ð: €U„]�_„_Øð<ð <ñ Ôñ „_ð<ð <ð <ð <ð <r>   rf   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚCsmMLPc                 ó¶  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        | _        t          j        | j        | j        |j        ¬¦  «        | _        t          j        | j        | j        |j        ¬¦  «        | _	        t          j        | j        | j        |j        ¬¦  «        | _
        t          |j                 | _        d S )N©Úbias)rG   rH   rh   rO   Úintermediate_sizerI   ÚLinearÚmlp_biasÚ	gate_projÚup_projÚ	down_projr   Ú
hidden_actÚact_fn©rN   rh   rP   s     €r?   rH   zCsmMLP.__init__º   s¯   ø€ Ý‰Œ×ÒÑÔÐØˆŒØ!Ô-ˆÔØ!'Ô!9ˆÔÝœ 4Ô#3°TÔ5KÐRXÔRaÐbÑbÔbˆŒÝ”y Ô!1°4Ô3IÐPVÔP_Ð`Ñ`Ô`ˆŒÝœ 4Ô#9¸4Ô;KÐRXÔRaÐbÑbÔbˆŒÝ˜VÔ.Ô/ˆŒˆˆr>   c                 ó¨   — |                       |                      |                      |¦  «        ¦  «        |                      |¦  «        z  ¦  «        }|S rš   )r©   r«   r§   r¨   )rN   r”   r©   s      r?   r\   zCsmMLP.forwardÄ   sA   € Ø—N’N 4§;¢;¨t¯~ª~¸aÑ/@Ô/@Ñ#AÔ#AÀDÇLÂLÐQRÁOÄOÑ#SÑTÔTˆ	ØÐr>   ©r5   r6   r7   rH   r\   rc   rd   s   @r?   r    r    ¹   sG   ø€ € € € € ð0ð 0ð 0ð 0ð 0ðð ð ð ð ð ð r>   r    c                 óœ   — | dd| j         d         dz  …f         }| d| j         d         dz  d…f         }t          j        | |fd¬¦  «        S )z*Rotates half the hidden dims of the input..NrR   r!   r‹   )r^   r9   r‘   )r”   Úx1Úx2s      r?   Úrotate_halfr²   É   s]   € à	
ˆ3Ð"�!”'˜"”+ Ñ"Ð"Ð"Ô	#€BØ	
ˆ3�”˜”˜qÑ Ð"Ð"Ð"Ô	#€BÝŒ9�r�c˜2�Y BÐ'Ñ'Ô'Ð'r>   Úrotary_pos_embc                 ó¾   — |                      |¦  «        }|                      |¦  «        }| |z  t          | ¦  «        |z  z   }||z  t          |¦  «        |z  z   }||fS )a…  Applies Rotary Position Embedding to the query and key tensors.

    Args:
        q (`torch.Tensor`): The query tensor.
        k (`torch.Tensor`): The key tensor.
        cos (`torch.Tensor`): The cosine part of the rotary embedding.
        sin (`torch.Tensor`): The sine part of the rotary embedding.
        unsqueeze_dim (`int`, *optional*, defaults to 1):
            The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and
            sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note
            that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and
            k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes
            cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have
            the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.
    Returns:
        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.
    )Ú	unsqueezer²   )ÚqÚkr’   r“   Úunsqueeze_dimÚq_embedÚk_embeds          r?   Úapply_rotary_pos_embr»   Ð   sc   € ð& �-Š-˜Ñ
&Ô
&€CØ
�-Š-˜Ñ
&Ô
&€CØ�3‰w�; q™>œ>¨CÑ/Ñ0€GØ�3‰w�; q™>œ>¨CÑ/Ñ0€GØ�GÐÐr>   r-   Ún_reprE   c                 ó¸   — | j         \  }}}}|dk    r| S | dd…dd…ddd…dd…f                              |||||¦  «        } |                      |||z  ||¦  «        S )zÔ
    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
    r#   N)r^   rŒ   Úreshape)r-   r¼   ÚbatchÚnum_key_value_headsÚslenr|   s         r?   Ú	repeat_kvrÂ   ê   s„   € ð
 2?Ô1DÑ.€EÐ  hØ�‚z€zØÐØ! ! ! ! Q Q Q¨¨a¨a¨a°°°Ð"2Ô3×:Ò:¸5ÐBUÐW\Ð^bÐdlÑmÔm€MØ× Ò  Ð(;¸eÑ(CÀTÈ8ÑTÔTÐTr>   ç        ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚscalingÚdropoutÚkwargsc                 ó  — t          || j        ¦  «        }t          || j        ¦  «        }	t          j        ||                     dd¦  «        ¦  «        |z  }
|�|
|z   }
t
          j                             |
dt          j        ¬¦  «         	                    |j
        ¦  «        }
t
          j                             |
|| j        ¬¦  «        }
t          j        |
|	¦  «        }|                     dd¦  «                             ¦   «         }||
fS )Nr!   r   rR   )r„   rT   )ÚpÚtrainingr#   )rÂ   Únum_key_value_groupsr9   Úmatmulr�   rI   Ú
functionalÚsoftmaxrV   rU   rT   rÊ   rÎ   Ú
contiguous)rÄ   rÅ   rÆ   rÇ   rÈ   rÉ   rÊ   rË   Ú
key_statesÚvalue_statesÚattn_weightsÚattn_outputs               r?   Úeager_attention_forwardrØ   ö   sé   € õ ˜3 Ô ;Ñ<Ô<€JÝ˜U FÔ$?Ñ@Ô@€Lå”<  z×';Ò';¸A¸qÑ'AÔ'AÑBÔBÀWÑL€LØÐ!Ø# nÑ4ˆå”=×(Ò(¨¸2ÅUÄ]Ð(ÑSÔS×VÒVÐW\ÔWbÑcÔc€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€LÝ”,˜|¨\Ñ:Ô:€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€Kà˜Ð$Ð$r>   c                   óÎ   ‡ — e Zd ZdZdedefˆ fd„Z	 	 	 ddej        de	ej        ej        f         dz  dej        dz  d	e
dz  d
ee         de	ej        ej        f         fd„Zˆ xZS )ÚCsmAttentionz=Multi-headed attention from 'Attention Is All You Need' paperrh   Ú	layer_idxc                 ó®  •— t          ¦   «                              ¦   «          || _        || _        t	          |d|j        |j        z  ¦  «        | _        |j        |j        z  | _	        | j        dz  | _
        |j        | _        d| _        t          j        |j        |j        | j        z  |j        ¬¦  «        | _        t          j        |j        |j        | j        z  |j        ¬¦  «        | _        t          j        |j        |j        | j        z  |j        ¬¦  «        | _        t          j        |j        | j        z  |j        |j        ¬¦  «        | _        d S )Nr|   g      à¿Tr¢   )rG   rH   rh   rÛ   r   rO   r€   r|   rÀ   rÏ   rÉ   Úattention_dropoutÚ	is_causalrI   r¥   Úattention_biasÚq_projÚk_projÚv_projÚo_proj©rN   rh   rÛ   rP   s      €r?   rH   zCsmAttention.__init__  sB  ø€ Ý‰Œ×ÒÑÔÐØˆŒØ"ˆŒÝ ¨
°FÔ4FÈ&ÔJdÑ4dÑeÔeˆŒØ$*Ô$>À&ÔB\Ñ$\ˆÔ!Ø”} dÑ*ˆŒØ!'Ô!9ˆÔØˆŒå”iØÔ Ô :¸T¼]Ñ JÐQWÔQfð
ñ 
ô 
ˆŒõ ”iØÔ Ô :¸T¼]Ñ JÐQWÔQfð
ñ 
ô 
ˆŒõ ”iØÔ Ô :¸T¼]Ñ JÐQWÔQfð
ñ 
ô 
ˆŒõ ”iØÔ&¨¬Ñ6¸Ô8JÐQWÔQfð
ñ 
ô 
ˆŒˆˆr>   Nr-   Úposition_embeddingsrÈ   r,   rË   rE   c                 ó"  — |j         d d…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }	|                      |¦  «                             |¦  «                             dd¦  «        }
|\  }}t          ||	||¦  «        \  }}	|�|                     |	|
| j	        ¦  «        \  }	}
t          j        | j        j        t          ¦  «        } || ||	|
|f| j        sdn| j        | j        dœ|¤Ž\  }} |j        g |¢d‘R Ž                      ¦   «         }|                      |¦  «        }||fS )NrR   r#   r!   rÃ   )rÊ   rÉ   )r^   r|   rà   Úviewr�   rá   râ   r»   ÚupdaterÛ   r   Úget_interfacerh   Ú_attn_implementationrØ   rÎ   rÝ   rÉ   r¾   rÓ   rã   )rN   r-   rå   rÈ   r,   rË   Úinput_shapeÚhidden_shapeÚquery_statesrÔ   rÕ   r’   r“   Úattention_interfacer×   rÖ   s                   r?   r\   zCsmAttention.forward*  sÀ  € ð $Ô)¨#¨2¨#Ô.ˆØ8˜Ð8 bÐ8¨$¬-Ð8Ð8ˆà—{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆØ—[’[ Ñ/Ô/×4Ò4°\ÑBÔB×LÒLÈQÐPQÑRÔRˆ
Ø—{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆà&‰ˆˆSÝ#7¸ÀjÐRUÐWZÑ#[Ô#[Ñ ˆ�jàÐ&Ø'6×'=Ò'=¸jÈ,ÐX\ÔXfÑ'gÔ'gÑ$ˆJ˜å(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðð %8Ð$7ØØØØØð	%
ð  $œ}ÐH�C�C°$Ô2HØ”Lð	%
ð 	%
ð ð	%
ð 	%
Ñ!ˆ�\ð *�kÔ)Ð;¨;Ð;¸Ð;Ð;Ð;×FÒFÑHÔHˆØ—k’k +Ñ.Ô.ˆØ˜LÐ(Ð(r>   r›   )r5   r6   r7   r8   r$   r�   rH   r9   rb   r<   r	   r   r   r\   rc   rd   s   @r?   rÚ   rÚ     så   ø€ € € € € àGÐGð
˜yð 
°Sð 
ð 
ð 
ð 
ð 
ð 
ð4 IMØ.2Ø(,ð&)ð &)à”|ð&)ð # 5¤<°´Ð#=Ô>ÀÑEð&)ð œ tÑ+ð	&)ð
  ™ð&)ð Ð+Ô,ð&)ð 
ˆuŒ|˜Uœ\Ð)Ô	*ð&)ð &)ð &)ð &)ð &)ð &)ð &)ð &)r>   rÚ   c                   óÒ   ‡ — e Zd Zdedefˆ fd„Z	 	 	 	 	 ddej        dej        dz  dej        dz  d	e	dz  d
e
dz  deej        ej        f         dz  dee         dej        fd„Zˆ xZS )ÚCsmDecoderLayerrh   rÛ   c                 ó4  •— t          ¦   «                              ¦   «          |j        | _        t          ||¬¦  «        | _        t          |¦  «        | _        t          |j        |j        ¬¦  «        | _	        t          |j        |j        ¬¦  «        | _
        d S )N)rh   rÛ   ©rD   )rG   rH   rO   rÚ   Ú	self_attnr    ÚmlprB   Úrms_norm_epsÚinput_layernormÚpost_attention_layernormrä   s      €r?   rH   zCsmDecoderLayer.__init__T  s�   ø€ Ý‰Œ×ÒÑÔÐØ!Ô-ˆÔå%¨V¸yÐIÑIÔIˆŒå˜&‘>”>ˆŒÝ)¨&Ô*<À&ÔBUÐVÑVÔVˆÔÝ(2°6Ô3EÈ6ÔK^Ð(_Ñ(_Ô(_ˆÔ%Ð%Ð%r>   NFr-   rÈ   r•   r,   Ú	use_cacherå   rË   rE   c           
      óÎ   — |}|                       |¦  «        } | j        d||||||dœ|¤Ž\  }}	||z   }|}|                      |¦  «        }|                      |¦  «        }||z   }|S )N)r-   rÈ   r•   r,   rø   rå   r=   )rö   ró   r÷   rô   )
rN   r-   rÈ   r•   r,   rø   rå   rË   ÚresidualÚ_s
             r?   r\   zCsmDecoderLayer.forward^  s¡   € ð !ˆØ×,Ò,¨]Ñ;Ô;ˆà)˜4œ>ð 
Ø'Ø)Ø%Ø+ØØ 3ð
ð 
ð ð
ð 
Ñˆ�qð ! =Ñ0ˆð !ˆØ×5Ò5°mÑDÔDˆØŸš Ñ/Ô/ˆØ  =Ñ0ˆØÐr>   )NNNFN)r5   r6   r7   r$   r�   rH   r9   rb   Ú
LongTensorr	   Úboolr<   r   r   r\   rc   rd   s   @r?   rð   rð   S  sÿ   ø€ € € € € ð`˜yð `°Sð `ð `ð `ð `ð `ð `ð /3Ø04Ø(,Ø!&ØHLðð à”|ðð œ tÑ+ðð Ô&¨Ñ-ð	ð
  ™ðð ˜$‘;ðð # 5¤<°´Ð#=Ô>ÀÑEðð Ð+Ô,ðð 
Œðð ð ð ð ð ð ð r>   rð   z[
    The bare Csm Model outputting raw hidden-states without any specific head on top.
    c                   ó†   ‡ — e Zd ZU eed<   dZdZdZdgZdgZ	dZ
dZdZdZeedœZ ej        ¦   «         ˆ fd„¦   «         Zˆ xZS )	ÚCsmPreTrainedModelrh   Úmodel)ÚaudioÚtextTrð   r,   )r-   r.   c                 ó°  •— t          ¦   «                              |¦  «         t          |t          ¦  «        rD|j        }t          |dz
  ¦  «        D ](}t          j        |j        d| j	        j
        ¬¦  «         Œ)d S t          |t          ¦  «        rEt          j        |j        t          j        | j	        j        ¦  «        | j	        j        z  ¦  «         d S d S )Nr#   rÃ   )rX   Ústd)rG   Ú_init_weightsr�   ÚCsmCodebooksHeadÚnum_codebooksÚrangeÚinitÚnormal_rL   rh   Úinitializer_rangeÚCsmBackboneModelEmbeddingsÚcopy_Úaudio_tokens_offsetsr9   r�   Ú
vocab_size)rN   rÄ   r  ÚirP   s       €r?   r  z CsmPreTrainedModel._init_weights—  sØ   ø€ å‰Œ×Ò˜fÑ%Ô%Ð%Ý�fÕ.Ñ/Ô/ð 	vØ"Ô0ˆMÝ˜=¨1Ñ,Ñ-Ô-ð Yð Y�Ý”˜Vœ]°¸$¼+Ô:WÐXÑXÔXÐXÐXðYð Yå˜Õ :Ñ;Ô;ð 	vÝŒJ�vÔ2µE´LÀÄÔAZÑ4[Ô4[Ð^bÔ^iÔ^tÑ4tÑuÔuÐuÐuÐuð	vð 	vr>   )r5   r6   r7   r$   r;   Úbase_model_prefixÚinput_modalitiesÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_skip_keys_device_placementÚ_supports_flash_attnÚ_supports_sdpaÚ_can_compile_fullgraphÚ_supports_attention_backendrð   rÚ   Ú_can_record_outputsr9   rž   r  rc   rd   s   @r?   rÿ   rÿ   ~  sµ   ø€ € € € € € ð ÐÐÑØÐØ(ÐØ&*Ð#Ø*Ð+ÐØ#4Ð"5ÐØÐØ€Nð "ÐØ"&Ðà(Ø"ðð Ðð
 €U„]�_„_ðvð vð vð vñ „_ðvð vð vð vð vr>   rÿ   c                   ó  ‡ — e Zd ZU eed<   ˆ fd„Zeee	 	 	 	 	 	 	 dde	j
        dz  de	j        dz  de	j        dz  de	j
        dz  dedz  d	e	j        dz  d
edz  dee         deez  fd„¦   «         ¦   «         ¦   «         Zˆ xZS )ÚCsmDepthDecoderModelrh   c                 ó.  •‡— t          ¦   «                              ‰¦  «         ‰j        | _        ‰j        | _        t          j        ‰j        ‰j        z  ‰j        ¦  «        | _	        t          j
        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        t          ‰j        ‰j        ¬¦  «        | _        t%          ‰¬¦  «        | _        d| _        t          j        ‰j        ‰j        d¬¦  «        | _        |                      ¦   «          d S )Nc                 ó0   •— g | ]}t          ‰|¦  «        ‘ŒS r=   ©rð   ©Ú.0rÛ   rh   s     €r?   ú
<listcomp>z1CsmDepthDecoderModel.__init__.<locals>.<listcomp>¬  ó#   ø€ ÐaÐaÐa°I�_˜V YÑ/Ô/ÐaÐaÐar>   rò   ©rh   Fr¢   )rG   rH   Úpad_token_idÚpadding_idxr  rI   Ú	Embeddingr  Úbackbone_hidden_sizeÚembed_tokensÚ
ModuleListr  Únum_hidden_layersÚlayersrB   rO   rõ   Únormrf   Ú
rotary_embÚgradient_checkpointingr¥   Úinputs_embeds_projectorÚ	post_initr¬   s    `€r?   rH   zCsmDepthDecoderModel.__init__¦  s÷   øø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø!Ô.ˆÔØ Ô+ˆŒÝœL¨&Ô*>ÀÔARÑ*RÐU[ÔUpÑqÔqˆÔÝ”mØaÐaÐaÐaÅÀvÔG_ÑA`ÔA`ÐaÑaÔañ
ô 
ˆŒõ ˜vÔ1°vÔ7JÐKÑKÔKˆŒ	Ý,°FÐ;Ñ;Ô;ˆŒØ&+ˆÔ#Ý')¤y°Ô1LÈfÔN`ÐglÐ'mÑ'mÔ'mˆÔ$ð 	�ŠÑÔÐÐÐr>   NÚ	input_idsÚbackbone_last_hidden_staterÈ   r•   r,   Úinputs_embedsrø   rË   rE   c           
      óØ  — |�*t          ¦   «         st                               d¦  «         d}|du |duz  rt          d¦  «        ‚|r|€t	          | j        ¬¦  «        }|�|                     ¦   «         nd}	|�|j        d         n|j        d         }
|�|j        n|j        }t          j
        |	|	|
z   |¬¦  «        }|€}t          j        |dz
  d¬¦  «        }|| j        z  }|                      ||z   ¦  «        }|d         dk    }|�
||dd…df<   n*t          ¦   «         s|rt                               d	¦  «         |                      |¦  «        }t!          | j        ||||¬
¦  «        }|}|                     d¦  «        }|                      ||¬¦  «        }| j        d| j        j        …         D ]} ||f|||||dœ|¤Ž}Œ|                      |¦  «        }t-          ||r|nd¬¦  «        S )aJ  
        backbone_last_hidden_state (`torch.FloatTensor` of shape `(batch_size, backbone_hidden_size)`, *optional*):
            The last hidden state of the backbone model. Such input is required when the first codebook token (the one generated by the backbone model)
            is provided in the `input_ids` argument.
        NzÕCustom `position_ids` were provided but will be ignored. CSM depth decoder automatically determines position_ids and as it requires them to be identical across the batch, the provided position_ids will be ignored.z;You must specify exactly one of input_ids or inputs_embeds.r$  r   r#   ©rw   )ÚminzvWhen the first codebook token is provided, `backbone_last_hidden_state` should also be provided for correct inference.©rh   r4  rÈ   r,   r•   ©r•   )rÈ   r•   r,   rø   rå   ©Úlast_hidden_stater,   )r   ÚloggerÚwarning_onceÚ
ValueErrorr
   rh   Úget_seq_lengthr^   rw   r9   r�   Úclampr  r)  Úwarningr0  r   rµ   r.  r,  r+  r-  r   )rN   r2  r3  rÈ   r•   r,   r4  rø   rË   Úpast_seen_tokensÚinputs_seq_lengthrw   Úcodebook_idxsÚoffsetÚinput_ids_are_first_codebookÚcausal_maskr-   rå   Údecoder_layers                      r?   r\   zCsmDepthDecoderModel.forward¶  s’  € ð& Ð#Õ,DÑ,FÔ,FÐ#Ý×Òðwñô ð ð  ˆLØ˜Ð -°tÐ";Ñ<ð 	\ÝÐZÑ[Ô[Ð[àð 	?˜Ð0Ý*°$´+Ð>Ñ>Ô>ˆOà?NÐ?Z˜?×9Ò9Ñ;Ô;Ð;Ð`aÐØ6CÐ6O˜MÔ/°Ô2Ð2ÐU^ÔUdÐefÔUgÐØ)6Ð)B�Ô%Ð%È	ÔHXˆÝ”|Ð$4Ð6FÐIZÑ6ZÐciÐjÑjÔjˆàÐ Ý!œK¨°qÑ(8¸aÐ@Ñ@Ô@ˆMØ" T¤_Ñ4ˆFØ ×-Ò-¨i¸&Ñ.@ÑAÔAˆMà+7¸¬?¸aÒ+?Ð(Ø)Ð5Ø&@�˜a˜a˜a ˜dÑ#Ð#å/Ñ1Ô1ð Ð6Rð Ý—N’Nð Qñô ð ð ×4Ò4°]ÑCÔCˆå(Ø”;Ø'Ø)Ø+Ø%ð
ñ 
ô 
ˆð &ˆð $×-Ò-¨aÑ0Ô0ˆØ"Ÿošo¨mÈ,˜oÑWÔWÐà!œ[Ð)H¨4¬;Ô+HÐ)HÔIð 		ð 		ˆMØ)˜MØðà*Ø)Ø /Ø#Ø$7ðð ð ðð ˆMˆMð Ÿ	š	 -Ñ0Ô0ˆÝ&Ø+Ø/8ÐB˜O˜O¸dð
ñ 
ô 
ð 	
r>   )NNNNNNN)r5   r6   r7   r%   r;   rH   r   r    r   r9   rü   r:   rb   r	   rý   r   r   r<   r   r\   rc   rd   s   @r?   r  r  ¢  s<  ø€ € € € € € à!Ð!Ð!Ñ!ðð ð ð ð ð   ØØð .2Ø?CØ.2Ø04Ø(,Ø26Ø!%ðN
ð N
àÔ# dÑ*ðN
ð %*Ô$5¸Ñ$<ðN
ð œ tÑ+ð	N
ð
 Ô&¨Ñ-ðN
ð  ™ðN
ð Ô(¨4Ñ/ðN
ð ˜$‘;ðN
ð Ð+Ô,ðN
ð 
Ð(Ñ	(ðN
ð N
ð N
ñ „^ñ „_ñ  ÔðN
ð N
ð N
ð N
ð N
r>   r  c                   ó&   ‡ — e Zd Zˆ fd„Zdd„Zˆ xZS )r  c                 óÀ   •— t          ¦   «                              ¦   «          || _        t          j        t          j        | j        dz
  ||¦  «        ¦  «        | _        d S )Nr#   )rG   rH   r  rI   rJ   r9   ÚemptyrL   )rN   rO   r  r  rP   s       €r?   rH   zCsmCodebooksHead.__init__  sM   ø€ Ý‰Œ×ÒÑÔÐØ*ˆÔÝ”l¥5¤;¨tÔ/AÀAÑ/EÀ{ÐT^Ñ#_Ô#_Ñ`Ô`ˆŒˆˆr>   Nc                 ó¨   ‡‡— |dz
  }| j         |         Šˆˆfd„t          ‰j        d         ¦  «        D ¦   «         Št          j        ‰d¬¦  «        Š‰S )Nr#   c           	      ó€   •— g | ]:}t           j                             ‰d d …|d d …f         ‰|         j        ¦  «        ‘Œ;S rš   )rI   rÑ   ÚlinearÚT)r!  Úcodebook_idxÚcodebook_weightr-   s     €€r?   r"  z,CsmCodebooksHead.forward.<locals>.<listcomp>  sX   ø€ ð 
ð 
ð 
àõ ŒM× Ò  ¨q¨q¨q°,ÀÀÀÐ/AÔ!BÀOÐT`ÔDaÔDcÑdÔdð
ð 
ð 
r>   r   r‹   )rL   r  r^   r9   Ústack)rN   r-   Úcodebook_indicesrQ  s    ` @r?   r\   zCsmCodebooksHead.forward  su   øø€ à+¨aÑ/ÐØœ+Ð&6Ô7ˆð
ð 
ð 
ð 
ð 
å % oÔ&;¸AÔ&>Ñ ?Ô ?ð
ñ 
ô 
ˆõ œ M°qÐ9Ñ9Ô9ˆàÐr>   rš   r®   rd   s   @r?   r  r  
  sQ   ø€ € € € € ðað að að að að
ð ð ð ð ð ð ð r>   r  a$  
    The CsmDepthDecoder Model transformer, with a [`CsmCodebooksHead`] on top,
    which can be seen a position-specific language modeling head, allowing to use a different linear layer for each codebook
    (e.g. position 0 is the first codebook and uses the first codebook head, etc.)
    c                   óŒ  ‡ — e Zd ZdZdZdZˆ fd„Zee	 	 	 	 	 	 	 	 	 dde	j
        dz  de	j        dz  de	j        dz  de	j
        dz  dedz  d	e	j        dz  d
e	j
        dz  dedz  dee	j        z  dee         deez  fd„¦   «         ¦   «         Z	 	 	 	 	 dde	j
        dedz  dedz  de	j
        dz  d	e	j        dz  dedz  fˆ fd„Zˆ xZS )ÚCsmDepthDecoderForCausalLMNc                 óü   •— t          ¦   «                              |¦  «         t          |¦  «        | _        |j        | _        t          |j        |j        |j        ¦  «        | _        |  	                    ¦   «          d S rš   )
rG   rH   r  r   r  r  rO   r  Úcodebooks_headr1  r¬   s     €r?   rH   z#CsmDepthDecoderForCausalLM.__init__*  sj   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý)¨&Ñ1Ô1ˆŒ
Ø Ô+ˆŒÝ.¨vÔ/AÀ6ÔCWÐY_ÔYjÑkÔkˆÔð 	�ŠÑÔÐÐÐr>   r   r2  r3  rÈ   r•   r,   r4  Úlabelsrø   Úlogits_to_keeprË   rE   c
                 ó²  — |�|                      ¦   «         nd}|�|j        d         n|j        d         }|�|j        n|j        }t          j        ||¬¦  «        |z   } | j        d	|||||||dœ|
¤Ž}|d         }t          |	t          ¦  «        r)|	dk    rt          dd¦  «        }nt          |	 d¦  «        }n|	}|  	                    |dd…|dd…f         ||         ¦  «        }| 
                    ¦   «         }d}|�:|ddd…f          
                    ¦   «         } | j        d	|d| j        j        |dœ|
¤Ž}t          |||j        |j        |j        ¬¦  «        S )
aî  
        backbone_last_hidden_state (`torch.FloatTensor` of shape `(batch_size, backbone_hidden_size)`, *optional*):
            The last hidden state of the backbone model. Such input is required when the first codebook token (the one generated by the backbone model)
            is provided in the `input_ids` argument.
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
            config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
            (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
        Nr   r#   r6  )r2  r3  rÈ   r•   r,   r4  rø   .)r+   rX  r  Úshift_labels)r*   r+   r,   r-   r.   r=   )r?  r^   rw   r9   r�   r   r�   r�   ÚslicerW  rÓ   Úloss_functionrh   r  r   r,   r-   r.   )rN   r2  r3  rÈ   r•   r,   r4  rX  rø   rY  rË   rB  ry   rw   rS  Úoutputsr-   Úslice_indicesr+   r*   r[  s                        r?   r\   z"CsmDepthDecoderForCausalLM.forward3  sÌ  € ð0 @OÐ?Z˜?×9Ò9Ñ;Ô;Ð;Ð`aÐØ,9Ð,E�-Ô% aÔ(Ð(È9Ì?Ð[\ÔK]ˆØ)6Ð)B�Ô%Ð%È	ÔHXˆÝ œ<¨¸Ð?Ñ?Ô?ÐBRÑRÐà�$”*ð 	
ØØ'AØ)Ø%Ø+Ø'Øð	
ð 	
ð ð	
ð 	
ˆð   œ
ˆå�n¥cÑ*Ô*ð 	+Ø Ò"Ð"å % a¨¡¤��å % ~ o°tÑ <Ô <��à*ˆMà×$Ò$ ]°1°1°1°mÀQÀQÀQÐ3FÔ%GÐIYÐZgÔIhÑiÔiˆØ×"Ò"Ñ$Ô$ˆàˆØÐØ! # q r r 'œ?×5Ò5Ñ7Ô7ˆLØ%�4Ô%ð Ø d°t´{Ô7MÐ\hðð Ølrðð ˆDõ &ØØØ#Ô3Ø!Ô/ØÔ)ð
ñ 
ô 
ð 	
r>   FÚnext_sequence_lengthÚis_first_iterationc                 óœ   •—  t          ¦   «         j        |||||fi |¤Ž}|s|                     d¦  «         |                     d¦  «         |S )Nr3  r•   )rG   Úprepare_inputs_for_generationÚpop)
rN   r2  r`  r,   rÈ   r4  ra  rË   Úmodel_inputsrP   s
            €r?   rc  z8CsmDepthDecoderForCausalLM.prepare_inputs_for_generationx  sq   ø€ ð =•u‘w”wÔ<ØÐ+¨_¸nÈmð
ð 
Ø_eð
ð 
ˆð "ð 	;Ø×ÒÐ9Ñ:Ô:Ð:ð 	×Ò˜Ñ(Ô(Ð(àÐr>   )	NNNNNNNNr   )NNNNF)r5   r6   r7   Ú_tied_weights_keysÚ_tp_planÚ_pp_planrH   r   r   r9   rü   r:   rb   r	   rý   r�   r   r   r<   r   r\   rc  rc   rd   s   @r?   rU  rU    së  ø€ € € € € ð ÐØ€HØ€Hðð ð ð ð ð Øð .2Ø?CØ.2Ø04Ø(,Ø26Ø*.Ø!%Ø-.ðA
ð A
àÔ# dÑ*ðA
ð %*Ô$5¸Ñ$<ðA
ð œ tÑ+ð	A
ð
 Ô&¨Ñ-ðA
ð  ™ðA
ð Ô(¨4Ñ/ðA
ð Ô  4Ñ'ðA
ð ˜$‘;ðA
ð ˜eœlÑ*ðA
ð Ð+Ô,ðA
ð 
Ð'Ñ	'ðA
ð A
ð A
ñ „^ñ ÔðA
ðL ,0Ø(,Ø26Ø26Ø*/ðð àÔ#ðð " D™jðð  ™ð	ð
 Ô(¨4Ñ/ðð Ô(¨4Ñ/ðð ! 4™Kðð ð ð ð ð ð ð ð ð r>   rU  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )r  c                 ó  •— t          ¦   «                              ¦   «          t          j        |j        |j        z  |j        ¦  «        | _        |                      dt          j
        |j        ¦  «        |j        z  d¬¦  «         d S )Nr  Frl   )rG   rH   rI   r'  r  Úcodebook_sizerO   Úembed_audio_tokensru   r9   r�   r¬   s     €r?   rH   z#CsmBackboneModelEmbeddings.__init__�  s€   ø€ Ý‰Œ×ÒÑÔÐÝ"$¤,°Ô0DÀvÔG[Ñ0[Ð^dÔ^pÑ"qÔ"qˆÔØ×ÒØ"¥E¤L°Ô1EÑ$FÔ$FÈÔI]Ñ$]Ðjoð 	ñ 	
ô 	
ð 	
ð 	
ð 	
r>   c                 ól   — |                       || j        z   ¦  «        }|                     d¬¦  «        }|S )Nr!   r‹   )rl  r  Úsum)rN   r2  r4  s      r?   r\   z"CsmBackboneModelEmbeddings.forward—  s9   € Ø×/Ò/°	¸DÔ<UÑ0UÑVÔVˆØ%×)Ò)¨aÐ)Ñ0Ô0ˆØÐr>   r®   rd   s   @r?   r  r  �  sG   ø€ € € € € ð
ð 
ð 
ð 
ð 
ðð ð ð ð ð ð r>   r  c                   óÜ   ‡ — e Zd Zˆ fd„Zeee	 	 	 	 	 	 ddej        dz  dej	        dz  dej        dz  de
dz  dej        dz  dedz  d	ee         d
efd„¦   «         ¦   «         ¦   «         Zˆ xZS )ÚCsmBackboneModelc                 ó²  •‡— t          ¦   «                              ‰¦  «         ‰j        | _        ‰j        | _        t          ‰¦  «        | _        t          j        ˆfd„t          ‰j
        ¦  «        D ¦   «         ¦  «        | _        t          ‰j        ‰j        ¬¦  «        | _        t!          ‰¬¦  «        | _        d| _        |                      ¦   «          d S )Nc                 ó0   •— g | ]}t          ‰|¦  «        ‘ŒS r=   r  r   s     €r?   r"  z-CsmBackboneModel.__init__.<locals>.<listcomp>¥  r#  r>   rò   r$  F)rG   rH   r%  r&  r  r  r)  rI   r*  r  r+  r,  rB   rO   rõ   r-  rf   r.  r/  r1  r¬   s    `€r?   rH   zCsmBackboneModel.__init__Ÿ  sÄ   øø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø!Ô.ˆÔØ Ô+ˆŒÝ6°vÑ>Ô>ˆÔÝ”mØaÐaÐaÐaÅÀvÔG_ÑA`ÔA`ÐaÑaÔañ
ô 
ˆŒõ ˜vÔ1°vÔ7JÐKÑKÔKˆŒ	Ý,°FÐ;Ñ;Ô;ˆŒØ&+ˆÔ#ð 	�ŠÑÔÐÐÐr>   Nr2  rÈ   r•   r,   r4  rø   rË   rE   c           
      óH  — |du |duz  rt          d¦  «        ‚|€|                      |¦  «        }|r|€t          | j        ¬¦  «        }|€V|�|                     ¦   «         nd}t          j        |j        d         |j        ¬¦  «        |z   }| 	                    d¦  «        }t          | j        ||||¬¦  «        }	|}
|                      |
|¬¦  «        }| j        d| j        j        …         D ]} ||
f|	||||d	œ|¤Ž}
Œ|                      |
¦  «        }
t          |
|¬
¦  «        S )a&  
        input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length, num_codebooks) or (batch_size, sequence_length)`):
            1. (batch_size, sequence_length): corresponds to the input sequence prepared with the processor from the text prompt. Such input
            requires `input_values` to be provided so that audio can be encoded in codebook tokens and then merged with the text tokens.

            2. (batch_size, sequence_length, num_codebooks): codebook tokens generated during the autoregressive decoding. Such input is not meant to be used by end users.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.

            [What are input IDs?](../glossary#input-ids)
        Nz:You must specify exactly one of input_ids or inputs_embedsr$  r   r#   r6  r8  r9  )rÈ   rå   r•   r,   rø   r:  )r>  r)  r
   rh   r?  r9   r�   r^   rw   rµ   r   r.  r,  r+  r-  r   )rN   r2  rÈ   r•   r,   r4  rø   rË   rB  rG  r-   rå   rH  s                r?   r\   zCsmBackboneModel.forward®  sŠ  € ð2 ˜Ð -°tÐ";Ñ<ð 	[ÝÐYÑZÔZÐZàÐ Ø*.×*;Ò*;¸IÑ*FÔ*FˆMàð 	?˜Ð0Ý*°$´+Ð>Ñ>Ô>ˆOàÐØCRÐC^˜×=Ò=Ñ?Ô?Ð?ÐdeÐÝ œ<¨Ô(;¸AÔ(>À}ÔG[Ð\Ñ\Ô\Ð_oÑoˆLØ'×1Ò1°!Ñ4Ô4ˆLå(Ø”;Ø'Ø)Ø+Ø%ð
ñ 
ô 
ˆð &ˆØ"Ÿošo¨mÈ,˜oÑWÔWÐà!œ[Ð)H¨4¬;Ô+HÐ)HÔIð 		ð 		ˆMØ)˜MØðà*Ø$7Ø)Ø /Ø#ðð ð ðð ˆMˆMð Ÿ	š	 -Ñ0Ô0ˆÝ&Ø+Ø+ð
ñ 
ô 
ð 	
r>   )NNNNNN)r5   r6   r7   rH   r   r    r   r9   rü   rb   r	   r:   rý   r   r   r   r\   rc   rd   s   @r?   rp  rp  �  s  ø€ € € € € ðð ð ð ð ð  ØØð .2Ø.2Ø04Ø(,Ø26Ø!%ð>
ð >
àÔ# dÑ*ð>
ð œ tÑ+ð>
ð Ô&¨Ñ-ð	>
ð
  ™ð>
ð Ô(¨4Ñ/ð>
ð ˜$‘;ð>
ð Ð+Ô,ð>
ð 
!ð>
ð >
ð >
ñ „^ñ „_ñ  Ôð>
ð >
ð >
ð >
ð >
r>   rp  zË
    The Csm model consists of two llama-like auto-regressive transformer models: a backbone model that predicts the first codebook token and a depth decoder that predicts the other codebook tokens.
    c                   ó8  ‡ — e Zd ZddiZˆ fd„Zd„ Zd„ Zeˆ fd„¦   «         Zˆ fd„Z		 	 	 	 dd	e
j        dz  d
e
j        dz  de
j        dz  de
j        dz  de
j        dz  f
d„Z	 	 	 	 dd	e
j        dedz  dedz  de
j        dz  de
j        dz  f
ˆ fd„Zee	 	 	 	 	 	 	 	 	 	 dd	e
j        dz  d
e
j        dz  de
j        dz  de
j        dz  de
j        dz  dedz  de
j        dz  de
j        dz  dedz  dee
j        z  dee         deez  fd„¦   «         ¦   «         Zˆ xZS )ÚCsmForConditionalGenerationz5backbone_model.embed_tokens.embed_audio_tokens.weightz'depth_decoder.model.embed_tokens.weightc                 óà  •— t          ¦   «                              |¦  «         |j        | _        t          j        |j        |j        d¬¦  «        | _        t          j        |j        |j        ¦  «        | _	        t                               |¦  «        | _        t                               |j        ¦  «        | _        t!          j        |j        ¦  «        | _        |                      ¦   «          d S )NFr¢   )rG   rH   r  rI   r¥   rO   Úlm_headr'  Útext_vocab_sizeÚembed_text_tokensrp  Ú_from_configÚbackbone_modelrU  Údepth_decoder_configÚdepth_decoderr"   Úfrom_configÚcodec_configÚcodec_modelr1  r¬   s     €r?   rH   z$CsmForConditionalGeneration.__init__ü  s¸   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø Ô+ˆŒÝ”y Ô!3°VÔ5FÈUÐSÑSÔSˆŒÝ!#¤¨fÔ.DÀfÔFXÑ!YÔ!YˆÔÝ.×;Ò;¸FÑCÔCˆÔÝ7×DÒDÀVÔE`ÑaÔaˆÔÝ$Ô0°Ô1DÑEÔEˆÔØ�ŠÑÔÐÐÐr>   c                 ó   — | j         j        S rš   ©r{  r)  r_   s    r?   Úget_input_embeddingsz0CsmForConditionalGeneration.get_input_embeddings  s   € ØÔ"Ô/Ð/r>   c                 ó   — || j         _        d S rš   r‚  )rN   rÇ   s     r?   Úset_input_embeddingsz0CsmForConditionalGeneration.set_input_embeddings	  s   € Ø+0ˆÔÔ(Ð(Ð(r>   c                 óÖ  •‡‡— |                      dd¦  «        r t          ¦   «         j        |i |¤Ž\  }}n t          ¦   «         j        |i |¤Ž}dŠt          ‰¦  «        Šˆˆfd„t	          |j        ¦  «                             ¦   «         D ¦   «         }t	          |j        j        ¦  «                             ddi|¥¦  «         |D ]}t          |j        ‰|z   ¦  «         Œd|v r||fS |S )NÚoutput_loading_infoFÚdepth_decoder_c                 óV   •— i | ]%\  }}|                      ‰¦  «        ¯|‰d …         |“Œ&S rš   )Ú
startswith)r!  ÚattrrÇ   ÚprefixÚ
prefix_lens      €€r?   ú
<dictcomp>z?CsmForConditionalGeneration.from_pretrained.<locals>.<dictcomp>  sJ   ø€ ð 
ð 
ð 
á��eØ�Š˜vÑ&Ô&ð
Ø���Ô˜uð
ð 
ð 
r>   Ú_from_model_config)
ÚgetrG   Úfrom_pretrainedÚlenÚvarsÚgeneration_configÚitemsr}  rè   Údelattr)
ÚclsÚargsrË   r   Úloading_infoÚdepth_decoder_attrsr‹  rŒ  r�  rP   s
          @@€r?   r‘  z+CsmForConditionalGeneration.from_pretrained  s'  øøø€ à�:Š:Ð+¨UÑ3Ô3ð 	=Ø"9¥%¡'¤'Ô"9¸4Ð"JÀ6Ð"JÐ"JÑˆE�<�<à+•E‘G”GÔ+¨TÐ<°VÐ<Ð<ˆEð "ˆÝ˜‘[”[ˆ
ð
ð 
ð 
ð 
ð 
å# EÔ$;Ñ<Ô<×BÒBÑDÔDð
ñ 
ô 
Ðõ 	ˆUÔ Ô2Ñ3Ô3×:Ò:Ð<PÐRWÐ;oÐ[nÐ;oÑpÔpÐpð (ð 	<ð 	<ˆDÝ�EÔ+¨V°d©]Ñ;Ô;Ð;Ð;à  FÐ*Ð*Ø˜,Ð&Ð&àˆLr>   c                 ó  •— d}| j         j                             ¦   «         }|                     dd ¦  «         |                     ¦   «         D ]\  }}t          | j        ||z   |¦  «         Œ t          ¦   «         j        |i |¤Ž d S )Nrˆ  Útransformers_version)r}  r”  Úto_diff_dictrd  r•  ÚsetattrrG   Úsave_pretrained)rN   r˜  rË   rŒ  rš  r‹  rÇ   rP   s          €r?   rŸ  z+CsmForConditionalGeneration.save_pretrained'  s–   ø€ à!ˆØ"Ô0ÔB×OÒOÑQÔQÐØ×ÒÐ 6¸Ñ=Ô=Ð=Ø.×4Ò4Ñ6Ô6ð 	Bð 	B‰KˆD�%Ý�DÔ*¨F°T©M¸5ÑAÔAÐAÐAà�‰ŒÔ Ð0¨Ð0Ð0Ð0Ð0Ð0r>   Nr2  Úinput_valuesÚinput_values_cutoffsrX  rE   c                 óÖ  ‡— |                       |¦  «        }|��Lt          j                             |d¦  «        }||dk                                  ¦   «         }||dk             }t          j        |                     ¦   «         |j        ¬¦  «         	                    t          |¦  «        d¦  «        }||                     d¦  «        k     }t          j        ¦   «         5  g }t          ||¦  «        D ]³\  }	}
|
|
dk             }
t          |
j        d         dz
  ¦  «        D ]„}|
|         }|
|dz            }|	d||…f         }| j                             |                     d¦  «        ¦  «        }|j                             dd¦  «        }|                     |d         ¦  «         Œ…Œ´t          d„ |D ¦   «         ¦  «        Št          j        ˆfd	„|D ¦   «         ¦  «        }| j                             |¦  «        }ddd¦  «         n# 1 swxY w Y   | j        j        }||k    }| j                             |¦  «        }||         ||<   t          j        dd| j        j        f|j        t
          j        ¬
¦  «        | j        j        z  }| j                             |¦  «                             d¦  «        }|| j        j         k    }| !                    | "                    ¦   «         d¦  «        ||<   |�v|                     d¦  «         !                    dd| j        j        ¦  «        }||         ||<   |||<   |dk     #                    d¬¦  «        }d||d         |d         dd…f<   |}||dœS )a˜  
        Merges the input_ids and input_values to produce a single inputs_embeds tensor:
        1 - Infers the codec model on the input_values to retrieve codebook token.
        2 - Embeds codebook tokens and places them at the correct positions in the inputs_embeds tensor.
        3 - If labels are provided, expands them to match codebook dimensions and position the target codebook tokens in the inputs_embeds tensor.

        Args:
            input_ids (`torch.Tensor` of shape `(batch_size, sequence_length)`):
                The input ids to embed.
            input_values (`torch.Tensor` of shape `(batch_size, channels, audio_sequence_length)`):
                The audio input values to embed.
            input_values_cutoffs (`torch.Tensor` of shape `(batch_size, max_num_audio)`):
                The cutoffs of the audio input values relative to its batch index, padded with -1 when no audio.
        N©r#   r   r   r6  rR   r#   .c              3   ó0   K  — | ]}|j         d          V — ŒdS )r   N)r^   )r!  Úels     r?   ú	<genexpr>zQCsmForConditionalGeneration._merge_input_ids_with_input_values.<locals>.<genexpr>a  s(   è è € Ð&OÐ&O°r r¤x°¤{Ð&OÐ&OÐ&OÐ&OÐ&OÐ&Or>   c                 ót   •— g | ]4}t           j                             |d d d ‰|j        d          z
  f¦  «        ‘Œ5S )r   )rI   rÑ   Úpadr^   )r!  r¥  Úmax_audio_framess     €r?   r"  zRCsmForConditionalGeneration._merge_input_ids_with_input_values.<locals>.<listcomp>c  sB   ø€ ÐrÐrÐrÐZ\•R”]×&Ò& r¨A¨q°!Ð5EÈÌÐQRÌÑ5SÐ+TÑUÔUÐrÐrÐrr>   r~   i›ÿÿÿT©Úas_tupleéœÿÿÿ)r4  rX  )$ry  rI   rÑ   r¨  Údiffr9   r�   Úmaxrw   rŒ   r’  rµ   rž   Úzipr  r^   r€  ÚencodeÚaudio_codesr�   ÚappendrR  Úget_audio_codes_maskrh   Úaudio_token_idr{  r)  rK   r  ÚlongÚcodebook_eos_token_idÚsqueezeÚaudio_eos_token_idÚrepeatrn  Únonzero)rN   r2  r   r¡  rX  r4  Úaudio_lengthsÚinput_values_maskÚaudio_tokens_listÚbatch_input_valuesÚbatch_input_values_cutoffsr  Ú	start_idxÚend_idxÚaudio_batchÚcodec_outputsÚcodebook_idsÚbatched_audio_token_idsÚaudio_codes_maskr´  Úaudio_token_maskÚaudio_embedsÚaudio_eos_frame_idsÚaudio_eos_embedsÚaudio_eos_token_maskÚlabels_expandedÚ depth_decoder_ignore_frames_idxsr©  s                              @r?   Ú"_merge_input_ids_with_input_valuesz>CsmForConditionalGeneration._merge_input_ids_with_input_values1  s  ø€ ð* ×.Ò.¨yÑ9Ô9ˆàÑ#å#%¤=×#4Ò#4Ð5IÈ6Ñ#RÔ#RÐ Ø0Ð1EÈÒ1JÔK×PÒPÑRÔRˆMØ)¨-¸!Ò*;Ô<ˆMÝ %¤Ð-A×-EÒ-EÑ-GÔ-GÐP\ÔPcÐ dÑ dÔ d× kÒ kÝ�MÑ"Ô" Bñ!ô !Ðð !2°M×4KÒ4KÈAÑ4NÔ4NÒ NÐõ
 ”‘”ð \ð \Ø$&Ð!ÝFIÈ,ÐXlÑFmÔFmð Bð BÑBÐ&Ð(BØ1KÐLfÐjkÒLkÔ1lÐ.Ý"Ð#=Ô#CÀAÔ#FÈÑ#JÑKÔKð Bð B˜Ø$>¸qÔ$A˜	Ø"<¸QÀ¹UÔ"C˜Ø&8¸¸iÈÐ>OÐ9OÔ&P˜Ø(,Ô(8×(?Ò(?À×@UÒ@UÐVWÑ@XÔ@XÑ(YÔ(Y˜Ø'4Ô'@×'JÒ'JÈ1ÈbÑ'QÔ'Q˜Ø)×0Ò0°¸a´ÑAÔAÐAÐAðBõ $'Ð&OÐ&OÐ=NÐ&OÑ&OÔ&OÑ#OÔ#OÐ Ý*/¬+ØrÐrÐrÐrÐ`qÐrÑrÔrñ+ô +Ð'ð $(Ô#3×#HÒ#HÐIZÑ#[Ô#[Ð ð!\ð \ð \ñ \ô \ð \ð \ð \ð \ð \ð \øøøð \ð \ð \ð \ð$ "œ[Ô7ˆNØ(¨NÒ:ÐàÔ.×;Ò;Ð<SÑTÔTˆLØ.:Ð;KÔ.LˆMÐ*Ñ+õ ”
˜A˜q $¤+Ô";Ð<ÀYÔEUÕ]bÔ]gÐhÑhÔhØ”+Ô3ñ4ð  ð  $Ô2×?Ò?Ð@SÑTÔT×\Ò\Ð]^Ñ_Ô_Ðà#,°´Ô0NÒ#NÐ Ø2B×2IÒ2IÐJ^×JbÒJbÑJdÔJdÐfgÑ2hÔ2hˆMÐ.Ñ/ð Ð!Ø"(×"2Ò"2°2Ñ"6Ô"6×"=Ò"=¸aÀÀDÄKÔD]Ñ"^Ô"^�Ø4KÐL\Ô4]�Ð 0Ñ1Ø8K�Ð 4Ñ5à4:¸d²N×3KÒ3KÐUYÐ3KÑ3ZÔ3ZÐ0Øpt�Ð @ÀÔ CÐEeÐfgÔEhÐjkÐjlÐjlÐ lÑmØ(�à!.¸&ÐAÐAÐAs   ÃDHÈHÈHr`  r,   rÈ   r4  c           	      óx  •—  t          ¦   «         j        d	|||||dœ|¤Ž}|�—|j        dk    rŒ|                     d¦  «        €w|                      ||                     d¦  «        |                     d¦  «        |                     d¦  «        ¬¦  «        }|                     |d         |d         d dœ¦  «         |S )
N)r2  r`  r,   rÈ   r4  r!   r4  r   r¡  rX  )r2  r   r¡  rX  )r4  rX  r2  r=   )rG   rc  Úndimr�  rÎ  rè   )
rN   r2  r`  r,   rÈ   r4  rË   re  Úmerged_inputsrP   s
            €r?   rc  z9CsmForConditionalGeneration.prepare_inputs_for_generationƒ  sî   ø€ ð =•u‘w”wÔ<ð 
ØØ!5Ø+Ø)Ø'ð
ð 
ð ð
ð 
ˆð Ð  Y¤^°qÒ%8Ð%8¸\×=MÒ=MÈoÑ=^Ô=^Ð=fØ ×CÒCØ#Ø#ŸZšZ¨Ñ7Ô7Ø%+§Z¢ZÐ0FÑ%GÔ%GØ—z’z (Ñ+Ô+ð	 Dñ ô ˆMð ×ÒØ"/°Ô"@ÈMÐZbÔLcÐrvÐwÐwñô ð ð Ðr>   r   r•   rø   rY  rË   c                 óÆ  — |�5|j         dk    r*|                      ||||¦  «        }|d         }|d         }d} | j        d||||||	dœ|¤Ž}|d         }t          |
t          ¦  «        rt          |
 d¦  «        n|
}|                      |dd…|dd…f         ¦  «        }d}d}d}d}|�î|dd…dd…df         } | j        d||| j        j	        dœ|¤Ž}|dd…dd…dd…f         d	k     
                    d
¬¦  «         }||         dd| j        j        dz
  …f         }t          j                             |dd¬¦  «        }|                     d¬¦  «        }||d         |d         dz
  dd…f         }||         } | j        d|||	d|dœ|¤Ž}|j        }||z   }t%          |||||j        |j        |j        |�|j        nd|�|j        nd|�|j        nd|�|j        nd¬¦  «        S )a  
        input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length, num_codebooks) or (batch_size, sequence_length)`):
            1. (batch_size, sequence_length): corresponds to the input sequence prepared with the processor from the text prompt. Such input
            requires `input_values` to be provided so that audio can be encoded in codebook tokens and then merged with the text tokens.

            2. (batch_size, sequence_length, num_codebooks): codebook tokens generated during the autoregressive decoding. Such input is not meant to be used by end users.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.

            [What are input IDs?](../glossary#input-ids)
        input_values_cutoffs (`torch.Tensor` of shape `(batch_size, max_num_audio)`, *optional*):
            Specify the end positions of audio segments within each batch entry, relative to the concatenated audio input.
            If a batch entry has fewer segments than the maximum, it is padded with -1. For example, in a batch of 2 sequences
            where the first contains 2 audio segments of length l1, and the second contains 1 audio segment of length l2,
            the input_values_cutoffs would be: [[l1, 2 * l1], [l2, -1]].
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should be in `[config.audio_token_id, -100, -101]`.
            Requires targeted `input_values` to be provided as audio tokens will be inferred from it using the `codec_model`.
            - `config.audio_token_id` indicates an audio frames (considering sequence length elements as frames)
            - `-100` will be ignored in the loss computation
            - `-101` indicates the audio frame will be used only for the backbone model (using the first codebook token as labels)

            Such labels can be prepared using `output_labels=True` when calling [`CsmProcessor`].
        logits_to_keep (`int` or `torch.Tensor`, *optional*):
            Kept for compatibility. Does not support another value than:
            1. `0`, which is equivalent to keeping all logits, used in the training regime
            2. `1`, which is equivalent to keeping only the last logit, used in the generation regime

        Example:

        ```python
        >>> import torch
        >>> from transformers import CsmForConditionalGeneration, AutoProcessor
        >>> from datasets import load_dataset, Audio

        >>> model_id = "sesame/csm-1b"
        >>> torch_device = "cuda" if torch.cuda.is_available() else "cpu"

        >>> processor = AutoProcessor.from_pretrained(model_id)

        >>> ds = load_dataset("hf-internal-testing/dailytalk-dummy", split="train")
        >>> # ensure the audio is 24kHz
        >>> ds = ds.cast_column("audio", Audio(sampling_rate=24000))

        >>> conversation = []
        >>> # prepare a conversation with text and corresponding audio
        >>> for text, audio, speaker_id in zip(ds[:4]["text"], ds[:4]["audio"], ds[:4]["speaker_id"]):
        ...     conversation.append(
        ...         {
        ...             "role": f"{speaker_id}",
        ...             "content": [{"type": "text", "text": text}, {"type": "audio", "path": audio["array"]}],
        ...         }
        ...     )

        >>> inputs = processor.apply_chat_template(
        ...     conversation,
        ...     tokenize=True,
        ...     return_dict=True,
        ...     output_labels=True,
        ... ).to(torch_device)

        >>> model = CsmForConditionalGeneration.from_pretrained(model_id, device_map=torch_device)
        >>> output = model(**inputs)
        >>> output.loss.backward()
        ```Nr!   r4  rX  )r2  rÈ   r•   r,   r4  rø   r   )r+   rX  r  r#   r¬  rR   r‹   .r£  )rÇ   Trª  )r2  r3  rø   Úreturn_dictrX  )r*   r4   r/   r+   r,   r-   r.   r0   r1   r2   r3   r=   )rÐ  rÎ  r{  r�   r�   r\  rw  r]  rh   r  Úallr  rI   rÑ   r¨  rº  r}  r*   r)   r,   r-   r.   r+   )rN   r2  r   rÈ   r¡  r•   r,   r4  rX  rø   rY  rË   rÑ  Úbackbone_outputsÚbackbone_hidden_statesr_  Úbackbone_logitsr*   r4   r/   Údepth_decoder_outputsÚbackbone_labelsÚ
train_maskÚdepth_decoder_input_idsÚ
train_idxsÚbackbone_last_hidden_statesÚdepth_decoder_labelss                              r?   r\   z#CsmForConditionalGeneration.forward¢  sæ  € ðd Ð  Y¤^°qÒ%8Ð%8Ø ×CÒCØ˜<Ð)=¸vñô ˆMð *¨/Ô:ˆMØ" 8Ô,ˆFØˆIà.˜4Ô.ð 
ØØ)Ø%Ø+Ø'Øð
ð 
ð ð
ð 
Ðð "2°!Ô!4Ðå8BÀ>ÕSVÑ8WÔ8WÐk�˜~˜o¨tÑ4Ô4Ð4Ð]kˆØŸ,š,Ð'=¸a¸a¸aÀÐPQÐPQÐPQÐ>QÔ'RÑSÔSˆàˆØˆØ!ÐØ $ÐØÐà$ Q Q Q¨¨¨¨1 WœoˆOØ.˜DÔ.ð Ø&¨È4Ì;ÔKaðð Øekðð ˆMð " ! ! ! Q Q Q¨¨¨ (Ô+¨tÒ3×8Ò8¸RÐ8Ñ@Ô@Ð@ˆJØ&,¨ZÔ&8¸Ð>]ÀÄÔ@YÐ\]Ñ@]Ð>]Ð9]Ô&^Ð#å&(¤m×&7Ò&7Ð8OÐQWÐ_`Ð&7Ñ&aÔ&aÐ#à#×+Ò+°TÐ+Ñ:Ô:ˆJØ*@ÀÈAÄÐPZÐ[\ÔP]Ð`aÑPaÐcdÐcdÐcdÐAdÔ*eÐ'Ø#)¨*Ô#5Ð à$6 DÔ$6ð %Ø1Ø+FØ#Ø Ø+ð%ð %ð ð%ð %Ð!ð "7Ô!;ÐØ Ð#5Ñ5ˆDå ØØ'Ø1Ø"Ø,Ô<Ø*Ô8Ø'Ô2ØAVÐAbÐ!6Ô!=Ð!=Ðhlà$Ð0ð +@Ô*OÐ*Oàà$Ð0ð )>Ô(KÐ(KàØI^ÐIjÐ%:Ô%EÐ%EÐptð
ñ 
ô 
ð 	
r>   )NNNN)
NNNNNNNNNr   )r5   r6   r7   rf  rH   rƒ  r…  Úclassmethodr‘  rŸ  r9   rb   rÎ  rü   r�   r	   r:   rc  r   r   rý   r   r   r<   r)   r\   rc   rd   s   @r?   ru  ru  ò  sÛ  ø€ € € € € ð 	@ÐAjðÐðð ð ð ð ð0ð 0ð 0ð1ð 1ð 1ð ðð ð ð ñ „[ðð41ð 1ð 1ð 1ð 1ð *.Ø,0Ø48Ø&*ðPBð PBà”< $Ñ&ðPBð ”l TÑ)ðPBð $œl¨TÑ1ð	PBð
 ”˜tÑ#ðPBð 
Œ˜Ñ	ðPBð PBð PBð PBðj ,0Ø(,Ø26Ø26ðð àÔ#ðð " D™jðð  ™ð	ð
 Ô(¨4Ñ/ðð Ô(¨4Ñ/ðð ð ð ð ð ð> Øð .2Ø,0Ø.2Ø48Ø04Ø(,Ø26Ø*.Ø!%Ø-.ðY
ð Y
àÔ# dÑ*ðY
ð ”l TÑ)ðY
ð œ tÑ+ð	Y
ð
 $œl¨TÑ1ðY
ð Ô&¨Ñ-ðY
ð  ™ðY
ð Ô(¨4Ñ/ðY
ð Ô  4Ñ'ðY
ð ˜$‘;ðY
ð ˜eœlÑ*ðY
ð Ð+Ô,ðY
ð 
Ð"Ñ	"ðY
ð Y
ð Y
ñ „^ñ ÔðY
ð Y
ð Y
ð Y
ð Y
r>   ru  )rÿ   rp  r  rU  ru  )r#   )rÃ   )SÚcollections.abcr   Údataclassesr   Útypingr   r9   Útorch.nnrI   Ú r   r	  Úactivationsr   Úcache_utilsr	   r
   Ú
generationr   Úintegrationsr   r   r   Úmasking_utilsr   Úmodeling_layersr   Úmodeling_outputsr   r   Úmodeling_rope_utilsr   r   Úmodeling_utilsr   r   Úprocessing_utilsr   Úutilsr   r   r   r   r   Úutils.genericr   r   Úutils.import_utilsr   Úutils.output_capturingr    Úautor"   Úconfiguration_csmr$   r%   Úgeneration_csmr&   Ú
get_loggerr5   r<  r)   ÚModulerB   rf   r    r²   r»   rb   r�   rÂ   ra   rØ   rÚ   rð   rÿ   r  r  rU  r  rp  ru  Ú__all__r=   r>   r?   ú<module>rù     si  ðð* %Ð $Ð $Ð $Ð $Ð $Ø !Ð !Ð !Ð !Ð !Ð !Ø Ð Ð Ð Ð Ð à €€€Ø Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &Ø !Ð !Ð !Ð !Ð !Ð !Ø .Ð .Ð .Ð .Ð .Ð .Ð .Ð .Ø )Ð )Ð )Ð )Ð )Ð )Ø fÐ fÐ fÐ fÐ fÐ fÐ fÐ fÐ fÐ fØ /Ð /Ð /Ð /Ð /Ð /Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø OÐ OÐ OÐ OÐ OÐ OÐ OÐ OØ KÐ KÐ KÐ KÐ KÐ KÐ KÐ KØ FÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &Ø _Ð _Ð _Ð _Ð _Ð _Ð _Ð _Ð _Ð _Ð _Ð _Ð _Ð _Ø GÐ GÐ GÐ GÐ GÐ GÐ GÐ GØ :Ð :Ð :Ð :Ð :Ð :Ø 5Ð 5Ð 5Ð 5Ð 5Ð 5Ø Ð Ð Ð Ð Ð Ø ?Ð ?Ð ?Ð ?Ð ?Ð ?Ð ?Ð ?Ø .Ð .Ð .Ð .Ð .Ð .ð 
ˆÔ	˜HÑ	%Ô	%€ð €ððñ ô ð
 ð'3ð '3ð '3ð '3ð '3˜ñ '3ô '3ñ „ñô ð'3ðT Ð˜YÑ'Ô'ðJð Jð Jð Jð J�”ñ Jô Jñ (Ô'ðJð(><ð ><ð ><ð ><ð ><˜œñ ><ô ><ð ><ðBð ð ð ð ˆRŒYñ ô ð ð (ð (ð (ð ÐÐ*Ñ+Ô+ðð ð ñ ,Ô+ðð2	U˜Uœ\ð 	U°#ð 	U¸%¼,ð 	Uð 	Uð 	Uð 	Uð& ð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð ð%ð ð%ð Ð'Ô(ð%ð %ð %ð %ð2 ÐÐ)Ñ*Ô*ð@)ð @)ð @)ð @)ð @)�2”9ñ @)ô @)ñ +Ô*ð@)ðF(ð (ð (ð (ð (Ð0ñ (ô (ð (ðV €ððñ ô ð
 ðvð vð vð vð v˜ñ vô vñ „ñô ðvð< ðd
ð d
ð d
ð d
ð d
Ð-ñ d
ô d
ñ „ðd
ðNð ð ð ð �r”yñ ô ð ð( €ððñ ô ðgð gð gð gð gÐ!3°_ñ gô gñô ðgðTð ð ð ð  ¤ñ ô ð ð ðQ
ð Q
ð Q
ð Q
ð Q
Ð)ñ Q
ô Q
ñ „ðQ
ðh €ððñ ô ð
F
ð F
ð F
ð F
ð F
Ð"4Ð6Hñ F
ô F
ñô ð
F
ðR
ð ð €€€r>   