
    ^j                        d Z ddlZddlmZ ddlZddlmZ ddlmZ ddlm	Z
 ddlmZ dd	lmZ dd
lmZ ddlmZmZmZ ddlmZ ddlmZ ddlmZ ddlmZmZmZmZm Z  ddl!m"Z" ddl#m$Z$m%Z% ddl&m'Z'  e jP                  e)      Z* ed      e G d de                    Z+ G d dejX                        Z- G d dejX                        Z. G d dejX                        Z/ G d dejX                        Z0 G d  d!ejX                        Z1 G d" d#ejX                        Z2 G d$ d%ejX                        Z3 G d& d'ejX                        Z4 G d( d)ejX                        Z5 G d* d+e      Z6 G d, d-ejX                        Z7 G d. d/ejX                        Z8e G d0 d1e             Z9 G d2 d3e9      Z:e G d4 d5e9             Z;e G d6 d7e9             Z< ed8       G d9 d:e9             Z= ed;       G d< d=e9             Z>g d>Z?y)?zPyTorch Bros model.    N)	dataclass)nn)CrossEntropyLoss   )initialization)ACT2FN)create_bidirectional_mask)GradientCheckpointingLayer)"BaseModelOutputWithCrossAttentions,BaseModelOutputWithPoolingAndCrossAttentionsTokenClassifierOutput)PreTrainedModel)Unpack)apply_chunking_to_forward)ModelOutputTransformersKwargsauto_docstringcan_return_tuplelogging)merge_with_config_defaults)OutputRecordercapture_outputs   )
BrosConfigz@
    Base class for outputs of token classification models.
    )custom_introc                       e Zd ZU dZdZej                  dz  ed<   dZej                  dz  ed<   dZ	ej                  dz  ed<   dZ
eej                     dz  ed<   dZeej                     dz  ed<   y)BrosSpadeOutputa  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Classification loss.
    initial_token_logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.num_labels)`):
        Classification scores for entity initial tokens (before SoftMax).
    subsequent_token_logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, sequence_length+1)`):
        Classification scores for entity sequence tokens (before SoftMax).
    Nlossinitial_token_logitssubsequent_token_logitshidden_states
attentions)__name__
__module____qualname____doc__r   torchFloatTensor__annotations__r   r    r!   tupler"        q/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/bros/modeling_bros.pyr   r   ,   s~     &*D%

d
")59%++d298<U..5<59M5**+d2926Je''(4/6r,   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )BrosPositionalEmbedding1Dc                     t         |           |j                  | _        ddt        j                  d| j                  d      | j                  z  z  z  }| j                  d|       y )Nr   '                 @inv_freq)super__init__dim_bbox_sinusoid_emb_1dr'   arangeregister_buffer)selfconfigr4   	__class__s      r-   r6   z"BrosPositionalEmbedding1D.__init__F   s^    (.(G(G%ell3(E(EsKdNkNkkl
 	Z2r,   pos_seqreturnc                    |j                         }|\  }}}|j                  |||d      | j                  j                  ddd| j                  dz        z  }t	        j
                  |j                         |j                         gd      }|S )Nr      dim)sizeviewr4   r7   r'   catsincos)r:   r=   seq_sizeb1b2b3sinusoid_inppos_embs           r-   forwardz!BrosPositionalEmbedding1D.forwardP   s    <<>
B||BB2T]]5G5G1aQUQnQnrsQs5tt))\--/1A1A1CD"Mr,   r#   r$   r%   r6   r'   TensorrO   __classcell__r<   s   @r-   r/   r/   C   s#    3u||  r,   r/   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )BrosPositionalEmbedding2Dc                     t         |           |j                  | _        t        |      | _        t        |      | _        y N)r5   r6   dim_bboxr/   	x_pos_emb	y_pos_embr:   r;   r<   s     r-   r6   z"BrosPositionalEmbedding2D.__init__Y   s1    26:26:r,   bboxr>   c                    g }t        | j                        D ]U  }|dz  dk(  r&|j                  | j                  |d|f                1|j                  | j	                  |d|f                W t        j                  |d      }|S )Nr@   r   .rA   rB   )rangerX   appendrY   rZ   r'   rF   )r:   r\   stackibbox_pos_embs        r-   rO   z!BrosPositionalEmbedding2D.forward`   s|    t}}% 	;A1uzT^^DaL9:T^^DaL9:		;
 yyB/r,   rP   rS   s   @r-   rU   rU   X   s#    ;ELL U\\ r,   rU   c                   >     e Zd Z fdZdej
                  fdZ xZS )BrosBboxEmbeddingsc                     t         |           t        |      | _        t	        j
                  |j                  |j                  d      | _        y )NF)bias)	r5   r6   rU   bbox_sinusoid_embr   Lineardim_bbox_sinusoid_emb_2ddim_bbox_projectionbbox_projectionr[   s     r-   r6   zBrosBboxEmbeddings.__init__l   s=    !:6!B!yy)H)H&JdJdkpqr,   r\   c                     |j                  dd      }|d d d d d d d f   |d d d d d d d f   z
  }| j                  |      }| j                  |      }|S )Nr   r   )	transposerg   rk   )r:   r\   bbox_tbbox_posrb   s        r-   rO   zBrosBboxEmbeddings.forwardq   s\    1%$1a-(6!T1a-+@@--h7++L9r,   rP   rS   s   @r-   rd   rd   k   s    r
ELL r,   rd   c                        e Zd ZdZ fdZ	 	 	 	 d
dej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  f
d	Z xZS )BrosTextEmbeddingszGConstruct 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                        | _        | j#                  dt%        j&                  |j                        j)                  d             | j#                  dt%        j*                  | j,                  j/                         t$        j0                  | j,                  j2                        d       y )	N)padding_idxepsposition_idsr   rA   token_type_idsdtypedeviceF)
persistent)r5   r6   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dropoutr9   r'   r8   expandzerosrv   rD   longr{   r[   s     r-   r6   zBrosTextEmbeddings.__init__}   s#   !||F,=,=v?Q?Q_e_r_rs#%<<0N0NPVPbPb#c %'\\&2H2H&J\J\%]"f&8&8f>S>STzz&"<"<=^U\\&:X:X-Y-`-`ah-ijKK!!&&(jj((//
  	 	
r,   N	input_idsrx   rv   inputs_embedsr>   c                 6   ||j                         }n|j                         d d }|d   }|| j                  d d d |f   }|st        | d      r-| j                  d d d |f   }|j	                  |d   |      }|}n:t        j                  |t
        j                  | j                  j                        }|| j                  |      }| j                  |      }	||	z   }
| j                  |      }|
|z  }
| j                  |
      }
| j                  |
      }
|
S )NrA   r   rx   r   ry   )rD   rv   hasattrrx   r   r'   r   r   r{   r   r   r   r   r   )r:   r   rx   rv   r   input_shape
seq_lengthbuffered_token_type_ids buffered_token_type_ids_expandedr   
embeddingsr   s               r-   rO   zBrosTextEmbeddings.forward   s/     #..*K',,.s3K ^
,,Q^<L!t-.*.*=*=a*n*M'3J3Q3QR]^_R`bl3m0!A!&[

SWSdSdSkSk!l  00;M $ : :> J"%::
"66|D))
^^J/
\\*-
r,   )NNNN)	r#   r$   r%   r&   r6   r'   rQ   rO   rR   rS   s   @r-   rq   rq   z   sv    Q
. *..2,0-1#<<$&# t+# llT)	#
 ||d*# 
#r,   rq   c                        e Zd Z fdZ	 	 	 d
dej
                  dej
                  dej
                  dz  dej
                  dz  dej
                  dz  deej
                     fd	Z xZS )BrosSelfAttentionc                    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                        | _        |j"                  | _        y )Nr   embedding_sizezThe hidden size (z6) is not a multiple of the number of attention heads ())r5   r6   r   num_attention_headsr   
ValueErrorintattention_head_sizeall_head_sizer   rh   querykeyvaluer   attention_probs_dropout_probr   
is_decoderr[   s     r-   r6   zBrosSelfAttention.__init__   s)    : ::a?PVXhHi#F$6$6#7 8 445Q8 
 $*#=#= #&v'9'9F<V<V'V#W !558P8PPYYv1143E3EF
99V//1C1CDYYv1143E3EF
zz&"E"EF ++r,   Nr!   rb   attention_maskencoder_hidden_statesencoder_attention_maskr>   c                    |j                   d   d| j                  | j                  f}| j                  |      j	                  |      j                  dd      }|d u}|rc| j                  |      j	                  |      j                  dd      }	| j                  |      j	                  |      j                  dd      }
|}n`| j                  |      j	                  |      j                  dd      }	| j                  |      j	                  |      j                  dd      }
t        j                  ||	j                  dd            }|j                   \  }}}}|j	                  ||||      }|j                  g d      }t        j                  d||f      }||z   }|t        j                  | j                        z  }|||z   } t        j                  d      |      }| j!                  |      }t        j                  ||
      }|j                  dddd	      j#                         }|j%                         d d | j&                  fz   } |j                  | }||fS )
Nr   rA   r   r@   )r@   r   r   r   zbnid,bijd->bnijrB   r   )shaper   r   r   rE   rm   r   r   r'   matmulpermuteeinsummathsqrtr   Softmaxr   
contiguousrD   r   )r:   r!   rb   r   r   r   hidden_shapequery_layeris_cross_attention	key_layervalue_layerattention_scores
batch_sizen_headr   d_headbbox_pos_scoresattention_probscontext_layernew_context_layer_shapes                       r-   rO   zBrosSelfAttention.forward   sG    &++A.D4L4LdNfNfgjj/44\BLLQPQR
 3$>!67<<\JTTUVXYZI**%:;@@NXXYZ\]^K3N/44\BLLQPQRI**]388FPPQRTUVK !<<Y5H5HR5PQ 2=1B1B.
FJ#((ZVT#++L9,,'8;:UV+o=+dii8P8P.QQ%/.@ -"**,-=> ,,7_kB%--aAq9DDF"/"4"4"6s";t?Q?Q>S"S***,CDo--r,   NNN)	r#   r$   r%   r6   r'   rQ   r*   rO   rR   rS   s   @r-   r   r      s~    ,0 /3596:5.||5. ll5. t+	5.
  %||d25. !&t 35. 
u||	5.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 )BrosSelfOutputc                 (   t         |           t        j                  |j                  |j                        | _        t        j                  |j                  |j                        | _        t        j                  |j                        | _
        y Nrt   )r5   r6   r   rh   r   denser   r   r   r   r   r[   s     r-   r6   zBrosSelfOutput.__init__  s`    YYv1163E3EF
f&8&8f>S>STzz&"<"<=r,   r!   input_tensorr>   c                 r    | j                  |      }| j                  |      }| j                  ||z         }|S rW   r   r   r   r:   r!   r   s      r-   rO   zBrosSelfOutput.forward  7    

=1]3}|'CDr,   rP   rS   s   @r-   r   r     1    >U\\  RWR^R^ r,   r   c                        e Zd Z fdZ	 	 	 d
dej
                  dej
                  dej
                  dz  dej
                  dz  dej
                  dz  dej
                  fd	Z xZS )BrosAttentionc                 b    t         |           t        |      | _        t	        |      | _        y rW   )r5   r6   r   r:   r   outputr[   s     r-   r6   zBrosAttention.__init__  s&    %f-	$V,r,   Nr!   rb   r   r   r   r>   c                 `    |}| j                  |||||      \  }}| j                  ||      }|S )Nrb   r   r   r   )r:   r   )r:   r!   rb   r   r   r   residual_s           r-   rO   zBrosAttention.forward  sH     !99%)"7#9 % 
q M8<r,   r   rP   rS   s   @r-   r   r     sy    - /3596:|| ll t+	
  %||d2 !&t 3 
r,   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )BrosIntermediatec                    t         |           t        j                  |j                  |j
                        | _        t        |j                  t              rt        |j                     | _        y |j                  | _        y rW   )r5   r6   r   rh   r   intermediate_sizer   
isinstance
hidden_actstrr   intermediate_act_fnr[   s     r-   r6   zBrosIntermediate.__init__0  s]    YYv1163K3KL
f''-'-f.?.?'@D$'-'8'8D$r,   r!   r>   c                 J    | j                  |      }| j                  |      }|S rW   )r   r   )r:   r!   s     r-   rO   zBrosIntermediate.forward8  s&    

=100?r,   rP   rS   s   @r-   r   r   /  s#    9U\\ ell 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 )
BrosOutputc                 (   t         |           t        j                  |j                  |j
                        | _        t        j                  |j
                  |j                        | _        t        j                  |j                        | _        y r   )r5   r6   r   rh   r   r   r   r   r   r   r   r   r[   s     r-   r6   zBrosOutput.__init__?  s`    YYv779K9KL
f&8&8f>S>STzz&"<"<=r,   r!   r   r>   c                 r    | j                  |      }| j                  |      }| j                  ||z         }|S rW   r   r   s      r-   rO   zBrosOutput.forwardE  r   r,   rP   rS   s   @r-   r   r   >  r   r,   r   c                        e Zd Z fdZ	 	 	 ddej
                  dej
                  dej                  dz  dej                  dz  dej                  dz  dee   d	ej
                  fd
Z	d Z
 xZS )	BrosLayerc                 b   t         |           |j                  | _        d| _        t	        |      | _        |j                  | _        |j                  | _        | j                  r*| j                  st        |  d      t	        |      | _	        t        |      | _        t        |      | _        y )Nr   z> should be used as a decoder model if cross attention is added)r5   r6   chunk_size_feed_forwardseq_len_dimr   	attentionr   add_cross_attention	Exceptioncrossattentionr   intermediater   r   r[   s     r-   r6   zBrosLayer.__init__M  s    '-'E'E$&v. ++#)#=#= ##??4&(f ghh"/"7D,V4 (r,   Nr!   rb   r   r   r   kwargsr>   c                    | j                  |||      }| j                  r7|5t        | d      rt        d|  d       | j                  |f|||d|\  }}t        | j                  | j                  | j                  |      }|S )N)rb   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   r   r   feed_forward_chunkr   r   )r:   r!   rb   r   r   r   r   r   s           r-   rO   zBrosLayer.forward[  s     %) ' 
 ??4@t-.=dV  Dd  e   3t22 -&;'=	 
  M1 2##((	
 r,   c                 L    | j                  |      }| j                  ||      }|S rW   )r   r   )r:   attention_outputintermediate_outputlayer_outputs       r-   r   zBrosLayer.feed_forward_chunk  s,    "//0@A{{#68HIr,   r   )r#   r$   r%   r6   r'   rQ   r(   r   r   rO   r   rR   rS   s   @r-   r   r   L  s    )$ 48:>;?$||$ ll$ ))D0	$
  %0047$ !& 1 1D 8$ +,$ 
$Lr,   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )
BrosPoolerc                     t         |           t        j                  |j                  |j                        | _        t        j                         | _        y rW   )r5   r6   r   rh   r   r   Tanh
activationr[   s     r-   r6   zBrosPooler.__init__  s9    YYv1163E3EF
'')r,   r!   r>   c                 \    |d d df   }| j                  |      }| j                  |      }|S )Nr   )r   r   )r:   r!   first_token_tensorpooled_outputs       r-   rO   zBrosPooler.forward  s6     +1a40

#566r,   rP   rS   s   @r-   r   r     s#    $
U\\ ell r,   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )BrosRelationExtractorc                 R   t         |           |j                  | _        |j                  | _        |j                  | _        |j                  | _        t        j                  | j                        | _	        t        j                  | j                  | j                  | j
                  z        | _        t        j                  | j                  | j                  | j
                  z        | _        t        j                  t        j                  d| j                              | _        y )Nr   )r5   r6   n_relationsr   backbone_hidden_sizehead_hidden_sizeclassifier_dropout_probr   r   droprh   r   r   	Parameterr'   r   
dummy_noder[   s     r-   r6   zBrosRelationExtractor.__init__  s    !--$*$6$6! & 2 2'-'E'E$JJt;;<	YYt88$:J:JTMbMb:bc
99T668H8H4K`K`8`a,,u{{1d6O6O'PQr,   r   r   c           	         | j                  | j                  |            }| j                  j                  d      j	                  d|j                  d      d      }t        j                  ||gd      }| j                  | j                  |            }|j                  |j                  d      |j                  d      | j                  | j                        }|j                  |j                  d      |j                  d      | j                  | j                        }t        j                  |j                  dddd      |j                  dddd            }|S )Nr   r   axisr@   r   )r   r  r  	unsqueezerepeatrD   r'   rF   r   rE   r   r   r   r   )r:   r   r   	dummy_vecrelation_scores        r-   rO   zBrosRelationExtractor.forward  s   jj;!78OO--a0779>>!;LaP	IIy)41=	HHTYYy12	!&&Q!1!1!!4d6F6FH]H]
 NN9>>!#4innQ6GIYIY[_[p[pq	1a+Y->->q!Q-J
 r,   rP   rS   s   @r-   r   r     s$    R5<< ELL r,   r   c                        e Zd ZU eed<   dZe eedd       eedd      dZ	 e
j                         dej                  f fd	       Z xZS )
BrosPreTrainedModelr;   brosr   r   )index
layer_namer   )r!   r"   cross_attentionsmodulec                    t         |   |       | j                  j                  }t	        |t
              r"t        j                  |j                  |       yt	        |t              ryt        j                  |j                  t        j                  |j                  j                  d         j                  d             t        j                   |j"                         yt	        |t$              rUddt        j                  d|j&                  d      |j&                  z  z  z  }t        j                  |j(                  |       yy)	zInitialize the weights)stdrA   rw   r   r1   r2   r3   N)r5   _init_weightsr;   initializer_ranger   r   initnormal_r  rq   copy_rv   r'   r8   r   r   zeros_rx   r/   r7   r4   )r:   r  r  r4   r<   s       r-   r  z!BrosPreTrainedModel._init_weights  s     	f%kk++f34LL**4 23JJv**ELL9L9L9R9RSU9V,W,^,^_f,ghKK--. 9:%,,sF,K,KSQTZTsTsstH JJv1	 ;r,   )r#   r$   r%   r   r)   base_model_prefixr   r   r   _can_record_outputsr'   no_gradr   Moduler  rR   rS   s   @r-   r  r    s\    "$%6aKX*+<ARbc U]]_2BII 2 2r,   r  c                        e Zd Z fdZee	 	 	 ddej                  dej                  dej                  dz  dej                  dz  dej                  dz  de	e
   d	eej                     ez  fd
              Z xZS )BrosEncoderc                     t         |   |       t        j                  t	        |j
                        D cg c]  }t        |       c}      | _        | j                          y c c}w rW   )	r5   r6   r   
ModuleListr^   num_hidden_layersr   layer	post_init)r:   r;   r   r<   s      r-   r6   zBrosEncoder.__init__  sK     ]]uVE]E]?^#_!If$5#_`
 $`s   A&Nr!   rb   r   r   r   r   r>   c           	      X    | j                   D ]  } ||f||||d|} t        |      S )Nr   )last_hidden_state)r#  r   )r:   r!   rb   r   r   r   r   layer_modules           r-   rO   zBrosEncoder.forward  sQ     !JJ 	L()-&;'= M	 2+
 	
r,   r   )r#   r$   r%   r6   r   r   r'   rQ   r(   r   r   r*   r   rO   rR   rS   s   @r-   r  r    s    
  
 48:>;?
||
 ll
 ))D0	

  %0047
 !& 1 1D 8
 +,
 
u||	A	A
   
r,   r  c                   x    e Zd Zd fd	Zd Zd Zee	 	 	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  dej                  dz  d	ej                  dz  d
ej                  dz  dej                  dz  dej                  dz  de
e   deej                     ez  fd              Z xZS )	BrosModelc                     t         |   |       || _        t        |      | _        t        |      | _        t        |      | _        |rt        |      nd| _
        | j                          y)zv
        add_pooling_layer (bool, *optional*, defaults to `True`):
            Whether to add a pooling layer
        N)r5   r6   r;   rq   r   rd   bbox_embeddingsr  encoderr   poolerr$  )r:   r;   add_pooling_layerr<   s      r-   r6   zBrosModel.__init__  sX    
 	 ,V41&9"6*,=j(4r,   c                 .    | j                   j                  S rW   r   r   )r:   s    r-   get_input_embeddingszBrosModel.get_input_embeddings  s    ...r,   c                 &    || j                   _        y rW   r0  )r:   r   s     r-   set_input_embeddingszBrosModel.set_input_embeddings	  s    */'r,   Nr   r\   r   rx   rv   r   r   r   r   r>   c	                    |du |duz  rt        d      |t        d      | j                  ||||      }
|
j                  dd }|
j                  }|t	        j
                  ||      }t        | j                  |
|      }|t        | j                  |
||      }|j                  d   d	k(  r|ddddg d
f   }|| j                  j                  z  }| j                  |      } | j                  |
f||||d|	}|d   }| j                  | j                  |      nd}t        |||j                  |j                  |j                        S )a  
        bbox ('torch.FloatTensor' of shape '(batch_size, num_boxes, 4)'):
            Bounding box coordinates for each token in the input sequence. Each bounding box is a list of four values
            (x1, y1, x2, y2), where (x1, y1) is the top left corner, and (x2, y2) is the bottom right corner of the
            bounding box.

        Examples:

        ```python
        >>> import torch
        >>> from transformers import BrosProcessor, BrosModel

        >>> processor = BrosProcessor.from_pretrained("jinho8345/bros-base-uncased")

        >>> model = BrosModel.from_pretrained("jinho8345/bros-base-uncased")

        >>> encoding = processor("Hello, my dog is cute", add_special_tokens=False, return_tensors="pt")
        >>> bbox = torch.tensor([[[0, 0, 1, 1]]]).repeat(1, encoding["input_ids"].shape[-1], 1)
        >>> encoding["bbox"] = bbox

        >>> outputs = model(**encoding)
        >>> last_hidden_states = outputs.last_hidden_state
        ```Nz:You must specify exactly one of input_ids or inputs_embedszYou have to specify bbox)r   rv   rx   r   rA   )r{   )r;   r   r   )r;   r   r   r      )r   r   r@   r   r@   r   r   r   r   r   )r&  pooler_outputr!   r"   r  )r   r   r   r{   r'   onesr	   r;   
bbox_scaler+  r,  r-  r   r!   r"   r  )r:   r   r\   r   rx   rv   r   r   r   r   embedding_outputr   r{   scaled_bboxbbox_position_embeddingsencoder_outputssequence_outputr   s                     r-   rO   zBrosModel.forward  s   J -t";<YZZ<788??%)'	 + 
 ',,Sb1!((!"ZZFCN2;;*)
 "-%>{{.5&;	&" ::b>Q1667DT[[333#'#7#7#D >Jdll?
1)"7#9?
 ?
 *!,8<8OO4UY;-')77&11,==
 	
r,   )TNNNNNNNN)r#   r$   r%   r6   r1  r3  r   r   r'   rQ   r   r   r*   r   rO   rR   rS   s   @r-   r)  r)    s    /0  *.$(.2.2,0-1596:[
<<$&[
 llT![
 t+	[

 t+[
 llT)[
 ||d*[
  %||d2[
 !&t 3[
 +,[
 
u||	K	K[
  [
r,   r)  c                   p    e Zd ZdgZ fdZee	 	 	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  d	ej                  dz  d
ej                  dz  dej                  dz  de	e
   deej                     ez  fd              Z xZS )BrosForTokenClassificationr-  c                 `   t         |   |       |j                  | _        t        |      | _        t        |d      r|j                  n|j                  }t        j                  |      | _
        t        j                  |j                  |j                        | _        | j                          y Nclassifier_dropout)r5   r6   
num_labelsr)  r  r   rC  r   r   r   r   rh   r   
classifierr$  r:   r;   rC  r<   s      r-   r6   z#BrosForTokenClassification.__init__p  s      ++f%	)09M)NF%%TZTnTn 	 zz"45))F$6$68I8IJr,   Nr   r\   r   bbox_first_token_maskrx   rv   r   labelsr   r>   c	           	          | j                   |f|||||d|	}
|
d   }| j                  |      }| j                  |      }d}|t               }|J|j	                  d      } ||j	                  d| j
                        |   |j	                  d      |         }n2 ||j	                  d| j
                        |j	                  d            }t        |||
j                  |
j                        S )a  
        bbox ('torch.FloatTensor' of shape '(batch_size, num_boxes, 4)'):
            Bounding box coordinates for each token in the input sequence. Each bounding box is a list of four values
            (x1, y1, x2, y2), where (x1, y1) is the top left corner, and (x2, y2) is the bottom right corner of the
            bounding box.
        bbox_first_token_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Mask to indicate the first token of each bounding box. Mask values selected in `[0, 1]`:

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

        Examples:

        ```python
        >>> import torch
        >>> from transformers import BrosProcessor, BrosForTokenClassification

        >>> processor = BrosProcessor.from_pretrained("jinho8345/bros-base-uncased")

        >>> model = BrosForTokenClassification.from_pretrained("jinho8345/bros-base-uncased")

        >>> encoding = processor("Hello, my dog is cute", add_special_tokens=False, return_tensors="pt")
        >>> bbox = torch.tensor([[[0, 0, 1, 1]]]).repeat(1, encoding["input_ids"].shape[-1], 1)
        >>> encoding["bbox"] = bbox

        >>> outputs = model(**encoding)
        ```)r\   r   rx   rv   r   r   NrA   r   logitsr!   r"   )	r  r   rE  r   rE   rD  r   r!   r"   )r:   r   r\   r   rG  rx   rv   r   rH  r   outputsr=  rK  r   loss_fcts                  r-   rO   z"BrosForTokenClassification.forward}  s	   R AJ		A
))%'A
 A
 "!*,,71')H$0(=(B(B2(F%KKDOO45JKV[[Y[_]rMs  B @&++b/R$!//))	
 	
r,   r>  r#   r$   r%   "_keys_to_ignore_on_load_unexpectedr6   r   r   r'   rQ   r   r   r*   r   rO   rR   rS   s   @r-   r@  r@  l  s   *3&  *.$(.259.2,0-1&*F
<<$&F
 llT!F
 t+	F

  %||d2F
 t+F
 llT)F
 ||d*F
 t#F
 +,F
 
u||	4	4F
  F
r,   r@  a  
    Bros Model with a token classification head on top (initial_token_layers and subsequent_token_layer on top of the
    hidden-states output) e.g. for Named-Entity-Recognition (NER) tasks. The initial_token_classifier is used to
    predict the first token of each entity, and the subsequent_token_classifier is used to predict the subsequent
    tokens within an entity. Compared to BrosForTokenClassification, this model is more robust to serialization errors
    since it predicts next token from one token.
    c                       e Zd ZdgZ fdZee	 	 	 	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  d	ej                  dz  d
ej                  dz  dej                  dz  dej                  dz  de	e
   deej                     ez  fd              Z xZS )!BrosSpadeEEForTokenClassificationr-  c           	      f   t         |   |       || _        |j                  | _        |j                  | _        |j
                  | _        t        |      | _        t        |d      r|j                  n|j                  }t        j                  t        j                  |      t        j                  |j
                  |j
                        t        j                  |      t        j                  |j
                  |j                              | _        t#        |      | _        | j'                          y rB  )r5   r6   r;   rD  r   r   r   r)  r  r   rC  r   r   
Sequentialr   rh   initial_token_classifierr   subsequent_token_classifierr$  rF  s      r-   r6   z*BrosSpadeEEForTokenClassification.__init__  s      ++!--$*$6$6!f%	)09M)NF%%TZTnTn 	
 )+JJ)*IIf((&*<*<=JJ)*IIf((&*;*;<	)
% ,A+H(r,   Nr   r\   r   rG  rx   rv   r   initial_token_labelssubsequent_token_labelsr   r>   c
           
      d    | j                   d
||||||d|
}|d   }|j                  dd      j                         }| j                  |      j                  dd      j                         }| j	                  ||      j                  d      }d|z
  }|j                  \  }}|j                  }t        j                  |t        j                  |dg|j                  |      gd      j                         }|j                  |dddddf   t        j                  |j                        j                        }t        j                   ||dz         j#                  |t        j                        }|j                  |dddddf   t        j                  |j                        j                        }|j%                  d      j                         }d}||	t'               }|j%                  d      }|;|j%                  d      } ||j%                  d| j(                        |   ||         }n# ||j%                  d| j(                        |      }|	j%                  d      }	 ||j%                  d|dz         |   |	|         }||z   }t+        ||||j,                  |j.                  	      S )a>  
        bbox ('torch.FloatTensor' of shape '(batch_size, num_boxes, 4)'):
            Bounding box coordinates for each token in the input sequence. Each bounding box is a list of four values
            (x1, y1, x2, y2), where (x1, y1) is the top left corner, and (x2, y2) is the bottom right corner of the
            bounding box.
        bbox_first_token_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Mask to indicate the first token of each bounding box. Mask values selected in `[0, 1]`:

            - 1 for tokens that are **not masked**,
            - 0 for tokens that are **masked**.
        initial_token_labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for the initial token classification.
        subsequent_token_labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for the subsequent token classification.

        Examples:

        ```python
        >>> import torch
        >>> from transformers import BrosProcessor, BrosSpadeEEForTokenClassification

        >>> processor = BrosProcessor.from_pretrained("jinho8345/bros-base-uncased")

        >>> model = BrosSpadeEEForTokenClassification.from_pretrained("jinho8345/bros-base-uncased")

        >>> encoding = processor("Hello, my dog is cute", add_special_tokens=False, return_tensors="pt")
        >>> bbox = torch.tensor([[[0, 0, 1, 1]]]).repeat(1, encoding["input_ids"].shape[-1], 1)
        >>> encoding["bbox"] = bbox

        >>> outputs = model(**encoding)
        ```r   r\   r   rx   rv   r   r   r   ry   r  Nr{   rz   rA   )r   r   r    r!   r"   r+   )r  rm   r   rT  rU  squeezer   r{   r'   rF   r   rz   boolmasked_fillfinfomineyetorE   r   rD  r   r!   r"   )r:   r   r\   r   rG  rx   rv   r   rV  rW  r   rL  last_hidden_statesr   r    inv_attention_maskr   max_seq_lengthr{   invalid_token_maskself_token_masksubsequent_token_maskr   rM  initial_token_losssubsequent_token_losss                             r-   rO   z)BrosSpadeEEForTokenClassification.forward  s   \ AJ		 A
))%'A
 A
 %QZ/99!Q?JJL#<<=OPZZ[\^_`kkm"&"B"BCUWi"j"r"rst"u /%7%=%="
N#**"YYj!_DVD\D\ek!lmtu

$& 	 #:"E"Eq$z*EKK8O8U8U,V,Z,Z#
  ))NNQ4FGJJRX`e`j`jJk"9"E"ED!QJ'5L5R5R)S)W)W#
 !/ 3 3B 7 < < >+0G0S')H $8#<#<R#@ $0(=(B(B2(F%%-(--b$//BCXY()>?&"
 &..B.G.GDOO.\^r%s"&=&B&B2&F#$,',,R!1CDEZ['(=>%!
 &(==D!5$;!//))
 	
r,   )	NNNNNNNNN)r#   r$   r%   rO  r6   r   r   r'   rQ   r   r   r*   r   rO   rR   rS   s   @r-   rQ  rQ    s)    +4&2  *.$(.259.2,0-1487;h
<<$&h
 llT!h
 t+	h

  %||d2h
 t+h
 llT)h
 ||d*h
 $llT1h
 "'!4h
 +,h
 
u||		.h
  h
r,   rQ  z
    Bros Model with a token classification head on top (a entity_linker layer on top of the hidden-states output) e.g.
    for Entity-Linking. The entity_linker is used to predict intra-entity links (one entity to another entity).
    c                   p    e Zd ZdgZ fdZee	 	 	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  d	ej                  dz  d
ej                  dz  dej                  dz  de	e
   deej                     ez  fd              Z xZS )!BrosSpadeELForTokenClassificationr-  c                 @   t         |   |       || _        |j                  | _        |j                  | _        |j
                  | _        t        |      | _        t        |d      r|j                  n|j                   t        |      | _        | j                          y rB  )r5   r6   r;   rD  r   r   r   r)  r  r   rC  r   r   entity_linkerr$  r[   s     r-   r6   z*BrosSpadeELForTokenClassification.__init__c  s      ++!--$*$6$6!f%	&-f6J&K	"	"QWQkQk26:r,   Nr   r\   r   rG  rx   rv   r   rH  r   r>   c	           
          | j                   d
||||||d|	}
|
d   }|j                  dd      j                         }| j                  ||      j	                  d      }d}|et               }|j                  \  }}|j                  }t        j                  ||dz         j                  |t        j                        }|j                  d      }t        j                  | t        j                  |dgt        j                  |      gd      }|j                  |dddddf   t        j                   |j"                        j$                        }|j                  |dddddf   t        j                   |j"                        j$                        } ||j                  d|dz         |   |j                  d      |         }t'        |||
j(                  |
j*                  	      S )a  
        bbox ('torch.FloatTensor' of shape '(batch_size, num_boxes, 4)'):
            Bounding box coordinates for each token in the input sequence. Each bounding box is a list of four values
            (x1, y1, x2, y2), where (x1, y1) is the top left corner, and (x2, y2) is the bottom right corner of the
            bounding box.
        bbox_first_token_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Mask to indicate the first token of each bounding box. Mask values selected in `[0, 1]`:

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

        Examples:

        ```python
        >>> import torch
        >>> from transformers import BrosProcessor, BrosSpadeELForTokenClassification

        >>> processor = BrosProcessor.from_pretrained("jinho8345/bros-base-uncased")

        >>> model = BrosSpadeELForTokenClassification.from_pretrained("jinho8345/bros-base-uncased")

        >>> encoding = processor("Hello, my dog is cute", add_special_tokens=False, return_tensors="pt")
        >>> bbox = torch.tensor([[[0, 0, 1, 1]]]).repeat(1, encoding["input_ids"].shape[-1], 1)
        >>> encoding["bbox"] = bbox

        >>> outputs = model(**encoding)
        ```rY  r   r   NrZ  rA   ry   r  rJ  r+   )r  rm   r   rm  r[  r   r   r{   r'   r`  ra  r\  rE   rF   r   r]  r^  rz   r_  r   r!   r"   )r:   r   r\   r   rG  rx   rv   r   rH  r   rL  rb  rK  r   rM  r   rd  r{   rf  masks                       r-   rO   z)BrosSpadeELForTokenClassification.forwardq  s   R AJ		 A
))%'A
 A
 %QZ/99!Q?JJL##$68JKSSTUV')H)7)=)=&J#**F#ii8JKNNV\didndnNoO(--b1D$)II**KKQuzz&Q %! ''(=aqj(I5;;W]WcWcKdKhKhiF''a
(CU[[QWQ]Q]E^EbEbcFFKKNQ,>?Ev{{SUW[G\]D$!//))	
 	
r,   r>  rN  rS   s   @r-   rk  rk  Z  s    +4&  *.$(.259.2,0-1&*Q
<<$&Q
 llT!Q
 t+	Q

  %||d2Q
 t+Q
 llT)Q
 ||d*Q
 t#Q
 +,Q
 
u||	4	4Q
  Q
r,   rk  )r  r)  r@  rQ  rk  )@r&   r   dataclassesr   r'   r   torch.nnr    r   r  activationsr   masking_utilsr	   modeling_layersr
   modeling_outputsr   r   r   modeling_utilsr   processing_utilsr   pytorch_utilsr   utilsr   r   r   r   r   utils.genericr   utils.output_capturingr   r   configuration_brosr   
get_loggerr#   loggerr   r  r/   rU   rd   rq   r   r   r   r   r   r   r   r   r  r  r)  r@  rQ  rk  __all__r+   r,   r-   <module>r     s     !   % & ! 6 9 
 . & 6 _ _ 7 E * 
		H	% 
 7k 7 7"		 *		 & ; ;|J.		 J.\RYY BII 6ryy  8* 8x BII D 2/ 2 24
% 
@ t
# t
 t
n X
!4 X
 X
v F
(; F
F
R d
(; d
d
Nr,   