§
    ‚ŠtjA¢  ã            
       ó¤  — d Z ddlZddlZddlmZ ddlmZmZ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 ddlmZmZmZmZmZ ddlmZ ddlmZmZ ddl m!Z!  ej"        e#¦  «        Z$dej%        de&dej'        dej%        fd„Z(dej%        dej%        de)de*dej%        f
d„Z+dej%        dej%        fd„Z,dej%        dej%        dej%        fd„Z- G d„ dej.        j/        ¦  «        Z0 G d„ d ej1        ¦  «        Z2 G d!„ d"ej1        ¦  «        Z3 G d#„ d$ej1        ¦  «        Z4 G d%„ d&e¦  «        Z5e G d'„ d(e¦  «        ¦   «         Z6e G d)„ d*e6¦  «        ¦   «         Z7 ed+¬,¦  «         G d-„ d.e6e¦  «        ¦   «         Z8 ed/¬,¦  «         G d0„ d1e6¦  «        ¦   «         Z9e G d2„ d3e6¦  «        ¦   «         Z:e G d4„ d5e6¦  «        ¦   «         Z;g d6¢Z<dS )7zPyTorch BLOOM model.é    N)Únn)ÚBCEWithLogitsLossÚCrossEntropyLossÚ	LayerNormÚMSELoss)Ú
functionalé   )ÚCacheÚDynamicCacheÚStaticCache)ÚGenerationMixin)Úcreate_causal_mask)ÚGradientCheckpointingLayer)Ú)BaseModelOutputWithPastAndCrossAttentionsÚ!CausalLMOutputWithCrossAttentionsÚQuestionAnsweringModelOutputÚ SequenceClassifierOutputWithPastÚTokenClassifierOutput)ÚPreTrainedModel)Úauto_docstringÚloggingé   )ÚBloomConfigÚattention_maskÚ	num_headsÚdtypeÚreturnc                 óž  — | j         \  }}dt          j        t          j        |¦  «        ¦  «        z  }t	          j        ddt          j        |¦  «        dz
   z   z  | j        t          j        ¬¦  «        }t	          j        dd|z   | j        t          j	        ¬¦  «        }t	          j
        ||¦  «        }||k    r²t	          j        ddt          j        d|z  ¦  «        dz
   z   z  | j        t          j        ¬¦  «        }	t          |||z
  ¦  «        }
t	          j        ddd|
z  z   d| j        t          j	        ¬¦  «        }t	          j        |t	          j
        |	|¦  «        gd¬¦  «        }|                      d¬¦  «        dz
  | z  dd…ddd…f         }|d	         |z  }|                     ||z  d|¦  «                             |¦  «        S )
aŽ  
    Link to paper: https://huggingface.co/papers/2108.12409 Alibi tensor is not causal as the original paper mentions, it
    relies on a translation invariance of softmax for quick implementation: with l being a tensor, and a fixed value
    `softmax(l+a) = softmax(l)`. Based on
    https://github.com/ofirpress/attention_with_linear_biases/blob/a35aaca144e0eb6b789dfcb46784c4b8e31b7983/fairseq/models/transformer.py#L742
    TODO @thomasw21 this doesn't work as nicely due to the masking strategy, and so masking varies slightly.

    Args:
    Returns tensor shaped (batch_size * num_heads, 1, max_seq_len)
        attention_mask (`torch.Tensor`):
            Token-wise attention mask, this should be of shape (batch_size, max_seq_len).
        num_heads (`int`):
            number of heads
        dtype (`torch.dtype`, *optional*, default=`torch.bfloat16`):
            dtype of the output tensor
    é   r	   ©Údevicer   r   r   ©ÚdiméÿÿÿÿN).N)ÚshapeÚmathÚfloorÚlog2ÚtorchÚtensorr!   Úfloat32ÚarangeÚint32ÚpowÚminÚcatÚcumsumÚreshapeÚto)r   r   r   Ú
batch_sizeÚ
seq_lengthÚclosest_power_of_2ÚbaseÚpowersÚslopesÚ
extra_baseÚnum_remaining_headsÚextra_powersÚarange_tensorÚalibis                 úf/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/bloom/modeling_bloom.pyÚbuild_alibi_tensorr@   -   sÚ  € ð" ,Ô1Ñ€J�
Ø�dœj­¬°9Ñ)=Ô)=Ñ>Ô>Ñ>ÐÝŒ<Ø	�•t”yÐ!3Ñ4Ô4°qÑ8Ð9Ñ9Ð:Ñ;ÀNÔDYÕafÔanðñ ô €Dõ Œ\˜!˜QÐ!3Ñ3¸NÔ<QÕY^ÔYdÐeÑeÔe€FÝŒY�t˜VÑ$Ô$€Fà˜YÒ&Ð&Ý”\Ø�A�4œ9 QÐ);Ñ%;Ñ<Ô<¸qÑ@ÐAÑAÐBÑCÈNÔLaÕinÔivð
ñ 
ô 
ˆ
õ "Ð"4°iÐBTÑ6TÑUÔUÐÝ”| A q¨1Ð/BÑ+BÑ'BÀAÈnÔNcÕkpÔkvÐwÑwÔwˆÝ”˜F¥E¤I¨j¸,Ñ$GÔ$GÐHÈaÐPÑPÔPˆð %×+Ò+°Ð+Ñ3Ô3°aÑ7¸>ÑIÈ1È1È1ÈdÐTUÐTUÐTUÈ:ÔV€MØ�9Ô Ñ-€EØ�=Š=˜ iÑ/°°JÑ?Ô?×BÒBÀ5ÑIÔIÐIó    ÚxÚresidualÚprobÚtrainingc                 ó>   — t          j        | ||¬¦  «        }||z   }|S )a
  
    Dropout add function

    Args:
        x (`torch.tensor`):
            input tensor
        residual (`torch.tensor`):
            residual tensor
        prob (`float`):
            dropout probability
        training (`bool`):
            training mode
    )ÚprE   )ÚFÚdropout)rB   rC   rD   rE   Úouts        r?   Údropout_addrK   Y   s(   € õ Œ)�A˜¨Ð
1Ñ
1Ô
1€CØ
�S‰.€CØ€JrA   c                 óZ   — | dz  dt          j        d| z  dd| z  | z  z   z  ¦  «        z   z  S )zà
    Custom bias GELU function. Adapted from Megatron-DeepSpeed code. Here we use a simple implementation (inference) to
    make the model jitable.

    Args:
        x (`torch.tensor`):
            input hidden states
    ç      à?ç      ð?ç Þe3Eˆé?r   ç÷Hmâä¦?©r)   Útanh)rB   s    r?   Úbloom_gelu_forwardrS   l   s9   € ð ˆs‰7�c�EœJ z°A¡~¸¸XÈ¹\ÈAÑ=MÑ9MÑ'NÑOÔOÑOÑPÐPrA   Úgc                 ó¨   — |d         }t          j        d|z  dd|z  |z  z   z  ¦  «        }d|z  d||z  z
  dd|z  |z  z   z  z  dd|z   z  z   }|| z  S )a   
    gradient of tanh approximation of gelu gradient of actual gelu is: 0.5 * (1. + torch.erf(x * 0.70710678)) +
    0.3989423 * x * torch.exp(-0.5 * x * x)

    Args:
        g (`torch.tensor`):
            gradient output tensor
        x (`torch.tensor`):
            input tensor
    r   rO   r   rP   rM   g6”ü¾vf»?rQ   )rT   rB   Útanh_outÚffs       r?   Úbloom_gelu_backrX   x   s{   € ð 	
ˆ!Œ€AÝŒz˜* q™.¨A°¸1±¸qÑ0@Ñ,@ÑAÑBÔB€Hà	ˆq‰�Q˜ HÑ,Ñ,°¸lÈQÑ>NÐQRÑ>RÑ1RÑSÑ	TÐWZÐ^_ÐbjÑ^jÑWkÑ	k€BØ�‰6€MrA   c                   óv   — e Zd Zedej        dej        fd„¦   «         Zedej        dej        fd„¦   «         ZdS )ÚGeLUFunctionÚinputr   c                 óJ   — |                       |¦  «         t          |¦  «        S ©N)Úsave_for_backwardrS   )Úctxr[   s     r?   ÚforwardzGeLUFunction.forward‹   s$   € à×Ò˜eÑ$Ô$Ð$Ý! %Ñ(Ô(Ð(rA   Úgrad_outputc                 ó4   — | j         }t          ||¦  «        }|S r]   )Úsaved_tensorsrX   )r_   ra   r[   Útmps       r?   ÚbackwardzGeLUFunction.backward�   s   € àÔ!ˆÝ˜k¨5Ñ1Ô1ˆØˆ
rA   N)Ú__name__Ú
__module__Ú__qualname__Ústaticmethodr)   ÚTensorr`   re   © rA   r?   rZ   rZ   Š   sv   € € € € € Øð)˜EœLð )¨U¬\ð )ð )ð )ñ „\ð)ð ð 5¤<ð °E´Lð ð ð ñ „\ðð ð rA   rZ   c                   óF   ‡ — e Zd ZdZˆ fd„Zdej        dej        fd„Zˆ xZS )Ú	BloomGeluzN
    Partly copied from Megatron-DeepSpeed code and adapted for our needs
    c                 óH   •— t          ¦   «                              ¦   «          d S r]   )ÚsuperÚ__init__)ÚselfÚ	__class__s    €r?   rp   zBloomGelu.__init__œ   s   ø€ Ý‰Œ×ÒÑÔÐÐÐrA   rB   r   c                 ó6   — t                                |¦  «        S r]   )rZ   Úapply)rq   rB   s     r?   r`   zBloomGelu.forwardŸ   s   € Ý×!Ò! !Ñ$Ô$Ð$rA   )	rf   rg   rh   Ú__doc__rp   r)   rj   r`   Ú__classcell__©rr   s   @r?   rm   rm   —   sh   ø€ € € € € ðð ðð ð ð ð ð%˜œð %¨%¬,ð %ð %ð %ð %ð %ð %ð %ð %rA   rm   c                   óø   ‡ — e Zd Zddededz  fˆ fd„Zdej        deej        ej        ej        f         fd„Z	dej        dej        fd	„Z
	 	 	 ddej        dej        dej        dej        dedz  dedefd„Zˆ xZS )ÚBloomAttentionNÚconfigÚ	layer_idxc                 óø  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        |j        | _        |j        | _        | j        | j        z  | _        | j        | _        |j	        | _	        | j        | j        z  | j        k    r t          d| j        › d| j        › d�¦  «        ‚dt          j        | j        ¦  «        z  | _        d| _        || _        |€(t                                d| j        j        › d�¦  «         t)          j        | j        d| j        z  d¬	¦  «        | _        t)          j        | j        | j        ¦  «        | _        t)          j        |j        ¦  «        | _        d S )
NzA`hidden_size` must be divisible by num_heads (got `hidden_size`: z and `num_heads`: z).rN   zInstantiating z¹ without passing a `layer_idx` is not recommended and will lead to errors during the forward call if caching is used. Please make sure to provide a `layer_idx` when creating this class.r	   T©Úbias)ro   rp   Úpretraining_tpÚslow_but_exactÚhidden_sizeÚn_headr   Úhead_dimÚ
split_sizeÚhidden_dropoutÚ
ValueErrorr&   ÚsqrtÚinv_norm_factorÚbetar{   ÚloggerÚwarning_oncerr   rf   r   ÚLinearÚquery_key_valueÚdenseÚDropoutÚattention_dropout)rq   rz   r{   rr   s      €r?   rp   zBloomAttention.__init__¤   sx  ø€ Ý‰Œ×ÒÑÔÐà$Ô3ˆÔØ$Ô3ˆÔà!Ô-ˆÔØœˆŒØÔ(¨D¬NÑ:ˆŒØÔ*ˆŒØ$Ô3ˆÔàŒ=˜4œ>Ñ)¨TÔ-=Ò=Ð=Ýð'ÐTXÔTdð 'ð 'Ø”Nð'ð 'ð 'ñô ð ð  #¥T¤Y¨t¬}Ñ%=Ô%=Ñ=ˆÔØˆŒ	Ø"ˆŒØÐÝ×Òð, ¤Ô!8ð ,ð ,ð ,ñô ð õ  "œy¨Ô)9¸1¸tÔ?OÑ;OÐVZÐ[Ñ[Ô[ˆÔÝ”Y˜tÔ/°Ô1AÑBÔBˆŒ
Ý!#¤¨FÔ,DÑ!EÔ!EˆÔÐÐrA   Ú	fused_qkvr   c                 ó.  — |j         \  }}}|                     ||| j        d| j        ¦  «        }|dddd…f                              dd¦  «        }|dddd…f                              dd¦  «        }|dddd…f                              dd¦  «        }|||fS )a  
        Split the last dimension into (num_heads, head_dim) and reshapes to (bs, heads, len, dim) shape
        without making any copies, results share same memory storage as `fused_qkv`

        Args:
            fused_qkv (`torch.tensor`): [batch_size, seq_length, num_heads * 3 * head_dim]

        Returns:
            query: [batch_size, num_heads, seq_length, head_dim]
            key: [batch_size, num_heads, seq_length, head_dim]
            value: [batch_size, num_heads, seq_length, head_dim]
        r	   .r   Nr   r   )r%   Úviewr   rƒ   Ú	transpose)rq   r‘   r4   r5   Úthree_times_hidden_sizeÚquery_layerÚ	key_layerÚvalue_layers           r?   Ú_reshapezBloomAttention._reshapeÅ   sª   € ð ;D¼/Ñ7ˆ
�JÐ 7Ø—N’N :¨z¸4¼>È1ÈdÌmÑ\Ô\ˆ	Ø  Q¨¨¨ 	Ô*×4Ò4°Q¸Ñ:Ô:ˆØ˜c 1 a a a˜iÔ(×2Ò2°1°aÑ8Ô8ˆ	Ø  Q¨¨¨ 	Ô*×4Ò4°Q¸Ñ:Ô:ˆØ˜I {Ð2Ð2rA   rB   c                 óè   — |j         \  }}}|| j        z  }|                     || j        || j        ¦  «        }|                     dddd¦  «        }|                     ||| j        | j        z  ¦  «        S )z÷
        Merge heads together over the last dimension

        Args:
            x (`torch.tensor`): [batch_size * num_heads, seq_length, head_dim]

        Returns:
            torch.tensor: [batch_size, seq_length, num_heads * head_dim]
        r   r   r   r	   )r%   r   r“   rƒ   Úpermuter2   )rq   rB   Úbatch_size_and_num_headsr5   Ú_r4   s         r?   Ú_merge_headszBloomAttention._merge_headsÙ   sv   € ð 34´'Ñ/Ð  *¨aØ-°´Ñ?ˆ
ð �FŠF�:˜tœ~¨z¸4¼=ÑIÔIˆð �IŠI�a˜˜A˜qÑ!Ô!ˆð �yŠy˜ Z°´À$Ä-Ñ1OÑPÔPÐPrA   FÚhidden_statesrC   r>   r   Ú
layer_pastÚ	use_cacheÚoutput_attentionsc                 ó|  — |j         \  }	}
}|                      |¦  «        }|                      |¦  «        \  }}}|�|                     ||| j        ¦  «        \  }}|                     |	| j        z  d| j        ¦  «        }|                     |	| j        z  d| j        ¦  «                             dd¦  «        }|                     |	| j        z  d| j        ¦  «        }| 	                    ||| j
        | j        ¬¦  «        }|                     |	| j        |
d¦  «        }|�||z   }t          j        |dt          j        ¬¦  «                             |j        ¦  «        }|                      |¦  «        }|                     |	| j        z  |
d¦  «        }t          j        ||¦  «        }|                      |¦  «        }| j        dk    rÅ| j        r¾| j        | j        z  }t          j        |¦  «        }t5          | j        ¦  «        D ]…}|t          j        |d d …d d …t9          ||z  ¦  «        t9          |dz   |z  ¦  «        …f         | j        j        d d …t9          ||z  ¦  «        t9          |dz   |z  ¦  «        …f         ¦  «        z   }Œ†n|                      |¦  «        }t?          ||| j         | j!        ¦  «        }||fS )Nr$   éþÿÿÿ)Úbatch1Úbatch2r‰   Úalpha)r#   r   r   )"r%   r�   r™   Úupdater{   r2   r   rƒ   r”   Úbaddbmmr‰   rˆ   r“   rH   Úsoftmaxr)   r+   r3   r   r�   Úbmmrž   r   r€   r�   Ú
zeros_likeÚrangeÚlinearÚintrŽ   ÚweightrK   r…   rE   )rq   rŸ   rC   r>   r   r    r¡   r¢   Úkwargsr4   Úq_lengthr�   r‘   r–   r—   r˜   Úattention_scoresÚattn_weightsÚattention_probsÚattention_probs_reshapedÚcontext_layerÚslicesÚoutput_tensorÚis                           r?   r`   zBloomAttention.forwardò   sÓ  € ð #0Ô"5Ñˆ
�H˜aØ×(Ò(¨Ñ7Ô7ˆ	à.2¯mªm¸IÑ.FÔ.FÑ+ˆ�Y àÐ!Ø%/×%6Ò%6°yÀ+ÈtÌ~Ñ%^Ô%^Ñ"ˆI�{ð "×)Ò)¨*°t´~Ñ*EÀrÈ4Ì=ÑYÔYˆØ×%Ò% j°4´>Ñ&AÀ2ÀtÄ}ÑUÔU×_Ò_Ð`bÐdfÑgÔgˆ	Ø!×)Ò)¨*°t´~Ñ*EÀrÈ4Ì=ÑYÔYˆð !Ÿ=š=ØØØ”ØÔ&ð	 )ñ 
ô 
Ðð (×,Ò,¨Z¸¼ÈÐSUÑVÔVˆØÐ%Ø'¨.Ñ8ˆLõ œ) L°bÅÄÐNÑNÔN×QÒQÐR]ÔRcÑdÔdˆð ×0Ò0°ÑAÔAˆð $3×#7Ò#7¸
ÀTÄ^Ñ8SÐU]Ð_aÑ#bÔ#bÐ õ œ	Ð":¸KÑHÔHˆð ×)Ò)¨-Ñ8Ô8ˆð Ô Ò"Ð" tÔ':Ð"ØÔ%¨Ô(;Ñ;ˆFÝ!Ô,¨]Ñ;Ô;ˆMÝ˜4Ô.Ñ/Ô/ð ð �Ø -µ´Ø! ! ! ! Q Q Q­¨A°©J©¬½#¸qÀ1¹uÈÑ>NÑ:OÔ:OÐ(OÐ"OÔPØ”JÔ% a a a­¨Q°©Z©¬½3ÀÀAÁÈÑ?OÑ;PÔ;PÐ)PÐ&PÔQñ1ô 1ñ !��ðð !ŸJšJ }Ñ5Ô5ˆMå# M°8¸TÔ=PÐRVÔR_Ñ`Ô`ˆØ˜oÐ-Ð-rA   r]   ©NFF)rf   rg   rh   r   r¯   rp   r)   rj   Útupler™   rž   r
   Úboolr`   rv   rw   s   @r?   ry   ry   £   sF  ø€ € € € € ðFð F˜{ð F°s¸T±zð Fð Fð Fð Fð Fð FðB3 %¤,ð 3°5¸¼ÀuÄ|ÐUZÔUaÐ9aÔ3bð 3ð 3ð 3ð 3ð(Q˜eœlð Q¨u¬|ð Qð Qð Qð Qð> $(ØØ"'ðA.ð A.à”|ðA.ð ”,ðA.ð Œ|ð	A.ð
 œðA.ð ˜D‘LðA.ð ðA.ð  ðA.ð A.ð A.ð A.ð A.ð A.ð A.ð A.rA   ry   c                   óV   ‡ — e Zd Zdefˆ fd„Zdej        dej        dej        fd„Zˆ xZS )ÚBloomMLPrz   c                 ó8  •— t          ¦   «                              ¦   «          |j        }|j        | _        |j        | _        t          j        |d|z  ¦  «        | _        t          ¦   «         | _	        t          j        d|z  |¦  «        | _
        |j        | _        d S )Né   )ro   rp   r�   r   r€   r   rŒ   Údense_h_to_4hrm   Ú	gelu_implÚdense_4h_to_hr…   )rq   rz   r�   rr   s      €r?   rp   zBloomMLP.__init__7  sƒ   ø€ Ý‰Œ×ÒÑÔÐØÔ(ˆà$Ô3ˆÔØ$Ô3ˆÔÝœY {°A¸±OÑDÔDˆÔÝ"™œˆŒÝœY q¨;¡¸ÑDÔDˆÔØ$Ô3ˆÔÐÐrA   rŸ   rC   r   c                 óx  — |                       |                      |¦  «        ¦  «        }| j        dk    rÕ| j        rÎt	          j        |¦  «        }| j        j        j        d         | j        z  }t          | j        ¦  «        D ]…}|t          j        |d d …d d …t          ||z  ¦  «        t          |dz   |z  ¦  «        …f         | j        j        d d …t          ||z  ¦  «        t          |dz   |z  ¦  «        …f         ¦  «        z   }Œ†n|                      |¦  «        }t          ||| j        | j        ¦  «        }|S )Nr   r$   )rÃ   rÂ   r   r€   r)   r¬   rÄ   r°   r%   r­   rH   r®   r¯   rK   r…   rE   )rq   rŸ   rC   Úintermediate_outputr¸   rº   Úoutputs          r?   r`   zBloomMLP.forwardB  sC  € ØŸš t×'9Ò'9¸-Ñ'HÔ'HÑIÔIˆàÔ Ò"Ð" tÔ':Ð"Ý"'Ô"2°8Ñ"<Ô"<ÐØÔ'Ô.Ô4°RÔ8¸4Ô;NÑNˆFÝ˜4Ô.Ñ/Ô/ð ð �Ø&9½A¼HØ! ! ! ! Q Q Q­¨A°©J©¬½#¸qÀ1¹uÈÑ>NÑ:OÔ:OÐ(OÐ"OÔPØÔ&Ô-¨a¨a¨aµ°Q¸±Z±´Å3ÈÈAÉÐQWÑGWÑCXÔCXÐ1XÐ.XÔYñ=ô =ñ 'Ð#Ð#ðð #'×"4Ò"4°]Ñ"CÔ"CÐåÐ0°(¸DÔ<OÐQUÔQ^Ñ_Ô_ˆàˆrA   )	rf   rg   rh   r   rp   r)   rj   r`   rv   rw   s   @r?   r¿   r¿   6  ss   ø€ € € € € ð	4˜{ð 	4ð 	4ð 	4ð 	4ð 	4ð 	4ð U¤\ð ¸U¼\ð ÈeÌlð ð ð ð ð ð ð ð rA   r¿   c                   ó|   ‡ — e Zd Zddededz  fˆ fd„Z	 	 	 ddej        dej        dej        d	edz  d
e	de	fd„Z
ˆ xZS )Ú
BloomBlockNrz   r{   c                 ó\  •— t          ¦   «                              ¦   «          |j        }t          ||j        ¬¦  «        | _        |j        | _        t          ||¦  «        | _	        t          ||j        ¬¦  «        | _
        t          |¦  «        | _        |j        | _        |j        | _        d S )N©Úeps)ro   rp   r�   r   Úlayer_norm_epsilonÚinput_layernormr‚   r   ry   Úself_attentionÚpost_attention_layernormr¿   ÚmlpÚ(apply_residual_connection_post_layernormr…   )rq   rz   r{   r�   rr   s       €r?   rp   zBloomBlock.__init__V  s—   ø€ Ý‰Œ×ÒÑÔÐØÔ(ˆå(¨¸&Ô:SÐTÑTÔTˆÔØœˆŒÝ,¨V°YÑ?Ô?ˆÔÝ(1°+À6ÔC\Ð(]Ñ(]Ô(]ˆÔ%å˜FÑ#Ô#ˆŒà8>Ô8gˆÔ5Ø$Ô3ˆÔÐÐrA   FrŸ   r>   r   r    r¡   r¢   c           	      óø   — |                       |¦  «        }| j        r|}	n|}	|                      ||	|||||¬¦  «        \  }
}|                      |
¦  «        }| j        r|}	n|
}	|                      ||	¦  «        }||fS )N)r    r   r>   r¡   r¢   )rÎ   rÒ   rÏ   rÐ   rÑ   )rq   rŸ   r>   r   r    r¡   r¢   r±   Úlayernorm_outputrC   Úattention_outputr´   rÇ   s                r?   r`   zBloomBlock.forwardd  s¹   € ð  ×/Ò/°Ñ>Ô>Ðð Ô8ð 	%Ø'ˆHˆHà$ˆHð *.×)<Ò)<ØØØ!Ø)ØØØ/ð *=ñ *
ô *
Ñ&Ð˜,ð  ×8Ò8Ð9IÑJÔJÐð Ô8ð 	(Ø'ˆHˆHà'ˆHð —’Ð*¨HÑ5Ô5ˆà�|Ð#Ð#rA   r]   r»   )rf   rg   rh   r   r¯   rp   r)   rj   r
   r½   r`   rv   rw   s   @r?   rÉ   rÉ   U  s¼   ø€ € € € € ð4ð 4˜{ð 4°s¸T±zð 4ð 4ð 4ð 4ð 4ð 4ð& $(ØØ"'ð+$ð +$à”|ð+$ð Œ|ð+$ð œð	+$ð
 ˜D‘Lð+$ð ð+$ð  ð+$ð +$ð +$ð +$ð +$ð +$ð +$ð +$rA   rÉ   c                   ó2   — e Zd ZU eed<   dZdZdgZdgZdZ	dS )ÚBloomPreTrainedModelrz   ÚtransformerTrÉ   Úpast_key_valuesN)
rf   rg   rh   r   Ú__annotations__Úbase_model_prefixÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_skip_keys_device_placementÚ_can_compile_fullgraphrk   rA   r?   r×   r×   ’  sA   € € € € € € àÐÐÑØ%ÐØ&*Ð#Ø%˜ÐØ#4Ð"5ÐØ!ÐÐÐrA   r×   c                   ó2  ‡ — e Zd Zdefˆ fd„Zdej        dedej        dej        fd„Z	d„ Z
d	ej        fd
„Ze	 	 	 	 	 	 	 	 d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dz  dedz  dedz  deej        df         ez  fd„¦   «         Zˆ xZS )Ú
BloomModelrz   c                 óè  •‡— t          ¦   «                              ‰¦  «         ‰j        | _        ‰j        | _        t          j        ‰j        | j        ¦  «        | _	        t          | j        ‰j        ¬¦  «        | _        t          j        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        t          | j        ‰j        ¬¦  «        | _        d| _        |                      ¦   «          d S )NrË   c                 ó2   •— g | ]}t          ‰|¬ ¦  «        ‘ŒS ))r{   )rÉ   )Ú.0rº   rz   s     €r?   ú
<listcomp>z'BloomModel.__init__.<locals>.<listcomp>©  s&   ø€ ÐiÐiÐiÀA¥
¨6¸QÐ ?Ñ ?Ô ?ÐiÐiÐirA   F)ro   rp   r�   Ú	embed_dimr‚   r   r   Ú	EmbeddingÚ
vocab_sizeÚword_embeddingsr   rÍ   Úword_embeddings_layernormÚ
ModuleListr­   Únum_hidden_layersÚhÚln_fÚgradient_checkpointingÚ	post_init©rq   rz   rr   s    `€r?   rp   zBloomModel.__init__ž  sÒ   øø€ Ý‰Œ×Ò˜Ñ Ô Ð àÔ+ˆŒØœˆŒõ  "œ|¨FÔ,=¸t¼~ÑNÔNˆÔÝ)2°4´>ÀvÔG`Ð)aÑ)aÔ)aˆÔ&õ ”ÐiÐiÐiÐiÍÈvÔOgÑIhÔIhÐiÑiÔiÑjÔjˆŒõ ˜dœn°&Ô2KÐLÑLÔLˆŒ	à&+ˆÔ#ð 	�ŠÑÔÐÐÐrA   r   r   r   r   c                 ó$   — t          |||¦  «        S r]   )r@   )rq   r   r   r   s       r?   r@   zBloomModel.build_alibi_tensor³  s   € Ý! .°)¸UÑCÔCÐCrA   c                 ó   — | j         S r]   ©ré   )rq   s    r?   Úget_input_embeddingszBloomModel.get_input_embeddings¶  s   € ØÔ#Ð#rA   Únew_embeddingsc                 ó   — || _         d S r]   rô   ©rq   rö   s     r?   Úset_input_embeddingszBloomModel.set_input_embeddings¹  s   € Ø-ˆÔÐÐrA   NÚ	input_idsrÙ   Úinputs_embedsr¡   r¢   Úoutput_hidden_statesÚreturn_dict.c	           	      ó  — |�|n| j         j        }|�|n| j         j        }|�|n| j         j        }|�|n| j         j        }|du |duz  rt          d¦  «        ‚| j        r%| j        r|rt           	                    d¦  «         d}|€|  
                    |¦  «        }|r|€t          | j         ¬¦  «        }|j        \  }
}}|�|                     ¦   «         nd}||z   }|                      |¦  «        }|rdnd}|rdnd}|€t          j        |
|f|j        ¬¦  «        }n|                     |j        ¦  «        }|                      || j        |j        ¬	¦  «        }t-          | j         |||¬
¦  «        }t/          | j        ¦  «        D ]4\  }}|r||fz   } |||||||¬¦  «        }|d         }|r||d         fz   }Œ5|                      |¦  «        }|r||fz   }|st5          d„ ||||fD ¦   «         ¦  «        S t7          ||||¬¦  «        S )á²  
        input_ids (`torch.LongTensor` of shape `(batch_size, input_ids_length)`):
            `input_ids_length` = `sequence_length` if `past_key_values` is `None` else `past_key_values.get_seq_length()`
            (`sequence_length` of input past key value states). Indices of input sequence tokens in the vocabulary.

            If `past_key_values` is used, only `input_ids` that do not have their past calculated should be passed as
            `input_ids`.

            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_embedszZ`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`...F)rz   r   rk   ©r!   )r   )rz   rû   r   rÙ   )r    r   r¡   r¢   r>   r   c              3   ó   K  — | ]}|®|V — Œ	d S r]   rk   )rä   Úvs     r?   ú	<genexpr>z%BloomModel.forward.<locals>.<genexpr>  s1   è è € ð ð ØÐghÐgt�ÐgtÐgtÐgtÐgtðð rA   )Úlast_hidden_staterÙ   rŸ   Ú
attentions)rz   r¢   rü   r¡   rý   r†   rï   rE   rŠ   r‹   ré   r   r%   Úget_seq_lengthrê   r)   Úonesr!   r3   r@   r   r   r   Ú	enumeraterí   rî   r¼   r   )rq   rú   rÙ   r   rû   r¡   r¢   rü   rý   r±   r4   r5   r�   Úpast_lengthÚseq_length_with_pastrŸ   Úall_self_attentionsÚall_hidden_statesr>   Úcausal_maskrº   ÚblockÚoutputss                          r?   r`   zBloomModel.forward¼  sõ  € ð4 2CÐ1NÐ-Ð-ÐTXÔT_ÔTqÐà$8Ð$DÐ Ð È$Ì+ÔJjð 	ð "+Ð!6�I�I¸D¼KÔ<Qˆ	Ø%0Ð%<�k�kÀ$Ä+ÔBYˆà˜Ð -°tÐ";Ñ<ð 	[ÝÐYÑZÔZÐZàÔ&ð 	¨4¬=ð 	¸Yð 	Ý×ÒØlñô ð ð ˆIàÐ Ø ×0Ò0°Ñ;Ô;ˆMàð 	?˜Ð0Ý*°$´+Ð>Ñ>Ô>ˆOà$1Ô$7Ñ!ˆ
�J Ø:IÐ:U�o×4Ò4Ñ6Ô6Ð6Ð[\ˆØ)¨KÑ7Ðà×6Ò6°}ÑEÔEˆà$5Ð?˜b˜b¸4ÐØ"6Ð@˜B˜B¸DÐð Ð!Ý"œZ¨Ð5IÐ(JÐS`ÔSgÐhÑhÔhˆNˆNà+×.Ò.¨}Ô/CÑDÔDˆNà×'Ò'¨¸¼ÈmÔNaÐ'ÑbÔbˆÝ(Ø”;Ø'Ø)Ø+ð	
ñ 
ô 
ˆõ " $¤&Ñ)Ô)ð 	Jð 	J‰HˆAˆuØ#ð IØ$5¸Ð8HÑ$HÐ!à�eØØ*Ø*Ø#Ø"3Øðñ ô ˆGð $ AœJˆMØ ð JØ&9¸WÀQ¼Z¸MÑ&IÐ#øð Ÿ	š	 -Ñ0Ô0ˆàð 	EØ 1°]Ð4DÑ DÐàð 	Ýð ð Ø)¨?Ð<MÐObÐcðñ ô ñ ô ð õ 9Ø+Ø+Ø+Ø*ð	
ñ 
ô 
ð 	
rA   ©NNNNNNNN)rf   rg   rh   r   rp   r)   rj   r¯   r   r@   rõ   rù   r   Ú
LongTensorr
   r½   r¼   r   r`   rv   rw   s   @r?   rá   rá   œ  s˜  ø€ € € € € ð˜{ð ð ð ð ð ð ð*D°´ð DÈ#ð DÐV[ÔVað DÐfkÔfrð Dð Dð Dð Dð$ð $ð $ð.°5´<ð .ð .ð .ð .ð ð .2Ø(,Ø.2Ø15Ø!%Ø)-Ø,0Ø#'ðg
ð g
àÔ# dÑ*ðg
ð  ™ðg
ð œ tÑ+ð	g
ð
 Ô'¨$Ñ.ðg
ð ˜$‘;ðg
ð   $™;ðg
ð # T™kðg
ð ˜D‘[ðg
ð 
ˆuŒ|˜SÐ Ô	!Ð$MÑ	Mðg
ð g
ð g
ñ „^ðg
ð g
ð g
ð g
ð g
rA   rá   zˆ
    The Bloom Model transformer with a language modeling head on top (linear layer with weights tied to the input
    embeddings).
    )Úcustom_introc                   ó<  ‡ — e Zd ZddiZdefˆ fd„Zdej        fd„Z	 	 	 	 	 dˆ fd
„	Z	e
	 	 	 	 	 	 	 	 	 	 ddej        dz  dedz  dej        dz  dej        dz  dej        dz  dedz  dedz  dedz  dedz  deej        z  deej                 ez  fd„¦   «         Zˆ xZS )ÚBloomForCausalLMzlm_head.weightz"transformer.word_embeddings.weightrz   c                 óæ   •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          j        |j        |j        d¬¦  «        | _        |  	                    ¦   «          d S ©NFr}   )
ro   rp   rá   rØ   r   rŒ   r�   rè   Úlm_headrð   rñ   s     €r?   rp   zBloomForCausalLM.__init__0  sa   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý% fÑ-Ô-ˆÔÝ”y Ô!3°VÔ5FÈUÐSÑSÔSˆŒð 	�ŠÑÔÐÐÐrA   rö   c                 ó   — || _         d S r]   )r  rø   s     r?   Úset_output_embeddingsz&BloomForCausalLM.set_output_embeddings8  s   € Ø%ˆŒˆˆrA   NTFc           	      ó:  •—  t          ¦   «         j        |f|||||dœ|¤Ž}t          |t          ¦  «        rd|�b|                     ¦   «         }	|j        \  }
}|	|z
  }t          j        |
||j        |j	        ¬¦  «        }t          j
        ||gd¬¦  «        }||d<   |S )N)rÙ   r   rû   r¡   Úis_first_iterationr    r$   r"   r   )ro   Úprepare_inputs_for_generationÚ
isinstancer   Úget_max_lengthr%   r)   Úzerosr!   r   r0   )rq   rú   rÙ   r   rû   r¡   r  r±   Úmodel_inputsÚtarget_lengthr4   r5   ÚdiffÚnew_attn_maskrr   s                 €r?   r  z.BloomForCausalLM.prepare_inputs_for_generation;  sÍ   ø€ ð =•u‘w”wÔ<Øð
à+Ø)Ø'ØØ1ð
ð 
ð ð
ð 
ˆõ �o¥{Ñ3Ô3ð 	<¸Ð8RØ+×:Ò:Ñ<Ô<ˆMØ%3Ô%9Ñ"ˆJ˜
Ø  :Ñ-ˆDå!œK¨
°DÀÔAVÐ^lÔ^rÐsÑsÔsˆMÝ"œY¨¸Ð'FÈBÐOÑOÔOˆNØ-;ˆLÐ)Ñ*àÐrA   r   rú   rÙ   r   rû   Úlabelsr¡   r¢   rü   rý   Úlogits_to_keepr   c           
      óî  — |	�|	n| j         j        }	|                      ||||||||	¬¦  «        }|d         }t          |
t          ¦  «        rt          |
 d¦  «        n|
}|                      |dd…|dd…f         ¦  «        }d}|�6|                      ||| j         j        | 	                    d¦  «        ¬¦  «        }|	s|f|dd…         z   }|�|f|z   n|S t          |||j        |j        |j        ¬¦  «        S )a\  
        input_ids (`torch.LongTensor` of shape `(batch_size, input_ids_length)`):
            `input_ids_length` = `sequence_length` if `past_key_values` is `None` else `past_key_values.get_seq_length()`
            (`sequence_length` of input past key value states). Indices of input sequence tokens in the vocabulary.

            If `past_key_values` is used, only `input_ids` that do not have their past calculated should be passed as
            `input_ids`.

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

            [What are input IDs?](../glossary#input-ids)
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for language modeling. Note that the labels **are shifted** inside the model, i.e. you can set
            `labels = input_ids` Indices are selected in `[-100, 0, ..., config.vocab_size]` All labels set to `-100`
            are ignored (masked), the loss is only computed for labels in `[0, ..., config.vocab_size]`
        N©rÙ   r   rû   r¡   r¢   rü   rý   r   Únum_items_in_batch)rè   r(  r   ©ÚlossÚlogitsrÙ   rŸ   r  )rz   rý   rØ   r  r¯   Úslicer  Úloss_functionrè   Úgetr   rÙ   rŸ   r  )rq   rú   rÙ   r   rû   r$  r¡   r¢   rü   rý   r%  r±   Útransformer_outputsrŸ   Úslice_indicesr+  r*  rÇ   s                     r?   r`   zBloomForCausalLM.forward^  sM  € ð@ &1Ð%<�k�kÀ$Ä+ÔBYˆà"×.Ò.ØØ+Ø)Ø'ØØ/Ø!5Ø#ð /ñ 	
ô 	
Ðð ,¨AÔ.ˆå8BÀ>ÕSVÑ8WÔ8WÐk�˜~˜o¨tÑ4Ô4Ð4Ð]kˆØ—’˜m¨A¨A¨A¨}¸a¸a¸aÐ,?Ô@ÑAÔAˆàˆØÐØ×%Ò%ØØØœ;Ô1Ø#)§:¢:Ð.BÑ#CÔ#Cð	 &ñ ô ˆDð ð 	FØ�YÐ!4°Q°R°RÔ!8Ñ8ˆFØ)-Ð)9�T�G˜fÑ$Ð$¸vÐEå0ØØØ/Ô?Ø-Ô;Ø*Ô5ð
ñ 
ô 
ð 	
rA   )NNNTF)
NNNNNNNNNr   )rf   rg   rh   Ú_tied_weights_keysr   rp   r)   rj   r  r  r   r  r
   r½   r¯   r¼   r   r`   rv   rw   s   @r?   r  r  '  s©  ø€ € € € € ð +Ð,PÐQÐð˜{ð ð ð ð ð ð ð&°E´Lð &ð &ð &ð &ð ØØØØ ð!ð !ð !ð !ð !ð !ðF ð .2Ø(,Ø.2Ø-1Ø&*Ø!%Ø)-Ø,0Ø#'Ø-.ðD
ð D
àÔ# dÑ*ðD
ð  ™ðD
ð œ tÑ+ð	D
ð
 ”| dÑ*ðD
ð ”˜tÑ#ðD
ð ˜$‘;ðD
ð   $™;ðD
ð # T™kðD
ð ˜D‘[ðD
ð ˜eœlÑ*ðD
ð 
ˆuŒ|Ô	Ð@Ñ	@ðD
ð D
ð D
ñ „^ðD
ð D
ð D
ð D
ð D
rA   r  aÖ  
    The Bloom Model transformer with a sequence classification head on top (linear layer).

    [`BloomForSequenceClassification`] uses the last token in order to do the classification, as other causal models
    (e.g. GPT-1) do.

    Since it does classification on the last token, it requires to know the position of the last token. If a
    `pad_token_id` is defined in the configuration, it finds the last token that is not a padding token in each row. If
    no `pad_token_id` is defined, it simply takes the last value in each row of the batch. Since it cannot guess the
    padding tokens when `inputs_embeds` are passed instead of `input_ids`, it does the same (take the last value in
    each row of the batch).
    c                   óò   ‡ — e Zd Zdefˆ fd„Ze	 	 	 	 	 	 	 	 	 ddej        dz  dedz  dej	        dz  dej	        dz  dej	        dz  d	e
dz  d
e
dz  de
dz  de
dz  deej	                 ez  fd„¦   «         Zˆ xZS )ÚBloomForSequenceClassificationrz   c                 óþ   •— t          ¦   «                              |¦  «         |j        | _        t          |¦  «        | _        t          j        |j        |j        d¬¦  «        | _        |  	                    ¦   «          d S r  )
ro   rp   Ú
num_labelsrá   rØ   r   rŒ   r�   Úscorerð   rñ   s     €r?   rp   z'BloomForSequenceClassification.__init__µ  sk   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø Ô+ˆŒÝ% fÑ-Ô-ˆÔÝ”Y˜vÔ1°6Ô3DÈ5ÐQÑQÔQˆŒ
ð 	�ŠÑÔÐÐÐrA   Nrú   rÙ   r   rû   r$  r¡   r¢   rü   rý   r   c
           
      ó¦  — |	�|	n| j         j        }	|                      ||||||||	¬¦  «        }|d         }|                      |¦  «        }|�|j        d         }n|j        d         }| j         j        €|dk    rt          d¦  «        ‚| j         j        €d}n¨|�}|| j         j        k                         |j        t          j
        ¦  «        }t          j        |j        d         |j        t          j
        ¬¦  «        }||z                       d¦  «        }n)d}t                               | j        j        › d�¦  «         |t          j        ||j        ¬	¦  «        |f         }d}|��.| j         j        €f| j        dk    rd
| j         _        nN| j        dk    r7|j        t          j        k    s|j        t          j        k    rd| j         _        nd| j         _        | j         j        d
k    rWt-          ¦   «         }| j        dk    r1 ||                     ¦   «         |                     ¦   «         ¦  «        }nb |||¦  «        }nU| j         j        dk    rt1          ¦   «         } |||¦  «        }n*| j         j        dk    rt3          ¦   «         } |||¦  «        }|	s|f|dd…         z   }|�|f|z   n|S t5          |||j        |j        |j        ¬¦  «        S )á6  
        input_ids (`torch.LongTensor` of shape `(batch_size, input_ids_length)`):
            `input_ids_length` = `sequence_length` if `past_key_values` is `None` else `past_key_values.get_seq_length()`
            (`sequence_length` of input past key value states). Indices of input sequence tokens in the vocabulary.

            If `past_key_values` is used, only `input_ids` that do not have their past calculated should be passed as
            `input_ids`.

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

            [What are input IDs?](../glossary#input-ids)
        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).
        Nr'  r   r   z=Cannot handle batch sizes > 1 if no padding token is defined.r$   r    zŠ will not detect padding tokens in `inputs_embeds`. Results may be unexpected if using padding tokens in conjunction with `inputs_embeds.`r   Ú
regressionÚsingle_label_classificationÚmulti_label_classificationr)  )rz   rý   rØ   r6  r%   Úpad_token_idr†   r3   r!   r)   r-   r,   ÚargmaxrŠ   r‹   rr   rf   Úproblem_typer5  r   Úlongr¯   r   Úsqueezer   r   r   rÙ   rŸ   r  )rq   rú   rÙ   r   rû   r$  r¡   r¢   rü   rý   r±   r/  rŸ   r+  r4   Úlast_non_pad_tokenÚnon_pad_maskÚtoken_indicesÚpooled_logitsr*  Úloss_fctrÇ   s                         r?   r`   z&BloomForSequenceClassification.forward¾  s  € ð> &1Ð%<�k�kÀ$Ä+ÔBYˆà"×.Ò.ØØ+Ø)Ø'ØØ/Ø!5Ø#ð /ñ 	
ô 	
Ðð ,¨AÔ.ˆØ—’˜MÑ*Ô*ˆàÐ Ø"œ¨Ô+ˆJˆJà&Ô,¨QÔ/ˆJàŒ;Ô#Ð+°
¸a²°ÝÐ\Ñ]Ô]Ð]ØŒ;Ô#Ð+Ø!#ÐÐØÐ"à%¨¬Ô)AÒA×EÒEÀfÄmÕUZÔU`ÑaÔaˆLÝ!œL¨¬¸Ô)<ÀVÄ]ÕZ_ÔZeÐfÑfÔfˆMØ"/°,Ñ">×!FÒ!FÀrÑ!JÔ!JÐÐà!#ÐÝ×ÒØ”>Ô*ð Zð Zð Zñô ð ð
 �uœ|¨J¸v¼}ÐMÑMÔMÐOaÐaÔbˆàˆØÑØŒ{Ô'Ð/Ø”? aÒ'Ð'Ø/;�D”KÔ,Ð,Ø”_ qÒ(Ð(¨f¬l½e¼jÒ.HÐ.HÈFÌLÕ\aÔ\eÒLeÐLeØ/L�D”KÔ,Ð,à/K�D”KÔ,àŒ{Ô'¨<Ò7Ð7Ý"™9œ9�Ø”? aÒ'Ð'Ø#˜8 M×$9Ò$9Ñ$;Ô$;¸V¿^º^Ñ=MÔ=MÑNÔN�D�Dà#˜8 M°6Ñ:Ô:�D�DØ”Ô)Ð-JÒJÐJÝ+Ñ-Ô-�Ø�x ¨vÑ6Ô6��Ø”Ô)Ð-IÒIÐIÝ,Ñ.Ô.�Ø�x ¨vÑ6Ô6�Øð 	FØ#Ð%Ð(;¸A¸B¸BÔ(?Ñ?ˆFØ)-Ð)9�T�G˜fÑ$Ð$¸vÐEå/ØØ Ø/Ô?Ø-Ô;Ø*Ô5ð
ñ 
ô 
ð 	
rA   ©	NNNNNNNNN)rf   rg   rh   r   rp   r   r)   r  r
   rj   r½   r¼   r   r`   rv   rw   s   @r?   r3  r3  ¦  s9  ø€ € € € € ð˜{ð ð ð ð ð ð ð ð .2Ø(,Ø.2Ø-1Ø&*Ø!%Ø)-Ø,0Ø#'ðe
ð e
àÔ# dÑ*ðe
ð  ™ðe
ð œ tÑ+ð	e
ð
 ”| dÑ*ðe
ð ”˜tÑ#ðe
ð ˜$‘;ðe
ð   $™;ðe
ð # T™kðe
ð ˜D‘[ðe
ð 
ˆuŒ|Ô	Ð?Ñ	?ðe
ð e
ð e
ñ „^ðe
ð e
ð e
ð e
ð e
rA   r3  c                   óò   ‡ — e Zd Zdefˆ fd„Ze	 	 	 	 	 	 	 	 	 ddej        dz  dedz  dej	        dz  dej	        dz  dej	        dz  d	e
dz  d
e
dz  de
dz  de
dz  deej	                 ez  fd„¦   «         Zˆ xZS )ÚBloomForTokenClassificationrz   c                 ó¬  •— t          ¦   «                              |¦  «         |j        | _        t          |¦  «        | _        t          |d¦  «        r|j        �|j        }n!t          |d¦  «        r|j        �|j        }nd}t          j	        |¦  «        | _
        t          j        |j        |j        ¦  «        | _        |                      ¦   «          d S )NÚclassifier_dropoutr…   gš™™™™™¹?)ro   rp   r5  rá   rØ   ÚhasattrrJ  r…   r   r�   rI   rŒ   r�   Ú
classifierrð   )rq   rz   rJ  rr   s      €r?   rp   z$BloomForTokenClassification.__init__)  sÌ   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø Ô+ˆŒå% fÑ-Ô-ˆÔÝ�6Ð/Ñ0Ô0ð 	%°VÔ5NÐ5ZØ!'Ô!:ÐÐÝ�VÐ-Ñ.Ô.ð 	%°6Ô3HÐ3TØ!'Ô!6ÐÐà!$ÐÝ”zÐ"4Ñ5Ô5ˆŒÝœ) FÔ$6¸Ô8IÑJÔJˆŒð 	�ŠÑÔÐÐÐrA   Nrú   rÙ   r   rû   r$  r¡   r¢   rü   rý   r   c
           
      ó  — |	�|	n| j         j        }	|                      ||||||||	¬¦  «        }|d         }|                      |¦  «        }|                      |¦  «        }d}|�p|                     |j        ¦  «        }|j        \  }}t          ¦   «         } || 	                    ||z  | j
        ¦  «        | 	                    ||z  ¦  «        ¦  «        }|	s|f|dd…         z   }|�|f|z   n|S t          |||j        |j        ¬¦  «        S )r8  Nr'  r   r   )r*  r+  rŸ   r  )rz   rý   rØ   rI   rL  r3   r!   r%   r   r“   r5  r   rŸ   r  )rq   rú   rÙ   r   rû   r$  r¡   r¢   rü   rý   r±   r/  rŸ   r+  r*  r4   r5   rE  rÇ   s                      r?   r`   z#BloomForTokenClassification.forward:  sL  € ð> &1Ð%<�k�kÀ$Ä+ÔBYˆà"×.Ò.ØØ+Ø)Ø'ØØ/Ø!5Ø#ð /ñ 	
ô 	
Ðð ,¨AÔ.ˆØŸš ]Ñ3Ô3ˆØ—’ Ñ/Ô/ˆàˆØÐà—Y’Y˜vœ}Ñ-Ô-ˆFØ%+¤\Ñ"ˆJ˜
Ý'Ñ)Ô)ˆHØ�8Ø—’˜J¨Ñ3°T´_ÑEÔEÀvÇ{Â{ÐS]Ð`jÑSjÑGkÔGkñô ˆDð ð 	FØ�YÐ!4°Q°R°RÔ!8Ñ8ˆFØ)-Ð)9�T�G˜fÑ$Ð$¸vÐEå$ØØØ-Ô;Ø*Ô5ð	
ñ 
ô 
ð 	
rA   rF  )rf   rg   rh   r   rp   r   r)   r  r
   rj   r½   r¼   r   r`   rv   rw   s   @r?   rH  rH  '  s9  ø€ € € € € ð˜{ð ð ð ð ð ð ð" ð .2Ø(,Ø.2Ø-1Ø&*Ø!%Ø)-Ø,0Ø#'ðB
ð B
àÔ# dÑ*ðB
ð  ™ðB
ð œ tÑ+ð	B
ð
 ”| dÑ*ðB
ð ”˜tÑ#ðB
ð ˜$‘;ðB
ð   $™;ðB
ð # T™kðB
ð ˜D‘[ðB
ð 
ˆuŒ|Ô	Ð4Ñ	4ðB
ð B
ð B
ñ „^ðB
ð B
ð B
ð B
ð B
rA   rH  c                   óÔ   ‡ — e Zd Zˆ fd„Z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dz  d
edz  de	e
z  fd„¦   «         Zˆ xZS )ÚBloomForQuestionAnsweringc                 óØ   •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          j        |j        d¦  «        | _        |                      ¦   «          d S )Nr   )	ro   rp   rá   rØ   r   rŒ   r�   Ú
qa_outputsrð   rñ   s     €r?   rp   z"BloomForQuestionAnswering.__init__‚  sY   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý% fÑ-Ô-ˆÔÝœ) FÔ$6¸Ñ:Ô:ˆŒð 	�ŠÑÔÐÐÐrA   Nrú   r   rû   Ústart_positionsÚend_positionsr¢   rü   rý   r   c	                 óª  — |�|n| j         j        }|                      ||||||¬¦  «        }
|
d         }|                      |¦  «        }|                     dd¬¦  «        \  }}|                     d¦  «                             ¦   «         }|                     d¦  «                             ¦   «         }d}|�ç|�åt          |                     ¦   «         ¦  «        dk    r|                     d¦  «        }t          |                     ¦   «         ¦  «        dk    r|                     d¦  «        }|                     d¦  «        }| 	                    d|¦  «        }| 	                    d|¦  «        }t          |¬¦  «        } |||¦  «        } |||¦  «        }||z   dz  }|s||f|
dd…         z   }|�|f|z   n|S t          ||||
j        |
j        ¬	¦  «        S )
rÿ   N)r   rû   r¢   rü   rý   r   r   r$   r"   )Úignore_indexr   )r*  Ústart_logitsÚ
end_logitsrŸ   r  )rz   rý   rØ   rQ  Úsplitr@  Ú
contiguousÚlenÚsizeÚclampr   r   rŸ   r  )rq   rú   r   rû   rR  rS  r¢   rü   rý   r±   r  Úsequence_outputr+  rV  rW  Ú
total_lossÚignored_indexrE  Ú
start_lossÚend_lossrÇ   s                        r?   r`   z!BloomForQuestionAnswering.forwardŠ  s  € ð4 &1Ð%<�k�kÀ$Ä+ÔBYˆà×"Ò"ØØ)Ø'Ø/Ø!5Ø#ð #ñ 
ô 
ˆð " !œ*ˆà—’ Ñ1Ô1ˆØ#)§<¢<°°r <Ñ#:Ô#:Ñ ˆ�jØ#×+Ò+¨BÑ/Ô/×:Ò:Ñ<Ô<ˆØ×'Ò'¨Ñ+Ô+×6Ò6Ñ8Ô8ˆ
àˆ
ØÐ&¨=Ð+Då�?×'Ò'Ñ)Ô)Ñ*Ô*¨QÒ.Ð.Ø"1×"9Ò"9¸"Ñ"=Ô"=�Ý�=×%Ò%Ñ'Ô'Ñ(Ô(¨1Ò,Ð,Ø -× 5Ò 5°bÑ 9Ô 9�à(×-Ò-¨aÑ0Ô0ˆMØ-×3Ò3°A°}ÑEÔEˆOØ)×/Ò/°°=ÑAÔAˆMå'°]ÐCÑCÔCˆHØ!˜ ,°Ñ@Ô@ˆJØ�x 
¨MÑ:Ô:ˆHØ$ xÑ/°1Ñ4ˆJàð 	RØ" JÐ/°'¸!¸"¸"´+Ñ=ˆFØ/9Ð/E�Z�M FÑ*Ð*È6ÐQå+ØØ%Ø!Ø!Ô/ØÔ)ð
ñ 
ô 
ð 	
rA   r  )rf   rg   rh   rp   r   r)   r  ÚFloatTensorr½   r¼   r   r`   rv   rw   s   @r?   rO  rO  €  s  ø€ € € € € ðð ð ð ð ð ð .2Ø37Ø26Ø37Ø15Ø)-Ø,0Ø#'ðF
ð F
àÔ# dÑ*ðF
ð Ô)¨DÑ0ðF
ð Ô(¨4Ñ/ð	F
ð
 Ô)¨DÑ0ðF
ð Ô'¨$Ñ.ðF
ð   $™;ðF
ð # T™kðF
ð ˜D‘[ðF
ð 
Ð-Ñ	-ðF
ð F
ð F
ñ „^ðF
ð F
ð F
ð F
ð F
rA   rO  )r  rá   r×   r3  rH  rO  )=ru   r&   r)   r   Útorch.nnr   r   r   r   r   rH   Úcache_utilsr
   r   r   Ú
generationr   Úmasking_utilsr   Úmodeling_layersr   Úmodeling_outputsr   r   r   r   r   Úmodeling_utilsr   Úutilsr   r   Úconfiguration_bloomr   Ú
get_loggerrf   rŠ   rj   r¯   r   r@   Úfloatr½   rK   rS   rX   ÚautogradÚFunctionrZ   ÚModulerm   ry   r¿   rÉ   r×   rá   r  r3  rH  rO  Ú__all__rk   rA   r?   ú<module>rr     sè  ðð Ð à €€€à €€€Ø Ð Ð Ð Ð Ð Ø LÐ LÐ LÐ LÐ LÐ LÐ LÐ LÐ LÐ LÐ LÐ LØ $Ð $Ð $Ð $Ð $Ð $à ;Ð ;Ð ;Ð ;Ð ;Ð ;Ð ;Ð ;Ð ;Ð ;Ø )Ð )Ð )Ð )Ð )Ð )Ø /Ð /Ð /Ð /Ð /Ð /Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9ðð ð ð ð ð ð ð ð ð ð ð ð ð ð .Ð -Ð -Ð -Ð -Ð -ðð ð ð ð ð ð ð ð -Ð ,Ð ,Ð ,Ð ,Ð ,ð 
ˆÔ	˜HÑ	%Ô	%€ð)J u¤|ð )JÀð )JÈEÌKð )JÐ\aÔ\hð )Jð )Jð )Jð )JðX�5”<ð ¨5¬<ð ¸uð ÐPTð ÐY^ÔYeð ð ð ð ð&	Q˜%œ,ð 	Q¨5¬<ð 	Qð 	Qð 	Qð 	Qð�u”|ð ¨¬ð ¸¼ð ð ð ð ð$
ð 
ð 
ð 
ð 
�5”>Ô*ñ 
ô 
ð 
ð	%ð 	%ð 	%ð 	%ð 	%�”	ñ 	%ô 	%ð 	%ðP.ð P.ð P.ð P.ð P.�R”Yñ P.ô P.ð P.ðfð ð ð ð ˆrŒyñ ô ð ð>:$ð :$ð :$ð :$ð :$Ð+ñ :$ô :$ð :$ðz ð"ð "ð "ð "ð "˜?ñ "ô "ñ „ð"ð ðG
ð G
ð G
ð G
ð G
Ð%ñ G
ô G
ñ „ðG
ðT €ððñ ô ðv
ð v
ð v
ð v
ð v
Ð+¨_ñ v
ô v
ñô ðv
ðr €ððñ ô ðp
ð p
ð p
ð p
ð p
Ð%9ñ p
ô p
ñô ðp
ðf ðU
ð U
ð U
ð U
ð U
Ð"6ñ U
ô U
ñ „ðU
ðp ðP
ð P
ð P
ð P
ð P
Ð 4ñ P
ô P
ñ „ðP
ðfð ð €€€rA   