a
    þdsJ ã                   @   sä  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 !e"¡Z#dZ$dZ%dZ&dZ'g d¢Z(dd„ Z)G dd„ de
j*ƒZ+G dd„ de
j*ƒZ,G dd„ de
j*ƒZ-G dd„ de
j*ƒZ.G dd„ de
j*ƒZ/G d d!„ d!e
j*ƒZ0G d"d#„ d#e
j*ƒZ1G d$d%„ d%e
j*ƒZ2G d&d'„ d'e
j*ƒZ3eG d(d)„ d)eƒƒZ4eG d*d+„ d+eƒƒZ5eG d,d-„ d-eƒƒZ6eG d.d/„ d/eƒƒZ7G d0d1„ d1e
j*ƒZ8G d2d3„ d3e
j*ƒZ9G d4d5„ d5e
j*ƒZ:G d6d7„ d7e
j*ƒZ;G d8d9„ d9e
j*ƒZ<d:Z=d;Z>G d<d=„ d=eƒZ?G d>d?„ d?e?ƒZ@ed@e=ƒG dAdB„ dBe?ƒƒZAedCe=ƒG dDdE„ dEe?ƒƒZBedFe=ƒG dGdH„ dHe?ƒƒZCedIe=ƒG dJdK„ dKe?ƒƒZDdLZEedMe=ƒG dNdO„ dOe?ƒƒZFdS )Pz 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é   )ÚRealmConfigú(google/realm-cc-news-pretrained-embedderú'google/realm-cc-news-pretrained-encoderú&google/realm-cc-news-pretrained-scorerr   )r   r   r   z&google/realm-cc-news-pretrained-openqazgoogle/realm-orqa-nq-openqazgoogle/realm-orqa-nq-readerzgoogle/realm-orqa-wq-openqazgoogle/realm-orqa-wq-readerc                 C   sŒ  zddl }ddl}ddl}W n ty:   t d¡ ‚ Y n0 tj |¡}t 	d|› ¡ |j
 |¡}g }g }	|D ]@\}
}t 	d|
› d|› ¡ |j
 ||
¡}| |
¡ |	 |¡ qpt||	ƒD ]È\}
}t| tƒröd|
vröt 	d|
› d	| jj› d
¡ q¼|
 d¡s|
 d¡r2t| tƒr2|
 dd¡}
|
 dd¡}
|
 d¡sJ|
 d¡rbt| tƒrb|
 dd¡}
|
 d¡rÜt| tƒr~dnd}|
 d|› d¡}
|
 d|› d¡}
|
 d|› d¡}
|
 d|› d¡}
|
 d|› d¡}
|
 d¡rjt| tƒrødnd}|
 d|› d¡}
|
 d|› d ¡}
|
 d!|› d"¡}
|
 d#|› d$¡}
|
 d%|› d¡}
|
 d&|› d$¡}
nD|
 d'¡r®t| tƒr†dnd}|
 d(|› d ¡}
|
 d)|› d"¡}
|
 d*¡}
td+d,„ |
D ƒƒrät 	dd* |
¡› ¡ q¼| }|
D ]Þ}| d-|¡r| d.|¡}n|g}|d d/ks.|d d0kr:t|d1ƒ}nl|d d2ksV|d d3krbt|d4ƒ}nDzt||d ƒ}W n0 ty¤   t 	dd* |
¡› ¡ Y qìY n0 t|ƒd5krìt|d6 ƒ}|| }qì|d7d… d8krêt|d1ƒ}n|d/krþ|  |¡}z,|j!|j!ks(J d9|j!› d:|j!› d;ƒ‚W n> t"yh } z$| j#|j!|j!f7  _#‚ W Y d}~n
d}~0 0 t 	d<|
› ¡ t$ %|¡|_&q¼| S )=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 parameterZ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                 s   s   | ]}|d v V  qdS ))Zadam_vZadam_mZAdamWeightDecayOptimizerZAdamWeightDecayOptimizer_1Zglobal_stepN© )Ú.0Únr   r   úq/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/transformers/models/realm/modeling_realm.pyÚ	<genexpr>|   s   ÿz+load_tf_weights_in_realm.<locals>.<genexpr>z[A-Za-z]+_\d+z_(\d+)ÚkernelÚgammaÚweightZoutput_biasÚbetaÚbiasé   r   iõÿÿÿZ_embeddingszPointer shape z and array shape z mismatchedzInitialize PyTorch weight )'ÚreÚnumpyZ
tensorflowÚImportErrorÚloggerÚerrorÚosÚpathÚabspathÚinfoZtrainZlist_variablesZ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ÚtorchZ
from_numpyÚdata)ÚmodelÚconfigZtf_checkpoint_pathr*   ÚnpÚtfZtf_pathZ	init_varsÚnamesZarraysÚnamerG   ÚarrayZreader_prefixZembedder_prefixÚpointerZm_nameZscope_namesÚnumÚer   r   r"   Úload_tf_weights_in_realm:   sª    ÿ
$$
þ



ÿþrV   c                       sT   e Zd ZdZ‡ fdd„Zd	eej eej eej eej e	ej
dœdd„Z‡  ZS )
ÚRealmEmbeddingszGConstruct the embeddings from word, position and token_type embeddings.c                    s¶   t ƒ  ¡  tj|j|j|jd| _t |j|j¡| _	t |j
|j¡| _tj|j|jd| _t |j¡| _t|ddƒ| _|  dt |j¡ d¡¡ | jdtj| j ¡ tjdd	d
 d S )N)Úpadding_idx©ZepsÚposition_embedding_typeÚabsoluteÚposition_ids)r   éÿÿÿÿÚtoken_type_ids©ÚdtypeF)Ú
persistent)ÚsuperÚ__init__r   Ú	EmbeddingÚ
vocab_sizeÚhidden_sizeZpad_token_idÚword_embeddingsÚmax_position_embeddingsÚposition_embeddingsZtype_vocab_sizeÚtoken_type_embeddingsÚ	LayerNormÚlayer_norm_epsÚDropoutÚhidden_dropout_probÚdropoutrB   rZ   Úregister_bufferrJ   ÚarangeÚexpandÚzerosr\   ÚsizeÚlong©ÚselfrM   ©r7   r   r"   rc   ©   s    
ÿzRealmEmbeddings.__init__Nr   )Ú	input_idsr^   r\   Úinputs_embedsÚpast_key_values_lengthÚreturnc                 C   sø   |d ur|  ¡ }n|  ¡ d d… }|d }|d u rL| jd d …||| …f }|d u r t| dƒrŠ| jd d …d |…f }| |d |¡}	|	}ntj|tj| jjd}|d u r²|  	|¡}|  
|¡}
||
 }| jdkrà|  |¡}||7 }|  |¡}|  |¡}|S )Nr]   r   r^   r   ©r`   Údevicer[   )rt   r\   Úhasattrr^   rr   rJ   rs   ru   r~   rg   rj   rZ   ri   rk   ro   )rw   ry   r^   r\   rz   r{   Úinput_shapeÚ
seq_lengthÚbuffered_token_type_idsÚ buffered_token_type_ids_expandedrj   Ú
embeddingsri   r   r   r"   Úforwardº   s,    







zRealmEmbeddings.forward)NNNNr   )r8   Ú
__module__Ú__qualname__Ú__doc__rc   r   rJ   Ú
LongTensorÚFloatTensorrE   ÚTensorr…   Ú__classcell__r   r   rx   r"   rW   ¦   s        úùrW   c                
       s‚   e Zd Zd‡ fdd„	Zejejdœdd„Zdejeej eej eej eej ee	e	ej   ee
 e	ej dœd	d
„Z‡  ZS )ÚRealmSelfAttentionNc                    sþ   t ƒ  ¡  |j|j dkr>t|dƒs>td|j› d|j› dƒ‚|j| _t|j|j ƒ| _| j| j | _t	 
|j| j¡| _t	 
|j| j¡| _t	 
|j| j¡| _t	 |j¡| _|p¸t|ddƒ| _| jdksÐ| jd	krò|j| _t	 d
|j d | j¡| _|j| _d S )Nr   Zembedding_sizezThe hidden size (z6) is not a multiple of the number of attention heads (ú)rZ   r[   Úrelative_keyÚrelative_key_queryr)   r   )rb   rc   rf   Únum_attention_headsr   Ú
ValueErrorrE   Úattention_head_sizeÚall_head_sizer   ÚLinearÚqueryÚkeyÚvaluerm   Zattention_probs_dropout_probro   rB   rZ   rh   rd   Údistance_embeddingÚ
is_decoder©rw   rM   rZ   rx   r   r"   rc   æ   s*    

ÿÿÿzRealmSelfAttention.__init__)Úxr|   c                 C   s6   |  ¡ d d… | j| jf }| |¡}| dddd¡S )Nr]   r   r)   r   r   )rt   r‘   r“   ÚviewÚpermute)rw   rœ   Znew_x_shaper   r   r"   Útranspose_for_scores   s    
z'RealmSelfAttention.transpose_for_scoresF©Úhidden_statesÚattention_maskÚ	head_maskÚencoder_hidden_statesÚencoder_attention_maskÚpast_key_valueÚoutput_attentionsr|   c                 C   sÒ  |   |¡}|d u}	|	r4|d ur4|d }
|d }|}n |	r^|  |  |¡¡}
|  |  |¡¡}|}nv|d ur´|  |  |¡¡}
|  |  |¡¡}tj|d |
gdd}
tj|d |gdd}n |  |  |¡¡}
|  |  |¡¡}|  |¡}|d u}| jrô|
|f}t ||
 dd¡¡}| j	dks | j	dkr|j
d |
j
d  }}|r^tj|d tj|jd	 dd¡}ntj|tj|jd	 dd¡}tj|tj|jd	 dd¡}|| }|  || j d ¡}|j|jd
}| j	dkrät d||¡}|| }n4| j	dkrt d||¡}t d|
|¡}|| | }|t | j¡ }|d ur:|| }tjj|dd}|  |¡}|d urf|| }t ||¡}| dddd¡ ¡ }| ¡ d d… | jf }| |¡}|r¶||fn|f}| jrÎ||f }|S )Nr   r   r)   ©Údimr]   éþÿÿÿr   r   r}   r_   zbhld,lrd->bhlrzbhrd,lrd->bhlrr   ) r–   rŸ   r—   r˜   rJ   Úcatrš   ÚmatmulrF   rZ   rG   Útensorru   r~   r   rq   r™   rh   Útor`   ÚeinsumÚmathÚsqrtr“   r   Z
functionalZsoftmaxro   rž   Ú
contiguousrt   r”   )rw   r¡   r¢   r£   r¤   r¥   r¦   r§   Zmixed_query_layerZis_cross_attentionZ	key_layerZvalue_layerZquery_layerÚ	use_cacheZattention_scoresZquery_lengthZ
key_lengthZposition_ids_lZposition_ids_rZdistanceZpositional_embeddingZrelative_position_scoresZrelative_position_scores_queryZrelative_position_scores_keyZattention_probsZcontext_layerZnew_context_layer_shapeÚoutputsr   r   r"   r…     sn    


ÿ





zRealmSelfAttention.forward)N)NNNNNF)r8   r†   r‡   rc   rJ   r‹   rŸ   r   rŠ   r   Úboolr…   rŒ   r   r   rx   r"   r   å   s$         ø÷r   c                       s4   e Zd Z‡ fdd„Zejejejdœdd„Z‡  ZS )ÚRealmSelfOutputc                    sB   t ƒ  ¡  t |j|j¡| _tj|j|jd| _t |j	¡| _
d S ©NrY   )rb   rc   r   r•   rf   Údenserk   rl   rm   rn   ro   rv   rx   r   r"   rc   m  s    
zRealmSelfOutput.__init__©r¡   Úinput_tensorr|   c                 C   s&   |   |¡}|  |¡}|  || ¡}|S ©N©r¸   ro   rk   ©rw   r¡   rº   r   r   r"   r…   s  s    

zRealmSelfOutput.forward©r8   r†   r‡   rc   rJ   r‹   r…   rŒ   r   r   rx   r"   r¶   l  s   r¶   c                
       sv   e Zd Zd
‡ fdd„	Zdd„ Zdejeej eej eej eej ee	e	ej   ee
 e	ej dœdd	„Z‡  ZS )ÚRealmAttentionNc                    s.   t ƒ  ¡  t||d| _t|ƒ| _tƒ | _d S )N©rZ   )rb   rc   r   rw   r¶   ÚoutputÚsetÚpruned_headsr›   rx   r   r"   rc   |  s    

zRealmAttention.__init__c                 C   s²   t |ƒdkrd S t|| jj| jj| jƒ\}}t| jj|ƒ| j_t| jj|ƒ| j_t| jj	|ƒ| j_	t| j
j|dd| j
_| jjt |ƒ | j_| jj| jj | j_| j |¡| _d S )Nr   r   r¨   )rD   r   rw   r‘   r“   rÃ   r   r–   r—   r˜   rÁ   r¸   r”   Úunion)rw   ÚheadsÚindexr   r   r"   Úprune_heads‚  s    ÿzRealmAttention.prune_headsFr    c              	   C   s<   |   |||||||¡}|  |d |¡}	|	f|dd …  }
|
S )Nr   r   )rw   rÁ   )rw   r¡   r¢   r£   r¤   r¥   r¦   r§   Zself_outputsÚattention_outputr´   r   r   r"   r…   ”  s    
ù	zRealmAttention.forward)N)NNNNNF)r8   r†   r‡   rc   rÇ   rJ   r‹   r   rŠ   r   rµ   r…   rŒ   r   r   rx   r"   r¿   {  s$         ø÷r¿   c                       s0   e Zd Z‡ fdd„Zejejdœdd„Z‡  ZS )ÚRealmIntermediatec                    sB   t ƒ  ¡  t |j|j¡| _t|jt	ƒr6t
|j | _n|j| _d S r»   )rb   rc   r   r•   rf   Úintermediate_sizer¸   r5   Ú
hidden_actÚstrr	   Úintermediate_act_fnrv   rx   r   r"   rc   ®  s
    
zRealmIntermediate.__init__©r¡   r|   c                 C   s   |   |¡}|  |¡}|S r»   )r¸   rÍ   ©rw   r¡   r   r   r"   r…   ¶  s    

zRealmIntermediate.forwardr¾   r   r   rx   r"   rÉ   ­  s   rÉ   c                       s4   e Zd Z‡ fdd„Zejejejdœdd„Z‡  ZS )ÚRealmOutputc                    sB   t ƒ  ¡  t |j|j¡| _tj|j|jd| _t 	|j
¡| _d S r·   )rb   rc   r   r•   rÊ   rf   r¸   rk   rl   rm   rn   ro   rv   rx   r   r"   rc   ¾  s    
zRealmOutput.__init__r¹   c                 C   s&   |   |¡}|  |¡}|  || ¡}|S r»   r¼   r½   r   r   r"   r…   Ä  s    

zRealmOutput.forwardr¾   r   r   rx   r"   rÐ   ½  s   rÐ   c                
       st   e Zd Z‡ fdd„Zd
ejeej eej eej eej eeeej   ee	 eej dœdd„Z
dd	„ Z‡  ZS )Ú
RealmLayerc                    sr   t ƒ  ¡  |j| _d| _t|ƒ| _|j| _|j| _| jrZ| jsLt| › dƒ‚t|dd| _	t
|ƒ| _t|ƒ| _d S )Nr   z> should be used as a decoder model if cross attention is addedr[   rÀ   )rb   rc   Úchunk_size_feed_forwardÚseq_len_dimr¿   Ú	attentionrš   Úadd_cross_attentionr’   ÚcrossattentionrÉ   ÚintermediaterÐ   rÁ   rv   rx   r   r"   rc   Í  s    


zRealmLayer.__init__NFr    c              	   C   s  |d ur|d d… nd }| j |||||d}	|	d }
| jrP|	dd… }|	d }n|	dd … }d }| jrÞ|d urÞt| dƒsˆtd| › dƒ‚|d urœ|d	d … nd }|  |
||||||¡}|d }
||dd…  }|d }|| }t| j| j| j|
ƒ}|f| }| jr||f }|S )
Nr)   ©r§   r¦   r   r   r]   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Ó   )rw   r¡   r¢   r£   r¤   r¥   r¦   r§   Zself_attn_past_key_valueZself_attention_outputsrÈ   r´   Zpresent_key_valueZcross_attn_present_key_valueZcross_attn_past_key_valueZcross_attention_outputsÚlayer_outputr   r   r"   r…   Û  sP    û


ÿù	ÿ

zRealmLayer.forwardc                 C   s   |   |¡}|  ||¡}|S r»   )r×   rÁ   )rw   rÈ   Zintermediate_outputrÚ   r   r   r"   rÙ     s    
zRealmLayer.feed_forward_chunk)NNNNNF)r8   r†   r‡   rc   rJ   r‹   r   rŠ   r   rµ   r…   rÙ   rŒ   r   r   rx   r"   rÑ   Ì  s$         ø÷ArÑ   c                       s†   e Zd Z‡ fdd„Zd	ejeej eej eej eej eeeej   ee	 ee	 ee	 ee	 e
eej ef dœdd„Z‡  ZS )
ÚRealmEncoderc                    s:   t ƒ  ¡  ˆ | _t ‡ fdd„tˆ jƒD ƒ¡| _d| _d S )Nc                    s   g | ]}t ˆ ƒ‘qS r   )rÑ   )r    Ú_©rM   r   r"   Ú
<listcomp>'  ó    z)RealmEncoder.__init__.<locals>.<listcomp>F)	rb   rc   rM   r   Z
ModuleListÚrangeÚnum_hidden_layersÚlayerÚgradient_checkpointingrv   rx   rÝ   r"   rc   $  s    
 zRealmEncoder.__init__NFT)r¡   r¢   r£   r¤   r¥   Úpast_key_valuesr³   r§   Úoutput_hidden_statesÚreturn_dictr|   c              	      st  |	rdnd }ˆ rdnd }ˆ r(| j jr(dnd }| jrJ| jrJ|rJt d¡ d}|rRdnd }t| jƒD ]Î\}}|	rv||f }|d ur†|| nd }|d urš|| nd ‰| jrÖ| jrÖ‡ ‡fdd„}tj	j
 
||ƒ|||||¡}n||||||ˆˆ ƒ}|d }|r||d f7 }ˆ r`||d f }| j jr`||d	 f }q`|	r@||f }|
sbtd
d„ |||||fD ƒƒS t|||||dS )Nr   zZ`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`...Fc                    s   ‡ ‡‡fdd„}|S )Nc                     s   ˆ g | ¢ˆ‘ˆ‘R Ž S r»   r   )Úinputs)Úmoduler§   r¦   r   r"   Úcustom_forwardM  s    zKRealmEncoder.forward.<locals>.create_custom_forward.<locals>.custom_forwardr   )rè   ré   rØ   )rè   r"   Úcreate_custom_forwardL  s    z3RealmEncoder.forward.<locals>.create_custom_forwardr   r]   r   r)   c                 s   s   | ]}|d ur|V  qd S r»   r   )r    Úvr   r   r"   r#   q  s   øz'RealmEncoder.forward.<locals>.<genexpr>)Úlast_hidden_staterä   r¡   Ú
attentionsÚcross_attentions)rM   rÕ   rã   Útrainingr-   Zwarning_onceÚ	enumeraterâ   rJ   ÚutilsÚ
checkpointÚtupler
   )rw   r¡   r¢   r£   r¤   r¥   rä   r³   r§   rå   ræ   Zall_hidden_statesZall_self_attentionsZall_cross_attentionsZnext_decoder_cacheÚiZlayer_moduleZlayer_head_maskrê   Zlayer_outputsr   rØ   r"   r…   *  sv    ÿ
ú	ù

ûþûzRealmEncoder.forward)	NNNNNNFFT)r8   r†   r‡   rc   rJ   r‹   r   rŠ   r   rµ   r   r
   r…   rŒ   r   r   rx   r"   rÛ   #  s.   	         õôrÛ   c                       s0   e Zd Z‡ fdd„Zejejdœdd„Z‡  ZS )ÚRealmPoolerc                    s*   t ƒ  ¡  t |j|j¡| _t ¡ | _d S r»   )rb   rc   r   r•   rf   r¸   ZTanhÚ
activationrv   rx   r   r"   rc   ‡  s    
zRealmPooler.__init__rÎ   c                 C   s(   |d d …df }|   |¡}|  |¡}|S )Nr   )r¸   rö   )rw   r¡   Zfirst_token_tensorÚpooled_outputr   r   r"   r…   Œ  s    

zRealmPooler.forwardr¾   r   r   rx   r"   rõ   †  s   rõ   c                   @   sL   e Zd ZU dZdZejed< dZe	e
ej  ed< dZe	e
ej  ed< dS )Ú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í   )r8   r†   r‡   rˆ   rù   rJ   rŠ   Ú__annotations__r¡   r   r   rí   r   r   r   r"   rø   •  s   
rø   c                   @   s<   e Zd ZU dZdZejed< dZejed< dZ	ejed< dS )Ú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)
r8   r†   r‡   rˆ   rü   rJ   rŠ   rú   rý   rþ   r   r   r   r"   rû   °  s   
rû   c                   @   s¼   e Zd ZU dZdZejed< dZejed< dZ	ejed< dZ
ejed< dZejed< dZejed< dZ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< dS )Ú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í   )r8   r†   r‡   rˆ   r   rJ   rŠ   rú   r  r  r  Ú
BoolTensorr  r  r‰   r  r  Úint32r  r¡   r   r   rí   r   r   r   r"   rÿ   Ã  s   
#rÿ   c                   @   s,   e Zd ZU dZdZeed< dZej	ed< dS )Ú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)
r8   r†   r‡   rˆ   r  Údictrú   r  rJ   r‰   r   r   r   r"   r  õ  s   
r  c                       s$   e Zd Z‡ fdd„Zdd„ Z‡  ZS )ÚRealmPredictionHeadTransformc                    sV   t ƒ  ¡  t |j|j¡| _t|jtƒr6t	|j | _
n|j| _
tj|j|jd| _d S r·   )rb   rc   r   r•   rf   r¸   r5   rË   rÌ   r	   Útransform_act_fnrk   rl   rv   rx   r   r"   rc     s    
z%RealmPredictionHeadTransform.__init__c                 C   s"   |   |¡}|  |¡}|  |¡}|S r»   )r¸   r  rk   rÏ   r   r   r"   r…     s    


z$RealmPredictionHeadTransform.forward©r8   r†   r‡   rc   r…   rŒ   r   r   rx   r"   r    s   	r  c                       s$   e Zd Z‡ fdd„Zdd„ Z‡  ZS )ÚRealmLMPredictionHeadc                    sL   t ƒ  ¡  t|ƒ| _tj|j|jdd| _t 	t
 |j¡¡| _| j| j_d S )NF)r(   )rb   rc   r  Ú	transformr   r•   rf   re   ÚdecoderÚ	ParameterrJ   rs   r(   rv   rx   r   r"   rc     s
    

zRealmLMPredictionHead.__init__c                 C   s   |   |¡}|  |¡}|S r»   )r  r  rÏ   r   r   r"   r…   %  s    

zRealmLMPredictionHead.forwardr  r   r   rx   r"   r    s   r  c                       s$   e Zd Z‡ fdd„Zdd„ Z‡  ZS )ÚRealmOnlyMLMHeadc                    s   t ƒ  ¡  t|ƒ| _d S r»   )rb   rc   r  Úpredictionsrv   rx   r   r"   rc   ,  s    
zRealmOnlyMLMHead.__init__c                 C   s   |   |¡}|S r»   )r  )rw   Úsequence_outputÚprediction_scoresr   r   r"   r…   0  s    
zRealmOnlyMLMHead.forwardr  r   r   rx   r"   r  +  s   r  c                       s$   e Zd Z‡ fdd„Zdd„ Z‡  ZS )ÚRealmScorerProjectionc                    s>   t ƒ  ¡  t|ƒ| _t |j|j¡| _tj	|j|j
d| _	d S r·   )rb   rc   r  r  r   r•   rf   Úretriever_proj_sizer¸   rk   rl   rv   rx   r   r"   rc   6  s    

zRealmScorerProjection.__init__c                 C   s   |   |¡}|  |¡}|S r»   )r¸   rk   rÏ   r   r   r"   r…   <  s    

zRealmScorerProjection.forwardr  r   r   rx   r"   r  5  s   r  c                       s$   e Zd Z‡ fdd„Zdd„ Z‡  ZS )ÚRealmReaderProjectionc                    sX   t ƒ  ¡  || _t |j|jd ¡| _t |jd¡| _tj	|j|j
d| _t ¡ | _d S )Nr)   r   rY   )rb   rc   rM   r   r•   rf   Zspan_hidden_sizeÚdense_intermediateÚdense_outputrk   Zreader_layer_norm_epsÚlayer_normalizationZReLUÚrelurv   rx   r   r"   rc   C  s    
zRealmReaderProjection.__init__c                    s¬   ‡ fdd„}t jfdd„}ˆ  |¡}|jddd\}}||ƒ\}}}	t j|d|d	}
t j|d|d	}|
| }ˆ  |¡}ˆ  |¡}ˆ  |¡ d¡}|||	|j	d
7 }|||fS )Nc                    s„   ˆj \}‰‡‡fdd„‰ t‡ fdd„tˆjjƒD ƒŽ \}}t |d¡}t |d¡}tjˆd|d}tjˆd|d}|| }|||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                    s6   t jˆ|  d ˆ jd}t j| d ˆˆ jd}||fS )Nr   ©r~   )rJ   rq   r~   )ÚwidthZcurrent_startsZcurrent_ends)ÚmasksÚmax_sequence_lenr   r"   Ú_spans_given_widthY  s    zRRealmReaderProjection.forward.<locals>.span_candidates.<locals>._spans_given_widthc                 3   s   | ]}ˆ |d  ƒV  qdS )r   Nr   )r    Úw)r%  r   r"   r#   ^  rß   zIRealmReaderProjection.forward.<locals>.span_candidates.<locals>.<genexpr>r   r]   ©r©   rÆ   )rG   r4   rà   rM   Úmax_span_widthrJ   r«   Úindex_select)r#  rÜ   ZstartsZendsZstart_masksZ	end_masksZ
span_masks©rw   )r%  r#  r$  r"   Úspan_candidatesL  s    
"z6RealmReaderProjection.forward.<locals>.span_candidatesc                 S   s   d|   |¡ t |¡j S ©Nç      ð?©ÚtyperJ   ZfinfoÚmin©Úmaskr`   r   r   r"   Úmask_to_scorek  s    z4RealmReaderProjection.forward.<locals>.mask_to_scorer)   r]   r¨   r   r'  r_   )
rJ   Úfloat32r  Úchunkr)  r   r  r  Úsqueezer`   )rw   r¡   Ú
block_maskr+  r3  Zstart_projectionZend_projectionÚcandidate_startsÚcandidate_endsZcandidate_maskZcandidate_start_projectionsZcandidate_end_projectionsZcandidate_hiddenÚreader_logitsr   r*  r"   r…   K  s    


zRealmReaderProjection.forwardr  r   r   rx   r"   r  B  s   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                   @   s2   e Zd ZdZeZeZdZdgZ	dd„ Z
dd„ ZdS )	ÚRealmPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    Úrealmr\   c                 C   s¤   t |tjƒr:|jjjd| jjd |jdur |jj 	¡  nft |tj
ƒrz|jjjd| jjd |jdur |jj|j  	¡  n&t |tjƒr |jj 	¡  |jj d¡ dS )zInitialize the weightsg        )ÚmeanZstdNr-  )r5   r   r•   r&   rK   Znormal_rM   Zinitializer_ranger(   Zzero_rd   rX   rk   Zfill_)rw   rè   r   r   r"   Ú_init_weightsÍ  s    

z"RealmPreTrainedModel._init_weightsc                 G   sT   g }|D ]F}|du r |  d¡ q|j}t|ƒdkrD| d|d f¡}|  |¡ q|S )z.Flatten inputs' shape to (-1, input_shape[-1])Nr)   r]   )r3   rG   rD   r   )rw   rç   Zflattened_inputsr­   r€   r   r   r"   Ú_flatten_inputsÝ  s    z$RealmPreTrainedModel._flatten_inputsN)r8   r†   r‡   rˆ   r   Úconfig_classrV   Zload_tf_weightsZbase_model_prefixÚ_keys_to_ignore_on_load_missingr>  r?  r   r   r   r"   r;  Â  s   r;  c                       sD   e Zd ZdZd‡ fdd„	Zdd„ Zdd„ Zd	d
„ Zddd„Z‡  Z	S )ÚRealmBertModelz?
    Same as the original BertModel but remove docstrings.
    Tc                    sD   t ƒ  |¡ || _t|ƒ| _t|ƒ| _|r2t|ƒnd | _|  	¡  d S r»   )
rb   rc   rM   rW   r„   rÛ   Úencoderrõ   ÚpoolerÚ	post_init)rw   rM   Zadd_pooling_layerrx   r   r"   rc   ð  s    

zRealmBertModel.__init__c                 C   s   | j jS r»   ©r„   rg   r*  r   r   r"   Úget_input_embeddingsý  s    z#RealmBertModel.get_input_embeddingsc                 C   s   || j _d S r»   rF  ©rw   r˜   r   r   r"   Úset_input_embeddings   s    z#RealmBertModel.set_input_embeddingsc                 C   s*   |  ¡ D ]\}}| jj| j |¡ qdS )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)ÚitemsrC  râ   rÔ   rÇ   )rw   Zheads_to_prunerâ   rÅ   r   r   r"   Ú_prune_heads  s    zRealmBertModel._prune_headsNc                 C   sR  |d ur|n| j j}|d ur |n| j j}|d ur4|n| j j}| j jrZ|
d urP|
n| j j}
nd}
|d urx|d urxtdƒ‚n4|d urŠ| ¡ }n"|d ur¤| ¡ d d… }ntdƒ‚|\}}|d urÂ|jn|j}|	d urâ|	d d j	d nd}|d u rt
j||| f|d}|d u rZt| jdƒrH| jjd d …d |…f }| ||¡}|}nt
j|t
j|d	}|  ||¡}| j jr´|d ur´| ¡ \}}}||f}|d u r¨t
j||d}|  |¡}nd }|  || j j¡}| j|||||d
}| j||||||	|
|||d
}|d }| jd ur|  |¡nd }|s6||f|dd …  S t|||j|j|j|jdS )NFzDYou cannot specify both input_ids and inputs_embeds at the same timer]   z5You have to specify either input_ids or inputs_embedsr   r)   r!  r^   r}   )ry   r\   r^   rz   r{   )	r¢   r£   r¤   r¥   rä   r³   r§   rå   ræ   r   )rì   Úpooler_outputrä   r¡   rí   rî   )rM   r§   rå   Úuse_return_dictrš   r³   r’   rt   r~   rG   rJ   Zonesr   r„   r^   rr   rs   ru   Zget_extended_attention_maskZinvert_attention_maskZget_head_maskrá   rC  rD  r   rä   r¡   rí   rî   )rw   ry   r¢   r^   r\   r£   rz   r¤   r¥   rä   r³   r§   rå   ræ   r€   Ú
batch_sizer   r~   r{   r‚   rƒ   Zextended_attention_maskZencoder_batch_sizeZencoder_sequence_lengthrÜ   Zencoder_hidden_shapeZencoder_extended_attention_maskZembedding_outputZencoder_outputsr  r÷   r   r   r"   r…     s‚    ÿ




ûöúzRealmBertModel.forward)T)NNNNNNNNNNNNN)
r8   r†   r‡   rˆ   rc   rG  rI  rK  r…   rŒ   r   r   rx   r"   rB  ë  s&   
             òrB  z`The embedder of REALM outputting projected score that will be used to calculate relevance score.c                       s¦   e Zd ZdgZ‡ fdd„Zdd„ Zdd„ Zee 	d¡ƒe
eed	deej eej eej eej eej eej ee ee ee eeef dœ
dd„ƒƒZ‡  ZS )r=   zcls.predictions.decoder.biasc                    s0   t ƒ  |¡ t| jƒ| _t| jƒ| _|  ¡  d S r»   )rb   rc   rB  rM   r<  r  r   rE  rv   rx   r   r"   rc   €  s    zRealmEmbedder.__init__c                 C   s
   | j jjS r»   ©r<  r„   rg   r*  r   r   r"   rG  ‡  s    z"RealmEmbedder.get_input_embeddingsc                 C   s   || j j_d S r»   rO  rH  r   r   r"   rI  Š  s    z"RealmEmbedder.set_input_embeddingsúbatch_size, sequence_length©Úoutput_typer@  N)
ry   r¢   r^   r\   r£   rz   r§   rå   ræ   r|   c
                 C   sn   |	dur|	n| j j}	| j|||||||||	d	}
|
d }|  |¡}|	sX|f|
dd…  S t||
j|
jdS d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
        ```
        N©r¢   r^   r\   r£   rz   r§   rå   ræ   r   r)   é   )rù   r¡   rí   )rM   rM  r<  r   rø   r¡   rí   )rw   ry   r¢   r^   r\   r£   rz   r§   rå   ræ   Zrealm_outputsrL  rù   r   r   r"   r…     s*    !÷
ýzRealmEmbedder.forward)	NNNNNNNNN)r8   r†   r‡   rA  rc   rG  rI  r   ÚREALM_INPUTS_DOCSTRINGÚformatr   rø   Ú_CONFIG_FOR_DOCr   rJ   r‰   rŠ   rµ   r   r   r…   rŒ   r   r   rx   r"   r=   y  s6   
         ö
õr=   zoThe scorer of REALM outputting relevance scores representing the score of document candidates (before softmax).c                       s¶   e Zd ZdZd
‡ fdd„	Zee d¡ƒee	e
ddeej eej eej eej eej eej eej eej eej eej ee ee ee eee	f dœdd	„ƒƒZ‡  ZS )ÚRealmScorerz­
    Args:
        query_embedder ([`RealmEmbedder`]):
            Embedder for input sequences. If not specified, it will use the same embedder as candidate sequences.
    Nc                    s8   t ƒ  |¡ t| jƒ| _|d ur$|n| j| _|  ¡  d S r»   )rb   rc   r=   rM   ÚembedderÚquery_embedderrE  )rw   rM   rZ  rx   r   r"   rc   Ö  s    zRealmScorer.__init__rP  rQ  )ry   r¢   r^   r\   Úcandidate_input_idsÚcandidate_attention_maskÚcandidate_token_type_idsÚcandidate_inputs_embedsr£   rz   r§   rå   ræ   r|   c                 C   sà   |dur|n| j j}|du r,|
du r,tdƒ‚|du rD|du rDtdƒ‚| j|||||	|
|||d	}|  |||¡\}}}| j|||||	||||d	}|d }|d }| d| j j| j j¡}t	 
d||¡}|sÒ|||fS t|||d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
        ```Nz5You have to specify either input_ids or input_embeds.zJYou have to specify either candidate_input_ids or candidate_inputs_embeds.rS  r   r]   z
bd,bnd->bn)rü   rý   rþ   )rM   rM  r’   rZ  r?  rY  r   Únum_candidatesr  rJ   r¯   rû   )rw   ry   r¢   r^   r\   r[  r\  r]  r^  r£   rz   r§   rå   ræ   Zquery_outputsÚflattened_input_idsÚflattened_attention_maskÚflattened_token_type_idsZcandidate_outputsrý   rþ   rü   r   r   r"   r…   ß  sN    I÷ÿ
÷
ÿzRealmScorer.forward)N)NNNNNNNNNNNNN)r8   r†   r‡   rˆ   rc   r   rU  rV  r   rû   rW  r   rJ   r‰   rŠ   rµ   r   r   r…   rŒ   r   r   rx   r"   rX  Ë  sB   	
             ò
ñrX  zrThe knowledge-augmented encoder of REALM outputting masked language model logits and marginal log-likelihood loss.c                       sÎ   e Zd ZdgZ‡ fdd„Zdd„ Zdd„ Zdd	„ Zd
d„ Ze	e
 d¡ƒeeeddeej eej eej eej eej eej eej eej eej ee ee ee eeef dœdd„ƒƒZ‡  ZS )r<   zcls.predictions.decoderc                    s0   t ƒ  |¡ t| jƒ| _t| jƒ| _|  ¡  d S r»   )rb   rc   rB  rM   r<  r  r   rE  rv   rx   r   r"   rc   f  s    z!RealmKnowledgeAugEncoder.__init__c                 C   s
   | j jjS r»   rO  r*  r   r   r"   rG  l  s    z-RealmKnowledgeAugEncoder.get_input_embeddingsc                 C   s   || j j_d S r»   rO  rH  r   r   r"   rI  o  s    z-RealmKnowledgeAugEncoder.set_input_embeddingsc                 C   s
   | j jjS r»   ©r   r  r  r*  r   r   r"   Úget_output_embeddingsr  s    z.RealmKnowledgeAugEncoder.get_output_embeddingsc                 C   s   || j j_d S r»   rc  )rw   Znew_embeddingsr   r   r"   Úset_output_embeddingsu  s    z.RealmKnowledgeAugEncoder.set_output_embeddingsz+batch_size, num_candidates, sequence_lengthrQ  N)ry   r¢   r^   r\   r£   rz   rü   ÚlabelsÚmlm_maskr§   rå   ræ   r|   c                 C   sz  |dur|n| j j}|  |||¡\}}}| j|||||||
||d	}|d }|  |¡}|}d}|dur6|du rxtdƒ‚| ¡ \}}|	du ržtj|tj	d}	n|	 
tj	¡}	tdd}| d| j j¡}| d	| j j¡ d¡}|||ƒ || j j|¡ }| d¡ d¡}|| }| d	¡}t t ||	 ¡t |	¡ ¡ }|sf|f|d
d…  }|durb|f| S |S t|||j|jd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
        ```NrS  r   zZYou have to specify `relevance_score` when `labels` is specified in order to compute loss.r_   Únone)Z	reductionr]   r   r)   rT  )r   Úlogitsr¡   rí   )rM   rM  r?  r<  r   r’   rt   rJ   Z	ones_liker4  r/  r   r   re   Ztiler_  Zlog_softmaxÚ	unsqueezeÚ	logsumexpZnansumÚsumr   r¡   rí   )rw   ry   r¢   r^   r\   r£   rz   rü   rf  rg  r§   rå   ræ   r`  ra  rb  Zjoint_outputsZjoint_outputr  rþ   Zmasked_lm_lossrN  r   Zloss_fctZ
mlm_logitsZmlm_targetsZmasked_lm_log_probZcandidate_log_probZjoint_gold_log_probZmarginal_gold_log_probsrÁ   r   r   r"   r…   x  s^    9ÿ
÷

ÿ


ÿ
 üz RealmKnowledgeAugEncoder.forward)NNNNNNNNNNNN)r8   r†   r‡   rA  rc   rG  rI  rd  re  r   rU  rV  r   r   rW  r   rJ   r‰   rŠ   rµ   r   r   r…   rŒ   r   r   rx   r"   r<   ^  sJ   ÿ
            ó
òr<   zThe reader of REALM.c                       sÀ   e Zd ZddgZ‡ fdd„Zee d¡ƒee	e
ddeej eej eej eej eej eej eej eej eej eej eej ee ee ee eee	f dœd	d
„ƒƒZ‡  ZS )r6   rD  r   c                    s>   t ƒ  |¡ |j| _t|ƒ| _t|ƒ| _t|ƒ| _|  	¡  d S r»   )
rb   rc   Z
num_labelsrB  r<  r  r   r  Ú
qa_outputsrE  rv   rx   r   r"   rc   ü  s    


zRealmReader.__init__z!reader_beam_size, sequence_lengthrQ  N)ry   r¢   r^   r\   r£   rz   rü   r7  Ústart_positionsÚend_positionsÚhas_answersr§   rå   ræ   r|   c           $      C   sL  |dur|n| j j}|du r$tdƒ‚|du r4tdƒ‚| d¡| j jk rNtdƒ‚| j|||||||||d	}|d }|  ||d| j j… ¡\}}}t 	|d| j j… d¡}||7 }t 
tj|dd	j¡}t 
tj|dd	j¡}tj|d|d
}tj|d|d
}d}d}d}d}d}|	durì|
durì|durìdd„ }dd„ }| d¡} |	 d| ¡}	|
 d| ¡}
|}t |¡}!||||	d| j j… |
d| j j… d}t |¡}"|||ƒ}|| d¡| d¡ƒ}||! tj¡9 }||" tj¡9 }||  ¡ }|s*||||f|dd…  }#|dur&|||||f|# S |#S t||||||||||j|jd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.rS  r   r]   r¨   r'  c                 S   s\   t  t  t  | d¡d¡t  |d¡¡}t  t  t  |d¡d¡t  |d¡¡}t  t  ||¡d¡S )zCompute correct span.r   r]   r   )rJ   Úeqrj  r?   Úlogical_and)r8  r9  Úgold_startsÚ	gold_endsZis_gold_startZis_gold_endr   r   r"   Úcompute_correct_candidates[  s    ÿÿz7RealmReader.forward.<locals>.compute_correct_candidatesc                 S   s@   t jfdd„}t j| ||| jd dd}t j| dd}|| S )z3Loss based on the negative marginal log-likelihood.c                 S   s   d|   |¡ t |¡j S r,  r.  r1  r   r   r"   r3  k  s    zERealmReader.forward.<locals>.marginal_log_loss.<locals>.mask_to_scorer_   r]   r¨   )rJ   r4  rk  r`   )ri  Z
is_correctr3  Zlog_numeratorZlog_denominatorr   r   r"   Úmarginal_log_lossh  s    z.RealmReader.forward.<locals>.marginal_log_loss)r8  r9  rs  rt  r)   )r   r  r  r  r  r  r  r  r  r¡   rí   )rM   rM  r’   rt   r(  r<  rm  Úreader_beam_sizerJ   rj  ZargmaxÚmaxÚvaluesr)  Úclampr?   r   r/  r4  r=  rÿ   r¡   rí   )$rw   ry   r¢   r^   r\   r£   rz   rü   r7  rn  ro  rp  r§   rå   ræ   r´   r  r:  r8  r9  Zretriever_logitsZpredicted_block_indexZpredicted_candidateZpredicted_startZpredicted_endZ
total_lossr  r  r  r  ru  rv  Zignored_indexZany_retriever_correctZany_reader_correctrÁ   r   r   r"   r…     s    &÷ÿ


ü

ÿýõzRealmReader.forward)NNNNNNNNNNNNNN)r8   r†   r‡   Z"_keys_to_ignore_on_load_unexpectedrc   r   rU  rV  r   rÿ   rW  r   rJ   r‰   rŠ   r	  rµ   r   r   r…   rŒ   r   r   rx   r"   r6   ø  sF   

              ñ
ðr6   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                       sŠ   e Zd Zd‡ fdd„	Zedd„ ƒZdd„ Zee 	d¡ƒe
eed	deej eej eej eej ee eeef d
œdd„ƒƒZ‡  ZS )r:   Nc              	      s`   t ƒ  |¡ t|ƒ| _t|ƒ| _|  dt d¡j	|j
|jftjt d¡d¡ || _|  ¡  d S )NÚ	block_embr   Úcpu)rt   r`   r~   )rb   rc   r=   rY  r6   r   rp   rJ   rs   Z	new_emptyZnum_block_recordsr  r4  r~   Ú	retrieverrE  )rw   rM   r}  rx   r   r"   rc   È  s    



ýþzRealmForOpenQA.__init__c                 C   s   | j r| jjS | jjS r»   )rï   rM   Úsearcher_beam_sizerw  r*  r   r   r"   r~  Ø  s    z!RealmForOpenQA.searcher_beam_sizec                 C   s   | j  |¡| _ dS )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®   )rw   r~   r   r   r"   Úblock_embedding_toÞ  s    z!RealmForOpenQA.block_embedding_toz1, sequence_lengthrQ  )ry   r¢   r^   Ú
answer_idsræ   r|   c                 C   sà  |dur|n| j j}|dur2|jd dkr2tdƒ‚| j|||dd}|d }t d| j| | jj	¡¡}tj
|| jdd	\}	}
|
 ¡ }
tj| jd|
d
}| j|
 ¡ ||| j jd\}}}}| | jj	¡}|j tj¡j| jj	d}| ¡  |j tj¡¡ |durDtj|tj| jj	d}tj|tj| jj	d}tj|tj| jj	d}t d| ¡ | | jj	¡¡}| j|jd| j j… |jd| j j… |jd| j j… |||||dd	}|j|j }||j|jd … }|sÔ||fS t ||d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
        ```Nr   r   z'The batch_size of the inputs must be 1.T)ry   r^   r¢   ræ   z	BD,QD->QBr]   )Úkr©   r'  )Ú
max_lengthr!  r}   zD,BD->B)	ry   r¢   r^   rü   r7  rp  rn  ro  ræ   )r  r  )!rM   rM  rG   r’   rY  rJ   r¯   r{  r®   r~   Ztopkr~  r6  r)  r}  r|  Zreader_seq_lenr   Zspecial_tokens_maskr/  rµ   Zlogical_not_Zlogical_and_r^   r­   ru   ry   rw  r¢   r  r  r  r  )rw   ry   r¢   r^   r€  ræ   Zquestion_outputsZquestion_projectionZbatch_scoresrÜ   Zretrieved_block_idsZretrieved_block_embrp  r  r  Zconcat_inputsr7  Zretrieved_logitsr  Zpredicted_blockr  r   r   r"   r…   è  sV    %ÿÿ
ÿ÷þzRealmForOpenQA.forward)N)NNNN)r8   r†   r‡   rc   Úpropertyr~  r  r   ÚREALM_FOR_OPEN_QA_DOCSTRINGrV  r   r  rW  r   rJ   r‰   rŠ   rµ   r   r   r…   rŒ   r   r   rx   r"   r:   Ã  s$   


    ú
ùr:   )Grˆ   r°   r/   Zdataclassesr   Útypingr   r   r   rJ   r   Ztorch.nnr   Zactivationsr	   Zmodeling_outputsr
   r   r   r   Zmodeling_utilsr   Zpytorch_utilsr   r   r   rñ   r   r   r   r   Zconfiguration_realmr   Z
get_loggerr8   r-   Z_EMBEDDER_CHECKPOINT_FOR_DOCZ_ENCODER_CHECKPOINT_FOR_DOCZ_SCORER_CHECKPOINT_FOR_DOCrW  Z#REALM_PRETRAINED_MODEL_ARCHIVE_LISTrV   ÚModulerW   r   r¶   r¿   rÉ   rÐ   rÑ   rÛ   rõ   rø   rû   rÿ   r  r  r  r  r  r  ZREALM_START_DOCSTRINGrU  r;  rB  r=   rX  r<   r6   r„  r:   r   r   r   r"   Ú<module>   s”   
l? 2Wc1
C2) þNþ ý  *!þ