Ë
    S^(hxF ã                   óì  — d Z ddlZddlZddlmZ ddlmZmZ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mZmZ dd
lmZ ddlmZmZmZ ddlmZmZmZmZ ddlm Z   ejB                  e"«      Z#dZ$dZ%dZ&dZ'd„ Z( G d„ de
jR                  «      Z* G d„ de
jR                  «      Z+ G d„ de
jR                  «      Z,de+iZ- G d„ de
jR                  «      Z. G d„ de
jR                  «      Z/ G d„ d e
jR                  «      Z0 G d!„ d"e
jR                  «      Z1 G d#„ d$e
jR                  «      Z2 G d%„ d&e
jR                  «      Z3e G d'„ d(e«      «       Z4e G d)„ d*e«      «       Z5e G d+„ d,e«      «       Z6e G d-„ d.e«      «       Z7 G d/„ d0e
jR                  «      Z8 G d1„ d2e
jR                  «      Z9 G d3„ d4e
jR                  «      Z: G d5„ d6e
jR                  «      Z; G d7„ d8e
jR                  «      Z<d9Z=d:Z> G d;„ d<e«      Z? G d=„ d>e?«      Z@ ed?e=«       G d@„ dAe?«      «       ZA edBe=«       G dC„ dDe?«      «       ZB edEe=«       G dF„ dGe?«      «       ZC edHe=«       G dI„ dJe?«      «       ZDdKZE edLe=«       G dM„ dNe?«      «       ZFy)OzPyTorch REALM model.é    N)Ú	dataclass)ÚOptionalÚTupleÚUnion)Únn)ÚCrossEntropyLossé   )ÚACT2FN)Ú)BaseModelOutputWithPastAndCrossAttentionsÚ,BaseModelOutputWithPoolingAndCrossAttentionsÚMaskedLMOutputÚModelOutput)ÚPreTrainedModel)Úapply_chunking_to_forwardÚ find_pruneable_heads_and_indicesÚprune_linear_layer)Úadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚloggingÚreplace_return_docstringsé   )ÚRealmConfigz(google/realm-cc-news-pretrained-embedderz'google/realm-cc-news-pretrained-encoderz&google/realm-cc-news-pretrained-scorerr   c           	      ó²
  — 	 ddl }ddl}ddl}t        j                  j                  |«      }t        j                  d|› �«       |j                  j                  |«      }g }g }	|D ]^  \  }
}t        j                  d|
› d|› �«       |j                  j                  ||
«      }|j                  |
«       |	j                  |«       Œ` t        ||	«      D �]é  \  }
}t        | t         «      r5d|
vr1t        j                  d|
› d	| j"                  j$                  › d
�«       ŒL|
j'                  d«      s|
j'                  d«      r4t        | t(        «      r$|
j+                  dd«      }
|
j+                  dd«      }
|
j'                  d«      s|
j'                  d«      r"t        | t,        «      r|
j+                  dd«      }
|
j'                  d«      r}t        | t         «      rdnd}|
j+                  d|› d�«      }
|
j+                  d|› d�«      }
|
j+                  d|› d�«      }
|
j+                  d|› d�«      }
|
j+                  d|› d�«      }
|
j'                  d«      r“t        | t.        «      rdnd}|
j+                  d|› d�«      }
|
j+                  d|› d �«      }
|
j+                  d!|› d"�«      }
|
j+                  d#|› d$�«      }
|
j+                  d%|› d�«      }
|
j+                  d&|› d$�«      }
nO|
j'                  d'«      r>t        | t.        «      rdnd}|
j+                  d(|› d �«      }
|
j+                  d)|› d"�«      }
|
j1                  d*«      }
t3        d+„ |
D «       «      r)t        j                  dd*j5                  |
«      › �«       �Œ³| }|
D ]–  }|j7                  d,|«      r|j1                  d-|«      }n|g}|d   d.k(  s|d   d/k(  rt9        |d0«      }n-|d   d1k(  s|d   d2k(  rt9        |d3«      }n	 t9        ||d   «      }t=        |«      d4k\  sŒ„t?        |d5   «      }||   }Œ˜ d6d d7k(  rt9        |d0«      }n|d.k(  r|jA                  |«      }	 |jB                  |jB                  k(  s"J d8|jB                  › d9|jB                  › d:�«       ‚	 t        j                  d;|
› �«       tI        jJ                  |«      |_&        �Œì | S # t        $ r t        j                  d«       ‚ w xY w# t:        $ r+ t        j                  dd*j5                  |
«      › �«       Y �ŒŽw xY w# tD        $ r1}|xjF                  |jB                  |jB                  fz  c_#        ‚ d}~ww xY w)<z'Load tf checkpoints in a pytorch model.r   Nz™Loading a TensorFlow model in PyTorch, requires TensorFlow to be installed. Please see https://www.tensorflow.org/install/ for installation instructions.z&Converting TensorFlow checkpoint from zLoading TF weight z with shape Úreaderz	Skipping z as it is not z's parameterÚbertÚclszbert/zreader/realm/zcls/zreader/cls/zrealm/Ú zreader/zreader/module/bert/zreader/module/cls/zreader/dense/zqa_outputs/dense_intermediate/zreader/dense_1/zqa_outputs/dense_output/zreader/layer_normalizationzqa_outputs/layer_normalizationzmodule/module/module/z	embedder/z!module/module/module/module/bert/zmodule/module/module/LayerNorm/zcls/LayerNorm/zmodule/module/module/dense/z
cls/dense/z,module/module/module/module/cls/predictions/zcls/predictions/zmodule/module/module/bert/z%module/module/module/cls/predictions/zmodule/module/zmodule/module/LayerNorm/zmodule/module/dense/ú/c              3   ó$   K  — | ]  }|d v –— Œ
 y­w))Úadam_vÚadam_mÚAdamWeightDecayOptimizerÚAdamWeightDecayOptimizer_1Úglobal_stepN© )Ú.0Úns     úq/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/deprecated/realm/modeling_realm.pyú	<genexpr>z+load_tf_weights_in_realm.<locals>.<genexpr>p   s   è ø€ ò 
àð ÐnÔnñ
ùó   ‚z[A-Za-z]+_\d+z_(\d+)ÚkernelÚgammaÚweightÚoutput_biasÚbetaÚbiasé   r   iõÿÿÿÚ_embeddingszPointer shape z and array shape z mismatchedzInitialize PyTorch weight )'ÚreÚnumpyÚ
tensorflowÚImportErrorÚloggerÚerrorÚosÚpathÚabspathÚinfoÚtrainÚlist_variablesÚload_variableÚappendÚzipÚ
isinstanceÚRealmReaderÚ	__class__Ú__name__Ú
startswithÚRealmForOpenQAÚreplaceÚRealmKnowledgeAugEncoderÚRealmEmbedderÚsplitÚanyÚjoinÚ	fullmatchÚgetattrÚAttributeErrorÚlenÚintÚ	transposeÚshapeÚAssertionErrorÚargsÚtorchÚ
from_numpyÚdata)ÚmodelÚconfigÚtf_checkpoint_pathr3   ÚnpÚtfÚtf_pathÚ	init_varsÚnamesÚarraysÚnamerT   ÚarrayÚreader_prefixÚembedder_prefixÚpointerÚm_nameÚscope_namesÚnumÚes                       r(   Úload_tf_weights_in_realmrl   .   sL  € ð
ÛãÛô �g‰g�o‰oÐ0Ó1€GÜ
‡K�KÐ8¸¸	ÐBÔCà—‘×'Ñ'¨Ó0€IØ€EØ€Fà ò ‰ˆˆeÜ�‰Ð(¨¨¨l¸5¸'ÐBÔCØ—‘×&Ñ& w°Ó5ˆØ�‰�TÔØ�‰�eÕð	ô ˜5 &Ó)ó M/‰ˆˆeÜ�eœ[Ô)¨h¸dÑ.BÜ�K‰K˜) D 6¨¸¿¹×8PÑ8PÐ7QÐQ]Ð^Ô_Øð �O‰O˜FÔ# t§¡°uÔ'=Ä:ÈeÔUcÔCdØ—<‘< ¨Ó9ˆDØ—<‘< ¨Ó6ˆDð �O‰O˜FÔ# t§¡°uÔ'=Ä:ÈeÔUmÔCnØ—<‘< ¨Ó2ˆDð �?‰?˜8Ô$Ü",¨U´KÔ"@™BÀiˆMØ—<‘<Ð 5¸-¸ÈÐ7OÓPˆDØ—<‘<Ð 4¸¸ÀtÐ6LÓMˆDØ—<‘< °M°?ÐB`Ð1aÓbˆDØ—<‘<Ð 1°m°_ÐD\Ð3]Ó^ˆDØ—<‘<Ð <ÀÀÐOmÐ>nÓoˆDð �?‰?Ð2Ô3Ü$.¨u´mÔ$D™bÈ+ˆOØ—<‘<Ð CÈÐGXÐX^ÐE_Ó`ˆDØ—<‘<Ð AÀoÐEVÐVdÐCeÓfˆDØ—<‘<Ð =À/ÐARÐR\Ð?]Ó^ˆDØ—<‘<Ð NÐSbÐRcÐcsÐPtÓuˆDØ—<‘<Ð <ÀÐ@QÐQWÐ>XÓYˆDØ—<‘<Ð GÈOÐK\Ð\lÐImÓn‰DØ�_‰_Ð-Ô.Ü$.¨u´mÔ$D™bÈ+ˆOØ—<‘<Ð :¸Ð>OÈ~Ð<^Ó_ˆDØ—<‘<Ð 6¸?Ð:KÈ:Ð8VÓWˆDà�z‰z˜#‹ˆô ñ 
àô
ô 
ô �K‰K˜) C§H¡H¨T£NÐ#3Ð4Ô5ÙØˆØò 	'ˆFØ�|‰|Ð,¨fÔ5Ø Ÿh™h y°&Ó9‘à%˜h�Ø˜1‰~ Ò)¨[¸©^¸wÒ-FÜ! '¨8Ó4‘Ø˜Q‘ =Ò0°KÀ±NÀfÒ4LÜ! '¨6Ó2‘ðÜ% g¨{¸1©~Ó>�Gô �;Ó 1Ó$Ü˜+ a™.Ó)�Ø! #™,‘ð#	'ð$ �#�$ˆ<˜=Ò(Ü˜g xÓ0‰GØ�xÒØ—L‘L Ó'ˆEð	Ø—=‘= E§K¡KÒ/ð Ø  §¡ Ð/@ÀÇÁÀÈ[ÐYóÑ/ô 	�‰Ð0°°Ð7Ô8Ü×'Ñ'¨Ó.ˆŽð[M/ð\ €LøôC ò Ü�‰ðQô	
ð 	ðûô\ &ò Ü—K‘K )¨C¯H©H°T«NÐ+;Ð <Ô=Úðûô ò 	Ø�FŠF�w—}‘} e§k¡kÐ2Ñ2�FØûð	ús5   ‚S Ï0S%Ñ;TÓ S"Ó%0TÔTÔ	UÔ%,UÕUc                   óÊ   ‡ — e Zd ZdZˆ fd„Z	 	 	 	 	 d
deej                     deej                     deej                     deej                     de	dej                  fd	„Zˆ xZS )ÚRealmEmbeddingszGConstruct the embeddings from word, position and token_type embeddings.c                 ó>  •— t         ‰| �  «        t        j                  |j                  |j
                  |j                  ¬«      | _        t        j                  |j                  |j
                  «      | _	        t        j                  |j                  |j
                  «      | _        t        j                  |j
                  |j                  ¬«      | _        t        j                  |j                  «      | _        t#        |dd«      | _        | j'                  dt)        j*                  |j                  «      j-                  d«      d¬«       | j'                  d	t)        j.                  | j0                  j3                  «       t(        j4                  ¬
«      d¬«       y )N)Úpadding_idx©ÚepsÚposition_embedding_typeÚabsoluteÚposition_ids)r   éÿÿÿÿF)Ú
persistentÚtoken_type_ids©Údtype)ÚsuperÚ__init__r   Ú	EmbeddingÚ
vocab_sizeÚhidden_sizeÚpad_token_idÚword_embeddingsÚmax_position_embeddingsÚposition_embeddingsÚtype_vocab_sizeÚtoken_type_embeddingsÚ	LayerNormÚlayer_norm_epsÚDropoutÚhidden_dropout_probÚdropoutrO   rs   Úregister_bufferrW   ÚarangeÚexpandÚzerosru   ÚsizeÚlong©Úselfr[   rD   s     €r(   r|   zRealmEmbeddings.__init__œ   s/  ø€ Ü‰ÑÔÜ!Ÿ|™|¨F×,=Ñ,=¸v×?QÑ?QÐ_e×_rÑ_rÔsˆÔÜ#%§<¡<°×0NÑ0NÐPV×PbÑPbÓ#cˆÔ Ü%'§\¡\°&×2HÑ2HÈ&×J\ÑJ\Ó%]ˆÔ"ô Ÿ™ f×&8Ñ&8¸f×>SÑ>SÔTˆŒÜ—z‘z &×"<Ñ"<Ó=ˆŒä'.¨vÐ7PÐR\Ó']ˆÔ$Ø×ÑØœEŸL™L¨×)GÑ)GÓH×OÑOÐPWÓXÐejð 	ô 	
ð 	×ÑØœeŸk™k¨$×*;Ñ*;×*@Ñ*@Ó*BÌ%Ï*É*ÔUÐbgð 	õ 	
ó    Ú	input_idsrx   ru   Úinputs_embedsÚpast_key_values_lengthÚreturnc                 óZ  — |�|j                  «       }n|j                  «       d d }|d   }|€| j                  d d …|||z   …f   }|€st        | d«      r-| j                  d d …d |…f   }|j	                  |d   |«      }	|	}n:t        j                  |t
        j                  | j                  j                  ¬«      }|€| j                  |«      }| j                  |«      }
||
z   }| j                  dk(  r| j                  |«      }||z  }| j                  |«      }| j                  |«      }|S )Nrv   r   rx   r   ©rz   Údevicert   )r�   ru   Úhasattrrx   r�   rW   rŽ   r�   rš   r�   r…   rs   rƒ   r†   rŠ   )r’   r”   rx   ru   r•   r–   Úinput_shapeÚ
seq_lengthÚbuffered_token_type_idsÚ buffered_token_type_ids_expandedr…   Ú
embeddingsrƒ   s                r(   ÚforwardzRealmEmbeddings.forward¯   sH  € ð Ð Ø#Ÿ.™.Ó*‰Kà'×,Ñ,Ó.¨s°Ð3ˆKà  ‘^ˆ
àÐØ×,Ñ,ªQÐ0FÈÐVlÑIlÐ0lÐ-lÑmˆLð
 Ð!Ü�tÐ-Ô.Ø*.×*=Ñ*=ºaÀÀ*À¸nÑ*MÐ'Ø3J×3QÑ3QÐR]Ð^_ÑR`ÐblÓ3mÐ0Ø!A‘ä!&§¡¨[ÄÇ
Á
ÐSW×SdÑSd×SkÑSkÔ!l�àÐ Ø ×0Ñ0°Ó;ˆMØ $× :Ñ :¸>Ó JÐà"Ð%:Ñ:ˆ
Ø×'Ñ'¨:Ò5Ø"&×":Ñ":¸<Ó"HÐØÐ-Ñ-ˆJØ—^‘^ JÓ/ˆ
Ø—\‘\ *Ó-ˆ
ØÐr“   )NNNNr   )rE   Ú
__module__Ú__qualname__Ú__doc__r|   r   rW   Ú
LongTensorÚFloatTensorrR   ÚTensorr¡   Ú__classcell__©rD   s   @r(   rn   rn   ™   s‹   ø„ ÙQô
ð* 15Ø59Ø37Ø59Ø&'ñ'à˜E×,Ñ,Ñ-ð'ð ! ×!1Ñ!1Ñ2ð'ð ˜u×/Ñ/Ñ0ð	'ð
   × 1Ñ 1Ñ2ð'ð !$ð'ð 
�‰÷'r“   rn   c                   óP  ‡ — e Zd Zdˆ fd„	Zdej
                  dej
                  fd„Z	 	 	 	 	 	 ddej
                  deej                     deej                     deej                     d	eej                     d
ee	e	ej                           dee
   de	ej
                     fd„Zˆ xZS )ÚRealmSelfAttentionc                 óâ  •— t         ‰| �  «        |j                  |j                  z  dk7  r2t	        |d«      s&t        d|j                  › d|j                  › d�«      ‚|j                  | _        t        |j                  |j                  z  «      | _        | j                  | j                  z  | _        t        j                  |j                  | j                  «      | _        t        j                  |j                  | j                  «      | _        t        j                  |j                  | j                  «      | _        t        j                  |j                  «      | _        |xs t#        |dd«      | _        | j$                  dk(  s| j$                  d	k(  rF|j&                  | _        t        j(                  d
|j&                  z  dz
  | j                  «      | _        |j,                  | _        y )Nr   Úembedding_sizezThe hidden size (z6) is not a multiple of the number of attention heads (ú)rs   rt   Úrelative_keyÚrelative_key_queryr1   r   )r{   r|   r   Únum_attention_headsr›   Ú
ValueErrorrR   Úattention_head_sizeÚall_head_sizer   ÚLinearÚqueryÚkeyÚvaluerˆ   Úattention_probs_dropout_probrŠ   rO   rs   r‚   r}   Údistance_embeddingÚ
is_decoder©r’   r[   rs   rD   s      €r(   r|   zRealmSelfAttention.__init__Ú   s�  ø€ Ü‰ÑÔØ×Ñ × :Ñ :Ñ:¸aÒ?ÌÐPVÐXhÔHiÜØ# F×$6Ñ$6Ð#7ð 8Ø ×4Ñ4Ð5°Qð8óð ð
 $*×#=Ñ#=ˆÔ Ü#& v×'9Ñ'9¸F×<VÑ<VÑ'VÓ#WˆÔ Ø!×5Ñ5¸×8PÑ8PÑPˆÔä—Y‘Y˜v×1Ñ1°4×3EÑ3EÓFˆŒ
Ü—9‘9˜V×/Ñ/°×1CÑ1CÓDˆŒÜ—Y‘Y˜v×1Ñ1°4×3EÑ3EÓFˆŒ
ä—z‘z &×"EÑ"EÓFˆŒØ'>ò (
Ä'ØÐ-¨zóC
ˆÔ$ð ×'Ñ'¨>Ò9¸T×=YÑ=YÐ]qÒ=qØ+1×+IÑ+IˆDÔ(Ü&(§l¡l°1°v×7UÑ7UÑ3UÐXYÑ3YÐ[_×[sÑ[sÓ&tˆDÔ#à ×+Ñ+ˆ�r“   Úxr—   c                 ó¤   — |j                  «       d d | j                  | j                  fz   }|j                  |«      }|j	                  dddd«      S )Nrv   r   r1   r   é   )r�   r±   r³   ÚviewÚpermute)r’   r½   Únew_x_shapes      r(   Útranspose_for_scoresz'RealmSelfAttention.transpose_for_scoresô   sL   € Ø—f‘f“h˜s �m t×'?Ñ'?À×AYÑAYÐ&ZÑZˆØ�F‰F�;ÓˆØ�y‰y˜˜A˜q !Ó$Ð$r“   Úhidden_statesÚattention_maskÚ	head_maskÚencoder_hidden_statesÚencoder_attention_maskÚpast_key_valueÚoutput_attentionsc                 ó$  — | j                  |«      }|d u}	|	r|�|d   }
|d   }|}�n |	rC| j                  | j                  |«      «      }
| j                  | j                  |«      «      }|}n»|�y| j                  | j                  |«      «      }
| j                  | j                  |«      «      }t	        j
                  |d   |
gd¬«      }
t	        j
                  |d   |gd¬«      }n@| j                  | j                  |«      «      }
| j                  | j                  |«      «      }| j                  |«      }|d u}| j                  r|
|f}t	        j                  ||
j                  dd«      «      }| j                  dk(  s| j                  dk(  �r—|j                  d   |
j                  d   }}|rDt	        j                  |dz
  t        j                  |j                  ¬	«      j                  dd«      }n@t	        j                  |t        j                  |j                  ¬	«      j                  dd«      }t	        j                  |t        j                  |j                  ¬	«      j                  dd«      }||z
  }| j!                  || j"                  z   dz
  «      }|j%                  |j&                  ¬
«      }| j                  dk(  rt	        j(                  d||«      }||z   }nE| j                  dk(  r6t	        j(                  d||«      }t	        j(                  d|
|«      }||z   |z   }|t+        j,                  | j.                  «      z  }|�||z   }t0        j2                  j5                  |d¬«      }| j7                  |«      }|�||z  }t	        j                  ||«      }|j9                  dddd«      j;                  «       }|j=                  «       d d | j>                  fz   }|j                  |«      }|r||fn|f}| j                  r||fz   }|S )Nr   r   r1   ©Údimrv   éþÿÿÿr¯   r°   r™   ry   zbhld,lrd->bhlrzbhrd,lrd->bhlrr¿   ) r¶   rÃ   r·   r¸   rW   Úcatr»   ÚmatmulrS   rs   rT   Útensorr�   rš   rÀ   rŒ   rº   r‚   Útorz   ÚeinsumÚmathÚsqrtr³   r   Ú
functionalÚsoftmaxrŠ   rÁ   Ú
contiguousr�   r´   )r’   rÄ   rÅ   rÆ   rÇ   rÈ   rÉ   rÊ   Úmixed_query_layerÚis_cross_attentionÚ	key_layerÚvalue_layerÚquery_layerÚ	use_cacheÚattention_scoresÚquery_lengthÚ
key_lengthÚposition_ids_lÚposition_ids_rÚdistanceÚpositional_embeddingÚrelative_position_scoresÚrelative_position_scores_queryÚrelative_position_scores_keyÚattention_probsÚcontext_layerÚnew_context_layer_shapeÚoutputss                               r(   r¡   zRealmSelfAttention.forwardù   sç  € ð !ŸJ™J }Ó5Ðð
 3¸$Ð>Ðá .Ð"<à& qÑ)ˆIØ(¨Ñ+ˆKØ3ŠNÙØ×1Ñ1°$·(±(Ð;PÓ2QÓRˆIØ×3Ñ3°D·J±JÐ?TÓ4UÓVˆKØ3‰NØÐ'Ø×1Ñ1°$·(±(¸=Ó2IÓJˆIØ×3Ñ3°D·J±J¸}Ó4MÓNˆKÜŸ	™	 >°!Ñ#4°iÐ"@ÀaÔHˆIÜŸ)™) ^°AÑ%6¸Ð$DÈ!ÔL‰Kà×1Ñ1°$·(±(¸=Ó2IÓJˆIØ×3Ñ3°D·J±J¸}Ó4MÓNˆKà×/Ñ/Ð0AÓBˆà"¨$Ð.ˆ	Ø�?Š?ð (¨Ð5ˆNô !Ÿ<™<¨°Y×5HÑ5HÈÈRÓ5PÓQÐà×'Ñ'¨>Ò9¸T×=YÑ=YÐ]qÓ=qØ'2×'8Ñ'8¸Ñ';¸Y¿_¹_ÈQÑ=O˜*ˆLÙÜ!&§¡¨j¸1©nÄEÇJÁJÐWd×WkÑWkÔ!l×!qÑ!qØ˜ó"‘ô "'§¡¨lÄ%Ç*Á*ÐUb×UiÑUiÔ!j×!oÑ!oÐprÐtuÓ!v�Ü"Ÿ\™\¨*¼E¿J¹JÈ}×OcÑOcÔd×iÑiÐjkÐmoÓpˆNØ%¨Ñ6ˆHà#'×#:Ñ#:¸8Àd×FbÑFbÑ;bÐefÑ;fÓ#gÐ Ø#7×#:Ñ#:À×ARÑARÐ#:Ó#SÐ à×+Ñ+¨~Ò=Ü+0¯<©<Ð8HÈ+ÐWkÓ+lÐ(Ø#3Ð6NÑ#NÑ Ø×-Ñ-Ð1EÒEÜ16·±Ð>NÐP[Ð]qÓ1rÐ.Ü/4¯|©|Ð<LÈiÐYmÓ/nÐ,Ø#3Ð6TÑ#TÐWsÑ#sÐ à+¬d¯i©i¸×8PÑ8PÓ.QÑQÐØÐ%à/°.Ñ@Ðô Ÿ-™-×/Ñ/Ð0@ÀbÐ/ÓIˆð Ÿ,™, Ó7ˆð Ð Ø-°	Ñ9ˆOäŸ™ _°kÓBˆà%×-Ñ-¨a°°A°qÓ9×DÑDÓFˆØ"/×"4Ñ"4Ó"6°s¸Ð";¸t×?QÑ?QÐ>SÑ"SÐØ%×*Ñ*Ð+BÓCˆá6G�= /Ñ2ÈmÐM]ˆà�?Š?Ø Ð 1Ñ1ˆGØˆr“   ©N©NNNNNF)rE   r¢   r£   r|   rW   r§   rÃ   r   r¦   r   Úboolr¡   r¨   r©   s   @r(   r«   r«   Ù   så   ø„ õ,ð4% e§l¡lð %°u·|±|ó %ð 7;Ø15Ø=AØ>BØDHØ,1ñcà—|‘|ðcð ! ×!2Ñ!2Ñ3ðcð ˜E×-Ñ-Ñ.ð	cð
  (¨×(9Ñ(9Ñ:ðcð !)¨×):Ñ):Ñ ;ðcð !  u¨U×->Ñ->Ñ'?Ñ!@ÑAðcð $ D™>ðcð 
ˆu�|‰|Ñ	÷cr“   r«   c                   ón   ‡ — e Zd Zˆ fd„Zdej
                  dej
                  dej
                  fd„Zˆ xZS )ÚRealmSelfOutputc                 ó(  •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        j                  |j                  |j                  ¬«      | _        t        j                  |j                  «      | _
        y ©Nrq   )r{   r|   r   rµ   r   Údenser†   r‡   rˆ   r‰   rŠ   r‘   s     €r(   r|   zRealmSelfOutput.__init__`  s`   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
ÜŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆŒÜ—z‘z &×"<Ñ"<Ó=ˆ�r“   rÄ   Úinput_tensorr—   c                 ór   — | j                  |«      }| j                  |«      }| j                  ||z   «      }|S rí   ©rô   rŠ   r†   ©r’   rÄ   rõ   s      r(   r¡   zRealmSelfOutput.forwardf  ó7   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆØŸ™ }°|Ñ'CÓDˆØÐr“   ©rE   r¢   r£   r|   rW   r§   r¡   r¨   r©   s   @r(   rñ   rñ   _  ó1   ø„ ô>ð U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r“   rñ   Úeagerc                   ó  ‡ — e Zd Zdˆ fd„	Zd„ Z	 	 	 	 	 	 ddej                  deej                     deej                     deej                     deej                     dee	e	ej                           d	ee
   d
e	ej                     fd„Zˆ xZS )ÚRealmAttentionc                 óž   •— t         ‰| �  «        t        |j                     ||¬«      | _        t        |«      | _        t        «       | _        y )N©rs   )	r{   r|   ÚREALM_SELF_ATTENTION_CLASSESÚ_attn_implementationr’   rñ   ÚoutputÚsetÚpruned_headsr¼   s      €r(   r|   zRealmAttention.__init__s  sC   ø€ Ü‰ÑÔÜ0°×1LÑ1LÑMØÐ,Cô
ˆŒ	ô & fÓ-ˆŒÜ›EˆÕr“   c                 ó>  — t        |«      dk(  ry t        || j                  j                  | j                  j                  | j
                  «      \  }}t        | j                  j                  |«      | j                  _        t        | j                  j                  |«      | j                  _        t        | j                  j                  |«      | j                  _	        t        | j                  j                  |d¬«      | j                  _        | j                  j                  t        |«      z
  | j                  _        | j                  j                  | j                  j                  z  | j                  _        | j
                  j                  |«      | _        y )Nr   r   rÌ   )rQ   r   r’   r±   r³   r  r   r¶   r·   r¸   r  rô   r´   Úunion)r’   ÚheadsÚindexs      r(   Úprune_headszRealmAttention.prune_heads{  s  € Üˆu‹:˜Š?ØÜ7Ø�4—9‘9×0Ñ0°$·)±)×2OÑ2OÐQU×QbÑQbó
‰ˆˆuô
 -¨T¯Y©Y¯_©_¸eÓDˆ�	‰	ŒÜ*¨4¯9©9¯=©=¸%Ó@ˆ�	‰	ŒÜ,¨T¯Y©Y¯_©_¸eÓDˆ�	‰	ŒÜ.¨t¯{©{×/@Ñ/@À%ÈQÔOˆ�‰Ôð )-¯	©	×(EÑ(EÌÈEË
Ñ(Rˆ�	‰	Ô%Ø"&§)¡)×"?Ñ"?À$Ç)Á)×B_ÑB_Ñ"_ˆ�	‰	ÔØ ×-Ñ-×3Ñ3°EÓ:ˆÕr“   rÄ   rÅ   rÆ   rÇ   rÈ   rÉ   rÊ   r—   c           	      óp   — | j                  |||||||«      }| j                  |d   |«      }	|	f|dd  z   }
|
S )Nr   r   )r’   r  )r’   rÄ   rÅ   rÆ   rÇ   rÈ   rÉ   rÊ   Úself_outputsÚattention_outputrì   s              r(   r¡   zRealmAttention.forward�  sW   € ð —y‘yØØØØ!Ø"ØØó
ˆð  Ÿ;™; |°A¡¸ÓFÐØ#Ð%¨°Q°RÐ(8Ñ8ˆØˆr“   rí   rî   )rE   r¢   r£   r|   r
  rW   r§   r   r¦   r   rï   r¡   r¨   r©   s   @r(   rþ   rþ   r  sÆ   ø„ õ"ò;ð* 7;Ø15Ø=AØ>BØDHØ,1ñà—|‘|ðð ! ×!2Ñ!2Ñ3ðð ˜E×-Ñ-Ñ.ð	ð
  (¨×(9Ñ(9Ñ:ðð !)¨×):Ñ):Ñ ;ðð !  u¨U×->Ñ->Ñ'?Ñ!@ÑAðð $ D™>ðð 
ˆu�|‰|Ñ	÷r“   rþ   c                   óV   ‡ — e Zd Zˆ fd„Zdej
                  dej
                  fd„Zˆ xZS )ÚRealmIntermediatec                 ó  •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        |j                  t        «      rt        |j                     | _        y |j                  | _        y rí   )r{   r|   r   rµ   r   Úintermediate_sizerô   rB   Ú
hidden_actÚstrr
   Úintermediate_act_fnr‘   s     €r(   r|   zRealmIntermediate.__init__¦  s]   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3KÑ3KÓLˆŒ
Ü�f×'Ñ'¬Ô-Ü'-¨f×.?Ñ.?Ñ'@ˆDÕ$à'-×'8Ñ'8ˆDÕ$r“   rÄ   r—   c                 óJ   — | j                  |«      }| j                  |«      }|S rí   )rô   r  ©r’   rÄ   s     r(   r¡   zRealmIntermediate.forward®  s&   € ØŸ
™
 =Ó1ˆØ×0Ñ0°Ó?ˆØÐr“   rú   r©   s   @r(   r  r  ¥  s#   ø„ ô9ð U§\¡\ð °e·l±l÷ r“   r  c                   ón   ‡ — e Zd Zˆ fd„Zdej
                  dej
                  dej
                  fd„Zˆ xZS )ÚRealmOutputc                 ó(  •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        j                  |j
                  |j                  ¬«      | _        t        j                  |j                  «      | _        y ró   )r{   r|   r   rµ   r  r   rô   r†   r‡   rˆ   r‰   rŠ   r‘   s     €r(   r|   zRealmOutput.__init__µ  s`   ø€ Ü‰ÑÔÜ—Y‘Y˜v×7Ñ7¸×9KÑ9KÓLˆŒ
ÜŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆŒÜ—z‘z &×"<Ñ"<Ó=ˆ�r“   rÄ   rõ   r—   c                 ór   — | j                  |«      }| j                  |«      }| j                  ||z   «      }|S rí   r÷   rø   s      r(   r¡   zRealmOutput.forward»  rù   r“   rú   r©   s   @r(   r  r  ´  rû   r“   r  c                   ó  ‡ — e Zd Zˆ fd„Z	 	 	 	 	 	 ddej
                  deej                     deej                     deej                     deej                     deeeej                           dee	   d	eej
                     fd
„Z
d„ Zˆ xZS )Ú
RealmLayerc                 óf  •— t         ‰| �  «        |j                  | _        d| _        t	        |«      | _        |j                  | _        |j                  | _        | j                  r,| j                  st        | › d�«      ‚t	        |d¬«      | _	        t        |«      | _        t        |«      | _        y )Nr   z> should be used as a decoder model if cross attention is addedrt   r   )r{   r|   Úchunk_size_feed_forwardÚseq_len_dimrþ   Ú	attentionr»   Úadd_cross_attentionr²   Úcrossattentionr  Úintermediater  r  r‘   s     €r(   r|   zRealmLayer.__init__Ã  s—   ø€ Ü‰ÑÔØ'-×'EÑ'EˆÔ$ØˆÔÜ'¨Ó/ˆŒØ ×+Ñ+ˆŒØ#)×#=Ñ#=ˆÔ Ø×#Ò#Ø—?’?Ü  D 6Ð)gÐ!hÓiÐiÜ"0°ÐQ[Ô"\ˆDÔÜ-¨fÓ5ˆÔÜ! &Ó)ˆ�r“   rÄ   rÅ   rÆ   rÇ   rÈ   rÉ   rÊ   r—   c           	      óÒ  — |�|d d nd }| j                  |||||¬«      }	|	d   }
| j                  r|	dd }|	d   }n|	dd  }d }| j                  rT|�Rt        | d«      st        d| › d�«      ‚|�|d	d  nd }| j	                  |
||||||«      }|d   }
||dd z   }|d   }|z   }t        | j                  | j                  | j                  |
«      }|f|z   }| j                  r|fz   }|S )
Nr1   )rÊ   rÉ   r   r   rv   r"  z'If `encoder_hidden_states` are passed, z` has to be instantiated with cross-attention layers by setting `config.add_cross_attention=True`rÎ   )	r   r»   r›   r²   r"  r   Úfeed_forward_chunkr  r  )r’   rÄ   rÅ   rÆ   rÇ   rÈ   rÉ   rÊ   Úself_attn_past_key_valueÚself_attention_outputsr  rì   Úpresent_key_valueÚcross_attn_present_key_valueÚcross_attn_past_key_valueÚcross_attention_outputsÚlayer_outputs                    r(   r¡   zRealmLayer.forwardÑ  s}  € ð :HÐ9S >°"°1Ñ#5ÐY]Ð Ø!%§¡ØØØØ/Ø3ð "0ó "
Ðð 2°!Ñ4Ðð �?Š?Ø,¨Q¨rÐ2ˆGØ 6°rÑ :Ñà,¨Q¨RÐ0ˆGà'+Ð$Ø�?Š?Ð4Ð@Ü˜4Ð!1Ô2Ü Ø=¸d¸Vð DDð Dóð ð @NÐ?Y¨°r°sÑ(;Ð_cÐ%Ø&*×&9Ñ&9Ø ØØØ%Ø&Ø)Ø!ó'Ð#ð  7°qÑ9ÐØÐ 7¸¸"Ð =Ñ=ˆGð ,CÀ2Ñ+FÐ(Ø 1Ð4PÑ PÐä0Ø×#Ñ# T×%AÑ%AÀ4×CSÑCSÐUeó
ˆð  �/ GÑ+ˆð �?Š?ØÐ!2Ð 4Ñ4ˆGàˆr“   c                 óL   — | j                  |«      }| j                  ||«      }|S rí   )r#  r  )r’   r  Úintermediate_outputr,  s       r(   r%  zRealmLayer.feed_forward_chunk  s,   € Ø"×/Ñ/Ð0@ÓAÐØ—{‘{Ð#6Ð8HÓIˆØÐr“   rî   )rE   r¢   r£   r|   rW   r§   r   r¦   r   rï   r¡   r%  r¨   r©   s   @r(   r  r  Â  sÇ   ø„ ô*ð" 7;Ø15Ø=AØ>BØDHØ,1ñ?à—|‘|ð?ð ! ×!2Ñ!2Ñ3ð?ð ˜E×-Ñ-Ñ.ð	?ð
  (¨×(9Ñ(9Ñ:ð?ð !)¨×):Ñ):Ñ ;ð?ð !  u¨U×->Ñ->Ñ'?Ñ!@ÑAð?ð $ D™>ð?ð 
ˆu�|‰|Ñ	ó?öBr“   r  c                   óD  ‡ — e Zd Zˆ fd„Z	 	 	 	 	 	 	 	 	 ddej
                  deej                     deej                     deej                     deej                     deeeej                           dee	   d	ee	   d
ee	   dee	   de
eej
                     ef   fd„Zˆ xZS )ÚRealmEncoderc                 óÐ   •— t         ‰| �  «        || _        t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        d| _	        y c c}w )NF)
r{   r|   r[   r   Ú
ModuleListÚrangeÚnum_hidden_layersr  ÚlayerÚgradient_checkpointing)r’   r[   Ú_rD   s      €r(   r|   zRealmEncoder.__init__  sN   ø€ Ü‰ÑÔØˆŒÜ—]‘]ÄÀf×F^ÑF^Ó@_Ö#`¸1¤J¨vÕ$6Ò#`ÓaˆŒ
Ø&+ˆÕ#ùò $as   ½A#rÄ   rÅ   rÆ   rÇ   rÈ   Úpast_key_valuesrÞ   rÊ   Úoutput_hidden_statesÚreturn_dictr—   c                 óš  — |	rdnd }|rdnd }|r| j                   j                  rdnd }| j                  r%| j                  r|rt        j                  d«       d}|rdnd }t        | j                  «      D ]¤  \  }}|	r||fz   }|�||   nd }|�||   nd }| j                  r/| j                  r#| j                  |j                  |||||||«      }n ||||||||«      }|d   }|r	||d   fz  }|sŒ|||d   fz   }| j                   j                  sŒœ||d   fz   }Œ¦ |	r||fz   }|
st        d„ |||||fD «       «      S t        |||||¬	«      S )
Nr%   zZ`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`...Fr   rv   r   r1   c              3   ó$   K  — | ]  }|�|–— Œ
 y ­wrí   r%   )r&   Úvs     r(   r)   z'RealmEncoder.forward.<locals>.<genexpr>a  s   è ø€ ò 
àð �=ô ñ
ùr*   )Úlast_hidden_stater8  rÄ   Ú
attentionsÚcross_attentions)r[   r!  r6  Útrainingr7   Úwarning_onceÚ	enumerater5  Ú_gradient_checkpointing_funcÚ__call__Útupler   )r’   rÄ   rÅ   rÆ   rÇ   rÈ   r8  rÞ   rÊ   r9  r:  Úall_hidden_statesÚall_self_attentionsÚall_cross_attentionsÚnext_decoder_cacheÚiÚlayer_moduleÚlayer_head_maskrÉ   Úlayer_outputss                       r(   r¡   zRealmEncoder.forward  sÎ  € ñ #7™B¸DÐÙ$5™b¸4ÐÙ%6¸4¿;¹;×;ZÒ;Z™rÐ`dÐà×&Ò&¨4¯=ª=ÙÜ×#Ñ#Øpôð "�	á#,™R°$ÐÜ(¨¯©Ó4ò #	V‰OˆAˆ|Ù#Ø$5¸Ð8HÑ$HÐ!à.7Ð.C˜i¨šlÈˆOØ3BÐ3N˜_¨QÒ/ÐTXˆNà×*Ò*¨t¯}ª}Ø $× AÑ AØ ×)Ñ)Ø!Ø"Ø#Ø)Ø*Ø"Ø%ó	!‘ñ !-Ø!Ø"Ø#Ø)Ø*Ø"Ø%ó!�ð *¨!Ñ,ˆMÙØ" }°RÑ'8Ð&:Ñ:Ð"Ú Ø&9¸]È1Ñ=MÐ<OÑ&OÐ#Ø—;‘;×2Ó2Ø+?À=ÐQRÑCSÐBUÑ+UÑ(ðG#	VñJ  Ø 1°]Ð4DÑ DÐáÜñ 
ð "Ø&Ø%Ø'Ø(ðô
ó 
ð 
ô 9Ø+Ø.Ø+Ø*Ø1ô
ð 	
r“   )	NNNNNNFFT)rE   r¢   r£   r|   rW   r§   r   r¦   r   rï   r   r   r¡   r¨   r©   s   @r(   r0  r0    s  ø„ ô,ð 7;Ø15Ø=AØ>BØEIØ$(Ø,1Ø/4Ø&*ñS
à—|‘|ðS
ð ! ×!2Ñ!2Ñ3ðS
ð ˜E×-Ñ-Ñ.ð	S
ð
  (¨×(9Ñ(9Ñ:ðS
ð !)¨×):Ñ):Ñ ;ðS
ð " %¨¨e×.?Ñ.?Ñ(@Ñ"AÑBðS
ð ˜D‘>ðS
ð $ D™>ðS
ð ' t™nðS
ð ˜d‘^ðS
ð 
ˆu�U—\‘\Ñ"Ð$MÐMÑ	N÷S
r“   r0  c                   óV   ‡ — e Zd Zˆ fd„Zdej
                  dej
                  fd„Zˆ xZS )ÚRealmPoolerc                 ó²   •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        j                  «       | _        y rí   )r{   r|   r   rµ   r   rô   ÚTanhÚ
activationr‘   s     €r(   r|   zRealmPooler.__init__v  s9   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
ÜŸ'™'›)ˆ�r“   rÄ   r—   c                 ó\   — |d d …df   }| j                  |«      }| j                  |«      }|S )Nr   )rô   rS  )r’   rÄ   Úfirst_token_tensorÚpooled_outputs       r(   r¡   zRealmPooler.forward{  s6   € ð +ª1¨a¨4Ñ0ÐØŸ
™
Ð#5Ó6ˆØŸ™¨Ó6ˆØÐr“   rú   r©   s   @r(   rP  rP  u  s#   ø„ ô$ð
 U§\¡\ð °e·l±l÷ r“   rP  c                   ó–   — e Zd ZU dZdZeej                     ed<   dZ	ee
ej                        ed<   dZee
ej                        ed<   y)ÚRealmEmbedderOutputa*  
    Outputs of [`RealmEmbedder`] models.

    Args:
        projected_score (`torch.FloatTensor` of shape `(batch_size, config.retriever_proj_size)`):

            Projected score.
        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 + 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 initial embedding outputs.
        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)`.

            Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
            heads.
    NÚprojected_scorerÄ   r?  )rE   r¢   r£   r¤   rY  r   rW   r¦   Ú__annotations__rÄ   r   r?  r%   r“   r(   rX  rX  „  sR   … ñð( 48€O�X˜e×/Ñ/Ñ0Ó7Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ô9r“   rX  c                   óŠ   — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eej                     ed<   y)ÚRealmScorerOutputa'  
    Outputs of [`RealmScorer`] models.

    Args:
        relevance_score (`torch.FloatTensor` of shape `(batch_size, config.num_candidates)`):
            The relevance score of document candidates (before softmax).
        query_score (`torch.FloatTensor` of shape `(batch_size, config.retriever_proj_size)`):
            Query score derived from the query embedder.
        candidate_score (`torch.FloatTensor` of shape `(batch_size, config.num_candidates, config.retriever_proj_size)`):
            Candidate score derived from the embedder.
    NÚrelevance_scoreÚquery_scoreÚcandidate_score)rE   r¢   r£   r¤   r]  r   rW   r¦   rZ  r^  r_  r%   r“   r(   r\  r\  Ÿ  sH   … ñ
ð 48€O�X˜e×/Ñ/Ñ0Ó7Ø/3€K�˜%×+Ñ+Ñ,Ó3Ø37€O�X˜e×/Ñ/Ñ0Ô7r“   r\  c                   ó¾  — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eej                     ed<   dZej                  ed<   dZej                  ed<   dZeej                     ed<   dZeej                     ed	<   dZej$                  ed
<   dZej$                  ed<   dZeeej                        ed<   dZeeej                        ed<   y)ÚRealmReaderOutputa+	  
    Outputs of [`RealmReader`] models.

    Args:
        loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `start_positions`, `end_positions`, `has_answers` are provided):
            Total loss.
        retriever_loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `start_positions`, `end_positions`, `has_answers` are provided):
            Retriever loss.
        reader_loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `start_positions`, `end_positions`, `has_answers` are provided):
            Reader loss.
        retriever_correct (`torch.BoolTensor` of shape `(config.searcher_beam_size,)`, *optional*):
            Whether or not an evidence block contains answer.
        reader_correct (`torch.BoolTensor` of shape `(config.reader_beam_size, num_candidates)`, *optional*):
            Whether or not a span candidate contains answer.
        block_idx (`torch.LongTensor` of shape `()`):
            The index of the retrieved evidence block in which the predicted answer is most likely.
        candidate (`torch.LongTensor` of shape `()`):
            The index of the retrieved span candidates in which the predicted answer is most likely.
        start_pos (`torch.IntTensor` of shape `()`):
            Predicted answer starting position in *RealmReader*'s inputs.
        end_pos (`torch.IntTensor` of shape `()`):
            Predicted answer ending position in *RealmReader*'s inputs.
        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 + 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 initial embedding outputs.
        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)`.

            Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
            heads.
    NÚlossÚretriever_lossÚreader_lossÚretriever_correctÚreader_correctÚ	block_idxÚ	candidateÚ	start_posÚend_posrÄ   r?  )rE   r¢   r£   r¤   rb  r   rW   r¦   rZ  rc  rd  re  Ú
BoolTensorrf  rg  r¥   rh  ri  Úint32rj  rÄ   r   r?  r%   r“   r(   ra  ra  ²  sä   … ñ!ðF )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø26€N�H˜U×.Ñ.Ñ/Ó6Ø/3€K�˜%×+Ñ+Ñ,Ó3Ø*.Ð�u×'Ñ'Ó.Ø'+€N�E×$Ñ$Ó+Ø,0€Iˆx˜×(Ñ(Ñ)Ó0Ø,0€Iˆx˜×(Ñ(Ñ)Ó0Ø!€Iˆu�{‰{Ó!Ø€GˆU�[‰[ÓØ8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ô9r“   ra  c                   óH   — e Zd ZU dZdZeed<   dZee	j                     ed<   y)ÚRealmForOpenQAOutputzï

    Outputs of [`RealmForOpenQA`] models.

    Args:
        reader_output (`dict`):
            Reader output.
        predicted_answer_ids (`torch.LongTensor` of shape `(answer_sequence_length)`):
            Predicted answer ids.
    NÚreader_outputÚpredicted_answer_ids)rE   r¢   r£   r¤   ro  ÚdictrZ  rp  r   rW   r¥   r%   r“   r(   rn  rn  ä  s)   … ñ	ð €M�4ÓØ7;Ð˜( 5×#3Ñ#3Ñ4Ô;r“   rn  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚRealmPredictionHeadTransformc                 óh  •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        |j                  t        «      rt        |j                     | _
        n|j                  | _
        t        j                  |j                  |j                  ¬«      | _        y ró   )r{   r|   r   rµ   r   rô   rB   r  r  r
   Útransform_act_fnr†   r‡   r‘   s     €r(   r|   z%RealmPredictionHeadTransform.__init__ö  s{   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
Ü�f×'Ñ'¬Ô-Ü$*¨6×+<Ñ+<Ñ$=ˆDÕ!à$*×$5Ñ$5ˆDÔ!ÜŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆ�r“   c                 ól   — | j                  |«      }| j                  |«      }| j                  |«      }|S rí   )rô   ru  r†   r  s     r(   r¡   z$RealmPredictionHeadTransform.forwardÿ  s4   € ØŸ
™
 =Ó1ˆØ×-Ñ-¨mÓ<ˆØŸ™ }Ó5ˆØÐr“   ©rE   r¢   r£   r|   r¡   r¨   r©   s   @r(   rs  rs  õ  s   ø„ ôUör“   rs  c                   ó*   ‡ — e Zd Zˆ fd„Zd„ Zd„ Zˆ xZS )ÚRealmLMPredictionHeadc                 óH  •— t         ‰| �  «        t        |«      | _        t	        j
                  |j                  |j                  d¬«      | _        t	        j                  t        j                  |j                  «      «      | _        | j                  | j                  _        y )NF)r0   )r{   r|   rs  Ú	transformr   rµ   r   r~   ÚdecoderÚ	ParameterrW   rŽ   r0   r‘   s     €r(   r|   zRealmLMPredictionHead.__init__  sm   ø€ Ü‰ÑÔÜ5°fÓ=ˆŒô —y‘y ×!3Ñ!3°V×5FÑ5FÈUÔSˆŒä—L‘L¤§¡¨V×->Ñ->Ó!?Ó@ˆŒ	ð !ŸI™Iˆ�‰Õr“   c                 ó:   — | j                   | j                  _         y rí   )r0   r|  ©r’   s    r(   Ú_tie_weightsz"RealmLMPredictionHead._tie_weights  s   € Ø ŸI™Iˆ�‰Õr“   c                 óJ   — | j                  |«      }| j                  |«      }|S rí   )r{  r|  r  s     r(   r¡   zRealmLMPredictionHead.forward  s$   € ØŸ™ }Ó5ˆØŸ™ ]Ó3ˆØÐr“   )rE   r¢   r£   r|   r€  r¡   r¨   r©   s   @r(   ry  ry    s   ø„ ô&ò&ör“   ry  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚRealmOnlyMLMHeadc                 óB   •— t         ‰| �  «        t        |«      | _        y rí   )r{   r|   ry  Úpredictionsr‘   s     €r(   r|   zRealmOnlyMLMHead.__init__  s   ø€ Ü‰ÑÔÜ0°Ó8ˆÕr“   c                 ó(   — | j                  |«      }|S rí   )r…  )r’   Úsequence_outputÚprediction_scoress      r(   r¡   zRealmOnlyMLMHead.forward"  s   € Ø ×,Ñ,¨_Ó=ÐØ Ð r“   rw  r©   s   @r(   rƒ  rƒ    s   ø„ ô9ö!r“   rƒ  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚRealmScorerProjectionc                 ó   •— t         ‰| �  «        t        |«      | _        t	        j
                  |j                  |j                  «      | _        t	        j                  |j                  |j                  ¬«      | _	        y ró   )r{   r|   ry  r…  r   rµ   r   Úretriever_proj_sizerô   r†   r‡   r‘   s     €r(   r|   zRealmScorerProjection.__init__(  sW   ø€ Ü‰ÑÔÜ0°Ó8ˆÔÜ—Y‘Y˜v×1Ñ1°6×3MÑ3MÓNˆŒ
ÜŸ™ f×&@Ñ&@Àf×F[ÑF[Ô\ˆ�r“   c                 óJ   — | j                  |«      }| j                  |«      }|S rí   )rô   r†   r  s     r(   r¡   zRealmScorerProjection.forward.  s$   € ØŸ
™
 =Ó1ˆØŸ™ }Ó5ˆØÐr“   rw  r©   s   @r(   rŠ  rŠ  '  s   ø„ ô]ör“   rŠ  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚRealmReaderProjectionc                 óp  •— t         ‰| �  «        || _        t        j                  |j
                  |j                  dz  «      | _        t        j                  |j                  d«      | _        t        j                  |j                  |j                  ¬«      | _        t        j                  «       | _        y )Nr1   r   rq   )r{   r|   r[   r   rµ   r   Úspan_hidden_sizeÚdense_intermediateÚdense_outputr†   Úreader_layer_norm_epsÚlayer_normalizationÚReLUÚrelur‘   s     €r(   r|   zRealmReaderProjection.__init__5  s   ø€ Ü‰ÑÔØˆŒÜ"$§)¡)¨F×,>Ñ,>À×@WÑ@WÐZ[Ñ@[Ó"\ˆÔÜŸI™I f×&=Ñ&=¸qÓAˆÔÜ#%§<¡<°×0GÑ0GÈV×MiÑMiÔ#jˆÔ Ü—G‘G“Iˆ�	r“   c                 óÀ  ‡ — ˆ fd„}t         j                  fd„}‰ j                  |«      }|j                  dd¬«      \  }} ||«      \  }}}	t        j                  |d|¬«      }
t        j                  |d|¬«      }|
|z   }‰ j                  |«      }‰ j                  |«      }‰ j                  |«      j                  d«      }| ||	|j                  ¬«      z  }|||fS )	Nc                 ób  •‡ ‡‡— ‰ j                   \  }Šˆ ˆfd„Št        ˆfd„t        ‰	j                  j                  «      D «       Ž \  }}t        j                  |d«      }t        j                  |d«      }t        j                  ‰ d|¬«      }t        j                  ‰ d|¬«      }||z  }|||fS )aK  
            Generate span candidates.

            Args:
                masks: <bool> [num_retrievals, max_sequence_len]

            Returns:
                starts: <int32> [num_spans] ends: <int32> [num_spans] span_masks: <int32> [num_retrievals, num_spans]
                whether spans locate in evidence block.
            c                 ó¤   •— t        j                  ‰| z
  dz   ‰j                  ¬«      }t        j                  | dz
  ‰‰j                  ¬«      }||fS )Nr   ©rš   )rW   rŒ   rš   )ÚwidthÚcurrent_startsÚcurrent_endsÚmasksÚmax_sequence_lens      €€r(   Ú_spans_given_widthzRRealmReaderProjection.forward.<locals>.span_candidates.<locals>._spans_given_widthK  sN   ø€ Ü!&§¡Ð.>ÀÑ.FÈÑ.JÐSX×S_ÑS_Ô!`�Ü$Ÿ|™|¨E°A©IÐ7GÐPU×P\ÑP\Ô]�Ø% |Ð3Ð3r“   c              3   ó4   •K  — | ]  } ‰|d z   «      –— Œ y­w)r   Nr%   )r&   Úwr¡  s     €r(   r)   zIRealmReaderProjection.forward.<locals>.span_candidates.<locals>.<genexpr>P  s   øè ø€ Ò f¸qÑ!3°A¸±E×!:Ñ fùs   ƒr   rv   ©rÍ   r	  )rT   rA   r3  r[   Úmax_span_widthrW   rÏ   Úindex_select)
rŸ  r7  ÚstartsÚendsÚstart_masksÚ	end_masksÚ
span_masksr¡  r   r’   s
   `      @@€r(   Úspan_candidatesz6RealmReaderProjection.forward.<locals>.span_candidates>  s¡   û€ ð #(§+¡+ÑˆAÐõ4ô
 Ó fÄEÈ$Ï+É+×JdÑJdÓDeÔ fÐg‰LˆF�Dô —Y‘Y˜v qÓ)ˆFÜ—9‘9˜T 1Ó%ˆDô  ×,Ñ,¨U¸À&ÔIˆKÜ×*Ñ*¨5°bÀÔEˆIØ$ yÑ0ˆJà˜4 Ð+Ð+r“   c                 ój   — d| j                  |«      z
  t        j                  |«      j                  z  S ©Nç      ð?©ÚtyperW   ÚfinfoÚmin©Úmaskrz   s     r(   Úmask_to_scorez4RealmReaderProjection.forward.<locals>.mask_to_score]  s*   € Ø˜$Ÿ)™) EÓ*Ñ*¬e¯k©k¸%Ó.@×.DÑ.DÑDÐDr“   r1   rv   rÌ   r   r¤  ry   )
rW   Úfloat32r’  Úchunkr¦  r—  r•  r“  Úsqueezerz   )r’   rÄ   Ú
block_maskr¬  r¶  Ústart_projectionÚend_projectionÚcandidate_startsÚcandidate_endsÚcandidate_maskÚcandidate_start_projectionsÚcandidate_end_projectionsÚcandidate_hiddenÚreader_logitss   `             r(   r¡   zRealmReaderProjection.forward=  sô   ø€ ô	,ô> ',§m¡mó 	Eð ×/Ñ/°Ó>ˆà+8×+>Ñ+>¸qÀbÐ+>Ó+IÑ(Ð˜.á;JÈ:Ó;VÑ8Ð˜.¨.ä&+×&8Ñ&8Ð9IÈqÐXhÔ&iÐ#Ü$)×$6Ñ$6°~È1ÐTbÔ$cÐ!Ø6Ð9RÑRÐð  Ÿ9™9Ð%5Ó6Ðà×3Ñ3Ð4DÓEÐà×)Ñ)Ð*:Ó;×CÑCÀBÓGˆà™ ~¸]×=PÑ=PÔQÑQˆàÐ.°Ð>Ð>r“   rw  r©   s   @r(   r�  r�  4  s   ø„ ôö7?r“   r�  aH  
    This model is a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) sub-class. Use
    it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage and
    behavior.

    Parameters:
        config ([`RealmConfig`]): Model configuration class with all the parameters of the model.
            Initializing with a config file does not load the weights associated with the model, only the
            configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.
a5
  
    Args:
        input_ids (`torch.LongTensor` of shape `({0})`):
            Indices of input sequence tokens in the vocabulary.

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

            [What are input IDs?](../glossary#input-ids)
        attention_mask (`torch.FloatTensor` of shape `({0})`, *optional*):
            Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:

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

            [What are attention masks?](../glossary#attention-mask)
        token_type_ids (`torch.LongTensor` of shape `({0})`, *optional*):
            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,
            1]`:

            - 0 corresponds to a *sentence A* token,
            - 1 corresponds to a *sentence B* token.

            [What are token type IDs?](../glossary#token-type-ids)
        position_ids (`torch.LongTensor` of shape `({0})`, *optional*):
            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.

            [What are position IDs?](../glossary#position-ids)
        head_mask (`torch.FloatTensor` of shape `(num_heads,)` or `(num_layers, num_heads)`, *optional*):
            Mask to nullify selected heads of the self-attention modules. Mask values selected in `[0, 1]`:

            - 1 indicates the head is **not masked**,
            - 0 indicates the head is **masked**.

        inputs_embeds (`torch.FloatTensor` of shape `({0}, hidden_size)`, *optional*):
            Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
            is useful if you want more control over how to convert *input_ids* indices into associated vectors than the
            model's internal embedding lookup matrix.
        output_attentions (`bool`, *optional*):
            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
            tensors for more detail.
        output_hidden_states (`bool`, *optional*):
            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
            more detail.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
c                   ó(   — e Zd ZdZeZeZdZd„ Z	d„ Z
y)ÚRealmPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    Úrealmc                 ó  — t        |t        j                  «      rm|j                  j                  j                  d| j                  j                  ¬«       |j                  �%|j                  j                  j                  «        yyt        |t        j                  «      rz|j                  j                  j                  d| j                  j                  ¬«       |j                  �2|j                  j                  |j                     j                  «        yyt        |t        j                  «      rJ|j                  j                  j                  «        |j                  j                  j                  d«       yy)zInitialize the weightsg        )ÚmeanÚstdNr¯  )rB   r   rµ   r-   rY   Únormal_r[   Úinitializer_ranger0   Úzero_r}   rp   r†   Úfill_)r’   Úmodules     r(   Ú_init_weightsz"RealmPreTrainedModel._init_weights¾  s  € ä�fœbŸi™iÔ(ð �M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ�{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ä˜¤§¡Ô-Ø�M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ×!Ñ!Ð-Ø—‘×"Ñ" 6×#5Ñ#5Ñ6×<Ñ<Õ>ð .ä˜¤§¡Ô-Ø�K‰K×Ñ×"Ñ"Ô$Ø�M‰M×Ñ×$Ñ$ SÕ)ð .r“   c                 óÂ   — g }|D ]W  }|€|j                  d«       Œ|j                  }t        |«      dkD  r|j                  d|d   f«      }|j                  |«       ŒY |S )z.Flatten inputs' shape to (-1, input_shape[-1])Nr1   rv   )r@   rT   rQ   rÀ   )r’   ÚinputsÚflattened_inputsrÑ   rœ   s        r(   Ú_flatten_inputsz$RealmPreTrainedModel._flatten_inputsÎ  sm   € àÐØò 	0ˆFØˆ~Ø ×'Ñ'¨Õ-à$Ÿl™l�Ü�{Ó# aÒ'Ø#Ÿ[™[¨"¨k¸"©oÐ)>Ó?�FØ ×'Ñ'¨Õ/ð	0ð  Ðr“   N)rE   r¢   r£   r¤   r   Úconfig_classrl   Úload_tf_weightsÚbase_model_prefixrÏ  rÓ  r%   r“   r(   rÅ  rÅ  ´  s#   „ ñð
 €LØ.€OØÐò*ó  r“   rÅ  c                   óX   ‡ — e Zd ZdZdˆ fd„	Zd„ Zd„ Zd„ Z	 	 	 	 	 	 	 	 	 	 	 	 	 dd„Zˆ xZ	S )	ÚRealmBertModelz?
    Same as the original BertModel but remove docstrings.
    c                 óº   •— t         ‰| �  |«       || _        t        |«      | _        t        |«      | _        |rt        |«      nd | _        | j                  «        y rí   )
r{   r|   r[   rn   r    r0  ÚencoderrP  ÚpoolerÚ	post_init)r’   r[   Úadd_pooling_layerrD   s      €r(   r|   zRealmBertModel.__init__á  sK   ø€ Ü‰Ñ˜Ô ØˆŒä)¨&Ó1ˆŒÜ# FÓ+ˆŒá->”k &Ô)ÀDˆŒð 	�‰Õr“   c                 ó.   — | j                   j                  S rí   ©r    r�   r  s    r(   Úget_input_embeddingsz#RealmBertModel.get_input_embeddingsî  s   € Ø�‰×.Ñ.Ð.r“   c                 ó&   — || j                   _        y rí   rß  ©r’   r¸   s     r(   Úset_input_embeddingsz#RealmBertModel.set_input_embeddingsñ  s   € Ø*/ˆ�‰Õ'r“   c                 ó˜   — |j                  «       D ]7  \  }}| j                  j                  |   j                  j	                  |«       Œ9 y)z�
        Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer} See base
        class PreTrainedModel
        N)ÚitemsrÚ  r5  r   r
  )r’   Úheads_to_pruner5  r  s       r(   Ú_prune_headszRealmBertModel._prune_headsô  sE   € ð
 +×0Ñ0Ó2ò 	C‰LˆE�5Ø�L‰L×Ñ˜uÑ%×/Ñ/×;Ñ;¸EÕBñ	Cr“   c                 óœ  — |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }| j                   j                  r|
�|
n| j                   j
                  }
nd}
|�|�t        d«      ‚|�#| j                  ||«       |j                  «       }n!|�|j                  «       d d }nt        d«      ‚|\  }}|�|j                  n|j                  }|	�|	d   d   j                  d   nd}|€t        j                  |||z   f|¬«      }|€pt        | j                  d«      r4| j                  j                  d d …d |…f   }|j!                  ||«      }|}n&t        j"                  |t        j$                  |¬	«      }| j'                  ||«      }| j                   j                  rE|�C|j                  «       \  }}}||f}|€t        j                  ||¬«      }| j)                  |«      }nd }| j+                  || j                   j,                  «      }| j                  |||||¬
«      }| j/                  ||||||	|
|||¬«
      }|d   }| j0                  �| j1                  |«      nd }|s
||f|dd  z   S t3        |||j4                  |j6                  |j8                  |j:                  ¬«      S )NFzDYou cannot specify both input_ids and inputs_embeds at the same timerv   z5You have to specify either input_ids or inputs_embedsr   r1   r›  rx   r™   )r”   ru   rx   r•   r–   )	rÅ   rÆ   rÇ   rÈ   r8  rÞ   rÊ   r9  r:  r   )r>  Úpooler_outputr8  rÄ   r?  r@  )r[   rÊ   r9  Úuse_return_dictr»   rÞ   r²   Ú%warn_if_padding_and_no_attention_maskr�   rš   rT   rW   Úonesr›   r    rx   r�   rŽ   r�   Úget_extended_attention_maskÚinvert_attention_maskÚget_head_maskr4  rÚ  rÛ  r   r8  rÄ   r?  r@  )r’   r”   rÅ   rx   ru   rÆ   r•   rÇ   rÈ   r8  rÞ   rÊ   r9  r:  rœ   Ú
batch_sizer�   rš   r–   rž   rŸ   Úextended_attention_maskÚencoder_batch_sizeÚencoder_sequence_lengthr7  Úencoder_hidden_shapeÚencoder_extended_attention_maskÚembedding_outputÚencoder_outputsr‡  rV  s                                  r(   r¡   zRealmBertModel.forwardü  s  € ð  2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà�;‰;×!Ò!Ø%.Ð%:™	ÀÇÁ×@UÑ@U‰IàˆIàÐ  ]Ð%>ÜÐcÓdÐdØÐ"Ø×6Ñ6°yÀ.ÔQØ#Ÿ.™.Ó*‰KØÐ&Ø'×,Ñ,Ó.¨s°Ð3‰KäÐTÓUÐUà!,Ñˆ
�JØ%.Ð%:�×!Ò!À×@TÑ@Tˆð DSÐC^ °Ñ!3°AÑ!6×!<Ñ!<¸QÒ!?ÐdeÐàÐ!Ü"ŸZ™Z¨*°jÐCYÑ6YÐ)ZÐdjÔkˆNàÐ!Ü�t—‘Ð(8Ô9Ø*.¯/©/×*HÑ*HÊÈKÈZÈKÈÑ*XÐ'Ø3J×3QÑ3QÐR\Ð^hÓ3iÐ0Ø!A‘ä!&§¡¨[ÄÇ
Á
ÐSYÔ!Z�ð 15×0PÑ0PÐQ_ÐalÓ0mÐð �;‰;×!Ò!Ð&;Ð&GØ=R×=WÑ=WÓ=YÑ:ÐÐ 7¸Ø$6Ð8OÐ#PÐ Ø%Ð-Ü).¯©Ð4HÐQWÔ)XÐ&Ø.2×.HÑ.HÐI_Ó.`Ñ+à.2Ð+ð ×&Ñ& y°$·+±+×2OÑ2OÓPˆ	àŸ?™?ØØ%Ø)Ø'Ø#9ð +ó 
Ðð Ÿ,™,ØØ2ØØ"7Ø#BØ+ØØ/Ø!5Ø#ð 'ó 
ˆð *¨!Ñ,ˆØ8<¿¹Ð8O˜Ÿ™ OÔ4ÐUYˆáØ# ]Ð3°oÀaÀbÐ6IÑIÐIä;Ø-Ø'Ø+×;Ñ;Ø)×7Ñ7Ø&×1Ñ1Ø,×=Ñ=ô
ð 	
r“   )T©NNNNNNNNNNNNN)
rE   r¢   r£   r¤   r|   rà  rã  rç  r¡   r¨   r©   s   @r(   rØ  rØ  Ü  sL   ø„ ñõò/ò0òCð ØØØØØØ"Ø#ØØØØ!Ø÷l
r“   rØ  z`The embedder of REALM outputting projected score that will be used to calculate relevance score.c                   óz  ‡ — e Zd ZdgZˆ fd„Zd„ Zd„ Z eej                  d«      «       e
ee¬«      	 	 	 	 	 	 	 	 	 ddeej                     deej                      d	eej                     d
eej                     deej                      deej                      dee   dee   dee   deeef   fd„«       «       Zˆ xZS )rJ   zcls.predictions.decoder.biasc                 ó¬   •— t         ‰| �  |«       t        | j                  «      | _        t        | j                  «      | _        | j                  «        y rí   )r{   r|   rØ  r[   rÆ  rŠ  r   rÜ  r‘   s     €r(   r|   zRealmEmbedder.__init__r  s:   ø€ Ü‰Ñ˜Ô ä# D§K¡KÓ0ˆŒ
Ü(¨¯©Ó5ˆŒØ�‰Õr“   c                 óB   — | j                   j                  j                  S rí   ©rÆ  r    r�   r  s    r(   rà  z"RealmEmbedder.get_input_embeddingsy  ó   € Ø�z‰z×$Ñ$×4Ñ4Ð4r“   c                 ó:   — || j                   j                  _        y rí   rü  râ  s     r(   rã  z"RealmEmbedder.set_input_embeddings|  ó   € Ø05ˆ�
‰
×ÑÕ-r“   úbatch_size, sequence_length©Úoutput_typerÔ  r”   rÅ   rx   ru   rÆ   r•   rÊ   r9  r:  r—   c
                 óð   — |	�|	n| j                   j                  }	| j                  |||||||||	¬«	      }
|
d   }| j                  |«      }|	s	|f|
dd z   S t	        ||
j
                  |
j                  ¬«      S )a  
        Returns:

        Example:

        ```python
        >>> from transformers import AutoTokenizer, RealmEmbedder
        >>> import torch

        >>> tokenizer = AutoTokenizer.from_pretrained("google/realm-cc-news-pretrained-embedder")
        >>> model = RealmEmbedder.from_pretrained("google/realm-cc-news-pretrained-embedder")

        >>> inputs = tokenizer("Hello, my dog is cute", return_tensors="pt")
        >>> outputs = model(**inputs)

        >>> projected_score = outputs.projected_score
        ```
        ©rÅ   rx   ru   rÆ   r•   rÊ   r9  r:  r   r1   r	   )rY  rÄ   r?  )r[   rê  rÆ  r   rX  rÄ   r?  )r’   r”   rÅ   rx   ru   rÆ   r•   rÊ   r9  r:  Úrealm_outputsré  rY  s                r(   r¡   zRealmEmbedder.forward  s�   € ðB &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàŸ
™
ØØ)Ø)Ø%ØØ'Ø/Ø!5Ø#ð #ó 

ˆð & aÑ(ˆàŸ(™( =Ó1ˆáØ#Ð%¨°a¸Ð(:Ñ:Ð:ä&Ø /Ø+×9Ñ9Ø(×3Ñ3ôð r“   )	NNNNNNNNN)rE   r¢   r£   Ú_tied_weights_keysr|   rà  rã  r   ÚREALM_INPUTS_DOCSTRINGÚformatr   rX  Ú_CONFIG_FOR_DOCr   rW   r¥   r¦   rï   r   r   r¡   r¨   r©   s   @r(   rJ   rJ   k  s-  ø„ ð
 9Ð9Ðôò5ò6ñ +Ð+A×+HÑ+HÐIfÓ+gÓhÙÐ+>È_Ô]ð 15Ø6:Ø59Ø37Ø15Ø59Ø,0Ø/3Ø&*ñ9à˜E×,Ñ,Ñ-ð9ð ! ×!2Ñ!2Ñ3ð9ð ! ×!1Ñ!1Ñ2ð	9ð
 ˜u×/Ñ/Ñ0ð9ð ˜E×-Ñ-Ñ.ð9ð   × 1Ñ 1Ñ2ð9ð $ D™>ð9ð ' t™nð9ð ˜d‘^ð9ð 
ˆuÐ)Ð)Ñ	*ò9ó ^ó iô9r“   rJ   zoThe scorer of REALM outputting relevance scores representing the score of document candidates (before softmax).c            !       óî  ‡ — e Zd ZdZdˆ fd„	Z eej                  d«      «       ee	e
¬«      	 	 	 	 	 	 	 	 	 	 	 	 	 ddeej                     deej                     deej                     deej                     d	eej                     d
eej                     deej                     deej                     deej                     deej                     dee   dee   dee   deee	f   fd„«       «       Zˆ xZS )ÚRealmScorerz­
    Args:
        query_embedder ([`RealmEmbedder`]):
            Embedder for input sequences. If not specified, it will use the same embedder as candidate sequences.
    c                 ó¢   •— t         ‰| �  |«       t        | j                  «      | _        |�|n| j                  | _        | j                  «        y rí   )r{   r|   rJ   r[   ÚembedderÚquery_embedderrÜ  )r’   r[   r  rD   s      €r(   r|   zRealmScorer.__init__È  s@   ø€ Ü‰Ñ˜Ô ä% d§k¡kÓ2ˆŒà0>Ð0J™nÐPT×P]ÑP]ˆÔà�‰Õr“   r   r  r”   rÅ   rx   ru   Úcandidate_input_idsÚcandidate_attention_maskÚcandidate_token_type_idsÚcandidate_inputs_embedsrÆ   r•   rÊ   r9  r:  r—   c                 óê  — |�|n| j                   j                  }|€|
€t        d«      ‚|€|€t        d«      ‚| j                  |||||	|
|||¬«	      }| j	                  |||«      \  }}}| j                  |||||	||||¬«	      }|d   }|d   }|j                  d| j                   j                  | j                   j                  «      }t        j                  d||«      }|s|||fS t        |||¬«      S )a÷
  
        candidate_input_ids (`torch.LongTensor` of shape `(batch_size, num_candidates, sequence_length)`):
            Indices of candidate input sequence tokens in the vocabulary.

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

            [What are input IDs?](../glossary#input-ids)
        candidate_attention_mask (`torch.FloatTensor` of shape `(batch_size, num_candidates, sequence_length)`, *optional*):
            Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:

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

            [What are attention masks?](../glossary#attention-mask)
        candidate_token_type_ids (`torch.LongTensor` of shape `(batch_size, num_candidates, sequence_length)`, *optional*):
            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,
            1]`:

            - 0 corresponds to a *sentence A* token,
            - 1 corresponds to a *sentence B* token.

            [What are token type IDs?](../glossary#token-type-ids)
        candidate_inputs_embeds (`torch.FloatTensor` of shape `(batch_size * num_candidates, sequence_length, hidden_size)`, *optional*):
            Optionally, instead of passing `candidate_input_ids` you can choose to directly pass an embedded
            representation. This is useful if you want more control over how to convert *candidate_input_ids* indices
            into associated vectors than the model's internal embedding lookup matrix.

        Returns:

        Example:

        ```python
        >>> import torch
        >>> from transformers import AutoTokenizer, RealmScorer

        >>> tokenizer = AutoTokenizer.from_pretrained("google/realm-cc-news-pretrained-scorer")
        >>> model = RealmScorer.from_pretrained("google/realm-cc-news-pretrained-scorer", num_candidates=2)

        >>> # batch_size = 2, num_candidates = 2
        >>> input_texts = ["How are you?", "What is the item in the picture?"]
        >>> candidates_texts = [["Hello world!", "Nice to meet you!"], ["A cute cat.", "An adorable dog."]]

        >>> inputs = tokenizer(input_texts, return_tensors="pt")
        >>> candidates_inputs = tokenizer.batch_encode_candidates(candidates_texts, max_length=10, return_tensors="pt")

        >>> outputs = model(
        ...     **inputs,
        ...     candidate_input_ids=candidates_inputs.input_ids,
        ...     candidate_attention_mask=candidates_inputs.attention_mask,
        ...     candidate_token_type_ids=candidates_inputs.token_type_ids,
        ... )
        >>> relevance_score = outputs.relevance_score
        ```z5You have to specify either input_ids or input_embeds.zJYou have to specify either candidate_input_ids or candidate_inputs_embeds.r  r   rv   z
bd,bnd->bn)r]  r^  r_  )r[   rê  r²   r  rÓ  r  rÀ   Únum_candidatesrŒ  rW   rÓ   r\  )r’   r”   rÅ   rx   ru   r  r  r  r  rÆ   r•   rÊ   r9  r:  Úquery_outputsÚflattened_input_idsÚflattened_attention_maskÚflattened_token_type_idsÚcandidate_outputsr^  r_  r]  s                         r(   r¡   zRealmScorer.forwardÑ  sH  € ðR &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ Ð!6ÜÐTÓUÐUàÐ&Ð+BÐ+JÜÐiÓjÐjà×+Ñ+ØØ)Ø)Ø%ØØ'Ø/Ø!5Ø#ð ,ó 

ˆð UY×ThÑThØÐ!9Ð;SóU
ÑQÐ	Ð6Ð8Pð !ŸM™MØØ3Ø3Ø%ØØ1Ø/Ø!5Ø#ð *ó 

Ðð $ AÑ&ˆà+¨AÑ.ˆà)×.Ñ.¨r°4·;±;×3MÑ3MÈtÏ{É{×OnÑOnÓoˆäŸ,™, |°[À/ÓRˆáØ" K°Ð@Ð@ä Ø+¸ÐVeô
ð 	
r“   rí   rø  )rE   r¢   r£   r¤   r|   r   r  r  r   r\  r	  r   rW   r¥   r¦   rï   r   r   r¡   r¨   r©   s   @r(   r  r  ½  s�  ø„ ñ
õñ +Ð+A×+HÑ+HÐIfÓ+gÓhÙÐ+<È?Ô[ð 15Ø6:Ø59Ø37Ø:>Ø@DØ?CØ?CØ15Ø59Ø,0Ø/3Ø&*ñz
à˜E×,Ñ,Ñ-ðz
ð ! ×!2Ñ!2Ñ3ðz
ð ! ×!1Ñ!1Ñ2ð	z
ð
 ˜u×/Ñ/Ñ0ðz
ð & e×&6Ñ&6Ñ7ðz
ð #+¨5×+<Ñ+<Ñ"=ðz
ð #+¨5×+;Ñ+;Ñ"<ðz
ð "*¨%×*;Ñ*;Ñ!<ðz
ð ˜E×-Ñ-Ñ.ðz
ð   × 1Ñ 1Ñ2ðz
ð $ D™>ðz
ð ' t™nðz
ð ˜d‘^ðz
ð 
ˆuÐ'Ð'Ñ	(òz
ó \ó iôz
r“   r  zrThe knowledge-augmented encoder of REALM outputting masked language model logits and marginal log-likelihood loss.c                   óæ  ‡ — e Zd ZdgZˆ fd„Zd„ Zd„ Zd„ Zd„ Z e	e
j                  d«      «       eee¬«      	 	 	 	 	 	 	 	 	 	 	 	 dd	eej"                     d
eej$                     deej"                     deej"                     deej$                     deej$                     deej$                     deej"                     deej"                     dee   dee   dee   deeef   fd„«       «       Zˆ xZS )rI   zcls.predictions.decoderc                 ó¬   •— t         ‰| �  |«       t        | j                  «      | _        t        | j                  «      | _        | j                  «        y rí   )r{   r|   rØ  r[   rÆ  rƒ  r   rÜ  r‘   s     €r(   r|   z!RealmKnowledgeAugEncoder.__init__X  s:   ø€ Ü‰Ñ˜Ô Ü# D§K¡KÓ0ˆŒ
Ü# D§K¡KÓ0ˆŒØ�‰Õr“   c                 óB   — | j                   j                  j                  S rí   rü  r  s    r(   rà  z-RealmKnowledgeAugEncoder.get_input_embeddings^  rý  r“   c                 ó:   — || j                   j                  _        y rí   rü  râ  s     r(   rã  z-RealmKnowledgeAugEncoder.set_input_embeddingsa  rÿ  r“   c                 óB   — | j                   j                  j                  S rí   )r   r…  r|  r  s    r(   Úget_output_embeddingsz.RealmKnowledgeAugEncoder.get_output_embeddingsd  s   € Ø�x‰x×#Ñ#×+Ñ+Ð+r“   c                 ó„   — || j                   j                  _        |j                  | j                   j                  _        y rí   )r   r…  r|  r0   )r’   Únew_embeddingss     r(   Úset_output_embeddingsz.RealmKnowledgeAugEncoder.set_output_embeddingsg  s,   € Ø'5ˆ�‰×ÑÔ$Ø$2×$7Ñ$7ˆ�‰×ÑÕ!r“   z+batch_size, num_candidates, sequence_lengthr  r”   rÅ   rx   ru   rÆ   r•   r]  ÚlabelsÚmlm_maskrÊ   r9  r:  r—   c                 ó0  — |�|n| j                   j                  }|�|€t        d«      ‚| j                  |||«      \  }}}| j	                  |||||||
||¬«	      }|d   }| j                  |«      }|}d}|��h|j                  «       \  }}|	€&t        j                  |t        j                  ¬«      }	n|	j                  t        j                  «      }	t        d¬«      }|j                  d| j                   j                  «      }|j                  d	| j                   j                  «      j                  d«      } |||«      j                  || j                   j                  |«       }|j!                  d«      j#                  d«      }||z   }|j%                  d	«      }t        j&                  t        j(                  ||	z  «      t        j(                  |	«      z  «       }|s|f|d
d z   }|�|f|z   S |S t+        |||j,                  |j.                  ¬«      S )aÕ  
        relevance_score (`torch.FloatTensor` of shape `(batch_size, num_candidates)`, *optional*):
            Relevance score derived from RealmScorer, must be specified if you want to compute the masked language
            modeling loss.

        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should be in `[-100, 0, ...,
            config.vocab_size]` (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]`

        mlm_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Mask to avoid calculating joint loss on certain positions. If not specified, the loss will not be masked.
            Mask values selected in `[0, 1]`:

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

        Returns:

        Example:

        ```python
        >>> import torch
        >>> from transformers import AutoTokenizer, RealmKnowledgeAugEncoder

        >>> tokenizer = AutoTokenizer.from_pretrained("google/realm-cc-news-pretrained-encoder")
        >>> model = RealmKnowledgeAugEncoder.from_pretrained(
        ...     "google/realm-cc-news-pretrained-encoder", num_candidates=2
        ... )

        >>> # batch_size = 2, num_candidates = 2
        >>> text = [["Hello world!", "Nice to meet you!"], ["The cute cat.", "The adorable dog."]]

        >>> inputs = tokenizer.batch_encode_candidates(text, max_length=10, return_tensors="pt")
        >>> outputs = model(**inputs)
        >>> logits = outputs.logits
        ```NzZYou have to specify `relevance_score` when `labels` is specified in order to compute loss.r  r   ry   Únone)Ú	reductionrv   r   r1   r	   )rb  ÚlogitsrÄ   r?  )r[   rê  r²   rÓ  rÆ  r   r�   rW   Ú	ones_liker·  r±  r   rÀ   r~   Útiler  Úlog_softmaxÚ	unsqueezeÚ	logsumexpÚnansumÚsumr   rÄ   r?  )r’   r”   rÅ   rx   ru   rÆ   r•   r]  r#  r$  rÊ   r9  r:  r  r  r  Újoint_outputsÚjoint_outputrˆ  r_  Úmasked_lm_lossrð  r�   Úloss_fctÚ
mlm_logitsÚmlm_targetsÚmasked_lm_log_probÚcandidate_log_probÚjoint_gold_log_probÚmarginal_gold_log_probsr  s                                  r(   r¡   z RealmKnowledgeAugEncoder.forwardk  s2  € ðp &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ /Ð"9ÜØlóð ð UY×ThÑThØ�~ ~óU
ÑQÐ	Ð6Ð8Pð Ÿ
™
ØØ3Ø3Ø%ØØ'Ø/Ø!5Ø#ð #ó 

ˆð % QÑ'ˆà ŸH™H \Ó2Ðà)ˆàˆØÑØ%+§[¡[£]Ñ"ˆJ˜
àÐÜ Ÿ?™?¨6¼¿¹ÔG‘à#Ÿ=™=¬¯©Ó7�ô (°&Ô9ˆHð +×/Ñ/°°D·K±K×4JÑ4JÓKˆJà Ÿ+™+ a¨¯©×)CÑ)CÓD×IÑIÈ"ÓMˆKá"*¨:°{Ó"C×"HÑ"HØ˜DŸK™K×6Ñ6¸
ó#ð "Ðð "1×!<Ñ!<¸RÓ!@×!JÑ!JÈ2Ó!NÐà"4Ð7IÑ"IÐà&9×&CÑ&CÀAÓ&FÐ#ä#Ÿl™l¬5¯9©9Ð5LÈxÑ5WÓ+XÔ[`×[dÑ[dÐemÓ[nÑ+nÓoÐoˆNáØ'Ð)¨M¸!¸AÐ,>Ñ>ˆFØ3AÐ3M�^Ð%¨Ñ.ÐYÐSYÐYäØØ$Ø'×5Ñ5Ø$×/Ñ/ô	
ð 	
r“   )NNNNNNNNNNNN)rE   r¢   r£   r  r|   rà  rã  r  r"  r   r  r  r   r   r	  r   rW   r¥   r¦   rï   r   r   r¡   r¨   r©   s   @r(   rI   rI   P  s�  ø„ ð 4Ð4Ðôò5ò6ò,ò8ñ +Ø×%Ñ%Ð&SÓTóñ ¨>ÈÔXð 15Ø6:Ø59Ø37Ø15Ø59Ø7;Ø-1Ø/3Ø,0Ø/3Ø&*ñx
à˜E×,Ñ,Ñ-ðx
ð ! ×!2Ñ!2Ñ3ðx
ð ! ×!1Ñ!1Ñ2ð	x
ð
 ˜u×/Ñ/Ñ0ðx
ð ˜E×-Ñ-Ñ.ðx
ð   × 1Ñ 1Ñ2ðx
ð " %×"3Ñ"3Ñ4ðx
ð ˜×)Ñ)Ñ*ðx
ð ˜5×+Ñ+Ñ,ðx
ð $ D™>ðx
ð ' t™nðx
ð ˜d‘^ðx
ð 
ˆu�nÐ$Ñ	%òx
ó Yóôx
r“   rI   zThe reader of REALM.c            #       ó  ‡ — e Zd Zˆ fd„Z eej                  d«      «       eee	¬«      	 	 	 	 	 	 	 	 	 	 	 	 	 	 dde
ej                     de
ej                     de
ej                     de
ej                     de
ej                     d	e
ej                     d
e
ej                     de
ej                     de
ej                     de
ej                     de
ej                     de
e   de
e   de
e   deeef   fd„«       «       Zˆ xZS )rC   c                 óÆ   •— t         ‰| �  |«       |j                  | _        t        |«      | _        t        |«      | _        t        |«      | _        | j                  «        y rí   )
r{   r|   Ú
num_labelsrØ  rÆ  rƒ  r   r�  Ú
qa_outputsrÜ  r‘   s     €r(   r|   zRealmReader.__init__ì  sK   ø€ Ü‰Ñ˜Ô Ø ×+Ñ+ˆŒä# FÓ+ˆŒ
Ü# FÓ+ˆŒÜ/°Ó7ˆŒà�‰Õr“   z!reader_beam_size, sequence_lengthr  r”   rÅ   rx   ru   rÆ   r•   r]  rº  Ústart_positionsÚend_positionsÚhas_answersrÊ   r9  r:  r—   c                 óþ  — |�|n| j                   j                  }|€t        d«      ‚|€t        d«      ‚|j                  d«      | j                   j                  k  rt        d«      ‚| j                  |||||||||¬«	      }|d   }| j                  ||d| j                   j                   «      \  }}}t        j                  |d| j                   j                   d«      }||z  }t        j                  t        j                  |d¬	«      j                  «      }t        j                  t        j                  |d¬	«      j                  «      }t        j                  |d|¬
«      }t        j                  |d|¬
«      }d}d}d}d}d}|	��.|
��+|��(d„ }d„ }|j                  d«      } |	j                  d| «      }	|
j                  d| «      }
|}t        j                  |«      }! ||||	d| j                   j                   |
d| j                   j                   ¬«      }t        j                  |«      }" |||«      } ||j!                  d«      |j!                  d«      «      }||!j#                  t        j$                  «      z  }||"j#                  t        j$                  «      z  }||z   j'                  «       }|s||||f|dd z   }#|�
|||||f|#z   S |#S t)        ||||||||||j*                  |j,                  ¬«      S )ar  
        relevance_score (`torch.FloatTensor` of shape `(searcher_beam_size,)`, *optional*):
            Relevance score, which must be specified if you want to compute the logits and marginal log loss.
        block_mask (`torch.BoolTensor` of shape `(searcher_beam_size, sequence_length)`, *optional*):
            The mask of the evidence block, which must be specified if you want to compute the logits and marginal log
            loss.
        start_positions (`torch.LongTensor` of shape `(searcher_beam_size,)`, *optional*):
            Labels for position (index) of the start of the labelled span for computing the token classification loss.
            Positions are clamped to the length of the sequence (`sequence_length`). Position outside of the sequence
            are not taken into account for computing the loss.
        end_positions (`torch.LongTensor` of shape `(searcher_beam_size,)`, *optional*):
            Labels for position (index) of the end of the labelled span for computing the token classification loss.
            Positions are clamped to the length of the sequence (`sequence_length`). Position outside of the sequence
            are not taken into account for computing the loss.
        has_answers (`torch.BoolTensor` of shape `(searcher_beam_size,)`, *optional*):
            Whether or not the evidence block has answer(s).

        Returns:
        NzCYou have to specify `relevance_score` to calculate logits and loss.zOYou have to specify `block_mask` to separate question block and evidence block.r   zQThe input sequence length must be greater than or equal to config.max_span_width.r  r   rv   rÌ   r¤  c                 óž  — t        j                  t        j                  t        j                  | d«      d«      t        j                  |d«      «      }t        j                  t        j                  t        j                  |d«      d«      t        j                  |d«      «      }t        j                  t        j                  ||«      d«      S )zCompute correct span.r   rv   r   )rW   Úeqr,  rL   Úlogical_and)r½  r¾  Úgold_startsÚ	gold_endsÚis_gold_startÚis_gold_ends         r(   Úcompute_correct_candidatesz7RealmReader.forward.<locals>.compute_correct_candidatesK  s”   € ô !&§¡Ü—O‘O¤E§O¡OÐ4DÀaÓ$HÈ!ÓLÌeÏoÉoÐ^iÐkmÓNnó!�ô $Ÿh™hÜ—O‘O¤E§O¡O°NÀAÓ$FÈÓJÌEÏOÉOÐ\eÐgiÓLjó�ô
 —y‘y¤×!2Ñ!2°=À+Ó!NÐPQÓRÐRr“   c                 ó¸   — t         j                  fd„}t        j                  |  ||| j                  ¬«      z   d¬«      }t        j                  | d¬«      }||z
  S )z3Loss based on the negative marginal log-likelihood.c                 ój   — d| j                  |«      z
  t        j                  |«      j                  z  S r®  r°  r´  s     r(   r¶  zERealmReader.forward.<locals>.marginal_log_loss.<locals>.mask_to_score[  s*   € Ø $§)¡)¨EÓ"2Ñ2´e·k±kÀ%Ó6H×6LÑ6LÑLÐLr“   ry   rv   rÌ   )rW   r·  r-  rz   )r(  Ú
is_correctr¶  Úlog_numeratorÚlog_denominators        r(   Úmarginal_log_lossz.RealmReader.forward.<locals>.marginal_log_lossX  sR   € ô /4¯m©mó Mô !&§¡°¹ÀzÐY_×YeÑYeÔ9fÑ0fÐlnÔ o�Ü"'§/¡/°&¸bÔ"A�Ø&¨Ñ6Ð6r“   )r½  r¾  rE  rF  r1   )rb  rc  rd  re  rf  rg  rh  ri  rj  rÄ   r?  )r[   rê  r²   r�   r¥  rÆ  r=  Úreader_beam_sizerW   r,  ÚargmaxÚmaxÚvaluesr¦  ÚclamprL   rÀ   r±  r·  rÈ  ra  rÄ   r?  )$r’   r”   rÅ   rx   ru   rÆ   r•   r]  rº  r>  r?  r@  rÊ   r9  r:  rì   r‡  rÃ  r½  r¾  Úretriever_logitsÚpredicted_block_indexÚpredicted_candidateÚpredicted_startÚpredicted_endÚ
total_lossrc  rd  re  rf  rI  rO  Úignored_indexÚany_retriever_correctÚany_reader_correctr  s$                                       r(   r¡   zRealmReader.forwardö  s7  € ðL &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ"ÜÐbÓcÐcØÐÜÐnÓoÐoØ×Ñ˜qÓ! D§K¡K×$>Ñ$>Ò>ÜÐpÓqÐqØ—*‘*ØØ)Ø)Ø%ØØ'Ø/Ø!5Ø#ð ó 

ˆð " !™*ˆð ;?¿/¹/Ø˜Z¨¨D¯K©K×,HÑ,HÐIó;
Ñ7ˆÐ'¨ô !Ÿ?™?¨?¸1¸t¿{¹{×?[Ñ?[Ð+\Ð^`ÓaÐàÐ)Ñ)ˆä %§¡¬U¯Y©Y°}È!Ô-L×-SÑ-SÓ TÐä#Ÿl™l¬5¯9©9°]ÈÔ+J×+QÑ+QÓRÐä×,Ñ,Ð-=À1ÐL_Ô`ˆä×*Ñ*¨>¸qÐH[Ô\ˆàˆ
ØˆØˆØ ÐØˆØÑ&¨=Ñ+DÈÑI`òSò	7ð ,×0Ñ0°Ó3ˆMØ-×3Ñ3°B¸ÓFˆOØ)×/Ñ/°°MÓBˆMà +ÐÜ$)§I¡IÐ.?Ó$@Ð!á7Ø!1Ø-Ø+¨A°·±×0LÑ0LÐMØ'¨¨D¯K©K×,HÑ,HÐIô	ˆNô "'§¡¨>Ó!:Ðá.¨Ð@QÓRˆNÙ+¨M×,>Ñ,>¸rÓ,BÀN×DWÑDWÐXZÓD[Ó\ˆKØÐ3×8Ñ8¼¿¹ÓGÑGˆNØÐ-×2Ñ2´5·=±=ÓAÑAˆKà(¨;Ñ6×<Ñ<Ó>ˆJáØ+Ð-@À/ÐS`ÐaÐdkÐlmÐlnÐdoÑoˆFð Ð)ð ˜n¨kÐ;LÈnÐ]Ð`fÑfðð ðô !ØØ)Ø#Ø/Ø)Ø+Ø)Ø%Ø!Ø!×/Ñ/Ø×)Ñ)ô
ð 	
r“   )NNNNNNNNNNNNNN)rE   r¢   r£   r|   r   r  r  r   ra  r	  r   rW   r¥   r¦   rk  rï   r   r   r¡   r¨   r©   s   @r(   rC   rC   ê  s¡  ø„ ôñ +Ð+A×+HÑ+HÐIlÓ+mÓnÙÐ+<È?Ô[ð 15Ø6:Ø59Ø37Ø15Ø59Ø7;Ø15Ø6:Ø48Ø26Ø,0Ø/3Ø&*ñW
à˜E×,Ñ,Ñ-ðW
ð ! ×!2Ñ!2Ñ3ðW
ð ! ×!1Ñ!1Ñ2ð	W
ð
 ˜u×/Ñ/Ñ0ðW
ð ˜E×-Ñ-Ñ.ðW
ð   × 1Ñ 1Ñ2ðW
ð " %×"3Ñ"3Ñ4ðW
ð ˜U×-Ñ-Ñ.ðW
ð " %×"2Ñ"2Ñ3ðW
ð   × 0Ñ 0Ñ1ðW
ð ˜e×.Ñ.Ñ/ðW
ð $ D™>ðW
ð ' t™nðW
ð ˜d‘^ðW
ð  
ˆuÐ'Ð'Ñ	(ò!W
ó \ó oôW
r“   rC   ay  
    Args:
        input_ids (`torch.LongTensor` of shape `({0})`):
            Indices of input sequence tokens in the vocabulary.

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

            [What are input IDs?](../glossary#input-ids)
        attention_mask (`torch.FloatTensor` of shape `({0})`, *optional*):
            Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:

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

            [What are attention masks?](../glossary#attention-mask)
        token_type_ids (`torch.LongTensor` of shape `({0})`, *optional*):
            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,
            1]`:

            - 0 corresponds to a *sentence A* token,
            - 1 corresponds to a *sentence B* token (should not be used in this model by design).

            [What are token type IDs?](../glossary#token-type-ids)
        answer_ids (`list` of shape `(num_answers, answer_length)`, *optional*):
            Answer ids for computing the marginal log-likelihood loss. Indices should be in `[-1, 0, ...,
            config.vocab_size]` (see `input_ids` docstring) Tokens with indices set to `-1` are ignored (masked), the
            loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
z?`RealmForOpenQA` for end-to-end open domain question answering.c                   ó&  ‡ — e Zd Zdˆ fd„	Zed„ «       Zd„ Z eej                  d«      «       e
ee¬«      	 	 	 	 ddeej                     deej                      deej                     d	eej                     d
ee   deeef   fd„«       «       Zˆ xZS )rG   c           
      ón  •— t         ‰| �  |«       t        |«      | _        t	        |«      | _        | j                  dt        j                  d«      j                  |j                  |j                  ft        j                  t        j                  d«      ¬«      «       || _        | j                  «        y )NÚ	block_embr%   Úcpu)r�   rz   rš   )r{   r|   rJ   r  rC   r   r‹   rW   rŽ   Ú	new_emptyÚnum_block_recordsrŒ  r·  rš   Ú	retrieverrÜ  )r’   r[   rd  rD   s      €r(   r|   zRealmForOpenQA.__init__¸  sŽ   ø€ Ü‰Ñ˜Ô Ü% fÓ-ˆŒÜ! &Ó)ˆŒØ×ÑØÜ�K‰K˜‹O×%Ñ%Ø×.Ñ.°×0JÑ0JÐKÜ—m‘mÜ—|‘| EÓ*ð &ó ô	
ð #ˆŒà�‰Õr“   c                 ór   — | j                   r| j                  j                  S | j                  j                  S rí   )rA  r[   Úsearcher_beam_sizerP  r  s    r(   rf  z!RealmForOpenQA.searcher_beam_sizeÈ  s)   € à�=Š=Ø—;‘;×1Ñ1Ð1Ø�{‰{×+Ñ+Ð+r“   c                 óD   — | j                   j                  |«      | _         y)z´Send `self.block_emb` to a specific device.

        Args:
            device (`str` or `torch.device`):
                The device to which `self.block_emb` will be sent.
        N)r`  rÒ   )r’   rš   s     r(   Úblock_embedding_toz!RealmForOpenQA.block_embedding_toÎ  s   € ð Ÿ™×*Ñ*¨6Ó2ˆ�r“   z1, sequence_lengthr  r”   rÅ   rx   Ú
answer_idsr:  r—   c                 ó@  — |�|n| j                   j                  }|�|j                  d   dk7  rt        d«      ‚| j	                  |||d¬«      }|d   }t        j                  d| j                  |j                  | j                  j                  «      «      }t        j                  || j                  d¬«      \  }	}
|
j                  «       }
t        j                  | j                  d|
¬	«      }| j                  |
j                  «       ||| j                   j                   ¬
«      \  }}}}|j                  | j"                  j                  «      }|j$                  j'                  t
        j(                  «      j                  | j"                  j                  ¬«      }|j+                  «       j-                  |j.                  j'                  t
        j(                  «      «       |�®t        j0                  |t
        j(                  | j"                  j                  ¬«      }t        j0                  |t
        j2                  | j"                  j                  ¬«      }t        j0                  |t
        j2                  | j"                  j                  ¬«      }t        j                  d|j                  «       |j                  | j"                  j                  «      «      }| j#                  |j4                  d| j                   j6                   |j8                  d| j                   j6                   |j.                  d| j                   j6                   |||||d¬«	      }|j4                  |j:                     }||j<                  |j>                  dz    }|s||fS tA        ||¬«      S )a  
        Returns:

        Example:

        ```python
        >>> import torch
        >>> from transformers import RealmForOpenQA, RealmRetriever, AutoTokenizer

        >>> retriever = RealmRetriever.from_pretrained("google/realm-orqa-nq-openqa")
        >>> tokenizer = AutoTokenizer.from_pretrained("google/realm-orqa-nq-openqa")
        >>> model = RealmForOpenQA.from_pretrained("google/realm-orqa-nq-openqa", retriever=retriever)

        >>> question = "Who is the pioneer in modern computer science?"
        >>> question_ids = tokenizer([question], return_tensors="pt")
        >>> answer_ids = tokenizer(
        ...     ["alan mathison turing"],
        ...     add_special_tokens=False,
        ...     return_token_type_ids=False,
        ...     return_attention_mask=False,
        ... ).input_ids

        >>> reader_output, predicted_answer_ids = model(**question_ids, answer_ids=answer_ids, return_dict=False)
        >>> predicted_answer = tokenizer.decode(predicted_answer_ids)
        >>> loss = reader_output.loss
        ```r   r   z'The batch_size of the inputs must be 1.T)r”   rx   rÅ   r:  z	BD,QD->QBrv   )ÚkrÍ   r¤  )Ú
max_lengthr›  r™   zD,BD->B)	r”   rÅ   rx   r]  rº  r@  r>  r?  r:  )ro  rp  )!r[   rê  rT   r²   r  rW   rÓ   r`  rÒ   rš   Útopkrf  r¹  r¦  rd  ra  Úreader_seq_lenr   Úspecial_tokens_maskr±  rï   Úlogical_not_Úlogical_and_rx   rÑ   r�   r”   rP  rÅ   rg  ri  rj  rn  )r’   r”   rÅ   rx   ri  r:  Úquestion_outputsÚquestion_projectionÚbatch_scoresr7  Úretrieved_block_idsÚretrieved_block_embr@  ri  rj  Úconcat_inputsrº  Úretrieved_logitsro  Úpredicted_blockrp  s                        r(   r¡   zRealmForOpenQA.forwardØ  s  € ðJ &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ  Y§_¡_°QÑ%7¸1Ò%<ÜÐFÓGÐGàŸ=™=Ø°È~Ðkoð )ó 
Ðð /¨qÑ1Ðô —|‘| K°·±ÐAT×AWÑAWÐX\×XfÑXf×XmÑXmÓAnÓoˆä!&§¡¨L¸D×<SÑ<SÐY[Ô!\ÑˆÐà1×9Ñ9Ó;Ðä#×0Ñ0°·±ÀQÐNaÔbÐð :>¿¹Ø×#Ñ#Ó% y°*ÈÏÉ×IcÑIcð :Hó :
Ñ6ˆ�Y ¨ð &×(Ñ(¨¯©×);Ñ);Ó<ˆØ"×6Ñ6×;Ñ;¼E¿J¹JÓG×JÑJÐRV×R]ÑR]×RdÑRdÐJÓeˆ
Ø×ÑÓ!×.Ñ.¨}×/KÑ/K×/PÑ/PÔQV×Q[ÑQ[Ó/\Ô]àÐ"ÜŸ,™, {¼%¿*¹*ÈTÏ[É[×M_ÑM_Ô`ˆKÜŸ™ Y´e·j±jÈÏÉ×I[ÑI[Ô\ˆIÜ—l‘l 7´%·*±*ÀTÇ[Á[×EWÑEWÔXˆGô !Ÿ<™<ØÐ*×2Ñ2Ó4Ð6I×6LÑ6LÈTÏ[É[×M_ÑM_Ó6`ó
Ðð Ÿ™Ø#×-Ñ-¨a°$·+±+×2NÑ2NÐOØ(×7Ñ7¸¸D¿K¹K×<XÑ<XÐYØ(×7Ñ7¸¸D¿K¹K×<XÑ<XÐYØ,Ø!Ø#Ø%Ø!Øð $ó 

ˆð (×1Ñ1°-×2IÑ2IÑJˆØ.¨}×/FÑ/FÈ×I^ÑI^ÐabÑIbÐcÐáØ Ð"6Ð6Ð6ä#Ø'Ø!5ô
ð 	
r“   rí   )NNNN)rE   r¢   r£   r|   Úpropertyrf  rh  r   ÚREALM_FOR_OPEN_QA_DOCSTRINGr  r   rn  r	  r   rW   r¥   r¦   rï   r   r   r¡   r¨   r©   s   @r(   rG   rG   ³  sä   ø„ õ
ð  ñ,ó ð,ò
3ñ +Ð+F×+MÑ+MÐNbÓ+cÓdÙÐ+?ÈoÔ^ð 7;Ø59Ø15Ø&*ña
à˜E×,Ñ,Ñ-ða
ð ! ×!2Ñ!2Ñ3ða
ð ! ×!1Ñ!1Ñ2ð	a
ð
 ˜U×-Ñ-Ñ.ða
ð ˜d‘^ða
ð 
ˆuÐ*Ð*Ñ	+òa
ó _ó eôa
r“   rG   )Gr¤   rÔ   r9   Údataclassesr   Útypingr   r   r   rW   r   Útorch.nnr   Úactivationsr
   Úmodeling_outputsr   r   r   r   Úmodeling_utilsr   Úpytorch_utilsr   r   r   Úutilsr   r   r   r   Úconfiguration_realmr   Ú
get_loggerrE   r7   Ú_EMBEDDER_CHECKPOINT_FOR_DOCÚ_ENCODER_CHECKPOINT_FOR_DOCÚ_SCORER_CHECKPOINT_FOR_DOCr	  rl   ÚModulern   r«   rñ   r  rþ   r  r  r  r0  rP  rX  r\  ra  rn  rs  ry  rƒ  rŠ  r�  ÚREALM_START_DOCSTRINGr  rÅ  rØ  rJ   r  rI   rC   r{  rG   r%   r“   r(   ú<module>r‹     sî  ðñ ã Û 	Ý !ß )Ñ )ã Ý Ý %å "÷ó õ /ß mÑ mß uÓ uÝ ,ð 
ˆ×	Ñ	˜HÓ	%€ØIÐ ØGÐ ØEÐ Ø€òhôV=�b—i‘iô =ô@C˜Ÿ™ô CôL�b—i‘iô ð Ðð Ð ô
0�R—Y‘Yô 0ôf˜Ÿ	™	ô ô�"—)‘)ô ôS�—‘ô SôlZ
�2—9‘9ô Z
ôz�"—)‘)ô ð ô:˜+ó :ó ð:ð4 ô8˜ó 8ó ð8ð$ ô.:˜ó .:ó ð.:ðb ô<˜;ó <ó ð<ô  2§9¡9ô ô"˜BŸI™Iô ô.!�r—y‘yô !ô
˜BŸI™Iô 
ô@?˜BŸI™Iô @?ðF	Ð ð/Ð ôd% ˜?ô % ôPL
Ð)ô L
ñ^ ØfØóôKÐ(ó Kó	ðKñ\ ØuØóôL
Ð&ó L
ó	ðL
ñ^ ðàóô
R
Ð3ó R
óð
R
ñj Ð,Ð.CÓDôd
Ð&ó d
ó Eðd
ðNÐ ñB ØEØóôD
Ð)ó D
ó	ñD
r“   