ó
    �®žjŸ*  ã                   óz  • S SK rS SKrS SKJr  S SKJs  Jr  S SKJ	r	  S SK
JrJr  S SKJr  S rS rS!S jrS rSS	 jrS
 rS rS r " S S\5      r\R2                  " 5       S"S j5       rS rS rS rS r " S S\R>                  5      r S#S\!S\"S\RF                  4S jjr$ " S S5      r%S\&4S jr'S\(4S  jr)g)$é    N)ÚImage)Ú
BasicBlockÚconv1x1)Úbox_areac                 ój   • Sn[        U 5       H  nUS:w  a    O	US-  nM     US:X  a  U S4$ U SU*  nX14$ )zÓ
Remove the trailing zeros from the provided input

Parameters
----------
list: List of integers
    Predicted sequence

Returns
-------
list: List of integers
    The part of the input before the zero padding

r   é   N)Úreversed)ÚseqÚpad_lenÚxÚ	un_paddeds       Úg/home/mande/repo/quber/.venv/lib/python3.13/site-packages/docling_ibm_models/tableformer/utils/utils.pyÚremove_paddingr   
   sR   € ð €GÜ�cŽ]ˆØ�‹6ÙØ�1‰Šñ ð �!ƒ|Ø�Aˆvˆà�I�g�X�€IØÐÐó    c                 ó0   • [         R                  " U SS9nU$ )a   
Convert probabilities to predictions

Parameters
----------
probabilities : Tensor[batch_size, vocab_size, seq_len]
    All log probabilities coming out at the last stage of the decoder

Returns
-------
predictions : tensor [batch_size, output_sequence_length]
    The prediceted trags

r   ©Údim)ÚtorchÚargmax)ÚprobabilitiesÚmax_idxs     r   Úprobabilities_to_predictionsr   %   s   € ô  �lŠl˜=¨aÑ0€GØ€Nr   c                 ó‚  • X   R                  5       nX   R                  5       nSnSnUb  X#   n[        [        U5      [        U5      5      n[        SR	                  UR                  US5      UR                  5       5      5        [        SR	                  UR                  US5      UR                  5       5      5        g)a‹  
For the Tags, print the target and predicted tensors for the specified batch index

We expect to have the batch size as the first dimension.
Only the specified batch is extractred and the remaining dimenions are flattened.
The results are printed as 2 lists with the target on top and the predictions below underlined

Parameters
---------
target : tensor [batch_size, output_sequence_length]
    The ground truth tags

predictions : tensor [batch_size, output_sequence_length]
    The prediceted trags

filenames : list of string
    The actual filename that provides the data

batch_idx : int
    Which index in the batch dimension will be printed
ÚtargetÚpredictNú{}: {}Ú )ÚflattenÚmaxÚlenÚprintÚformatÚljustÚtolist)	r   ÚpredictionsÚ	filenamesÚ	batch_idxÚtarget_flatÚpredictions_flatÚtarget_labelÚpredict_labelÚ	label_lens	            r   Úprint_target_predictr-   9   s«   € ð, Ñ#×+Ñ+Ó-€KØ"Ñ-×5Ñ5Ó7ÐØ€LØ€MØÑØ Ñ+ˆÜ”C˜Ó%¤s¨=Ó'9Ó:€IÜ	ˆ(�/‰/˜,×,Ñ,¨Y¸Ó<¸k×>PÑ>PÓ>RÓ
SÔTÜ	Ø�‰˜×+Ñ+¨I°sÓ;Ð=M×=TÑ=TÓ=VÓWõr   c                 óº   • [         R                  " U 5       n[        R                  " U5      nUR	                  SSS5      nSSS5        U$ ! , (       d  f       W$ = f)zâ
Load an image from the disk as a numpy array

Parameters
----------
full_fn : string
    The full path filename of the image

Results
-------
img : numpy array: (channels, width, height)
    The loaded image as a numpy array
é   r   r   N)r   ÚopenÚnpÚasarrayÚ	transpose)Úfull_fnÚfÚimgs      r   Ú
load_imager7   \   sM   € ô 
�Š�GÔ	 Ü�jŠj˜‹mˆØ�m‰m˜A˜q !Ó$ˆ÷ 
ð €J÷ 
Ô	ð €Jús   —*AÁ
Ac                 ó  • / n[         R                  " [        SSU 5      [         R                  " S5      5      nUR	                  [        SSX5      5        UR	                  [        SSS5      5        [         R                  " U6 $ )Né   i   r   )ÚnnÚ
Sequentialr   ÚBatchNorm2dÚappendr   )ÚstrideÚlayersÚ
downsamples      r   Úresnet_blockrA   p   sh   € Ø€FÜ—’Ü��S˜&Ó!Ü
�Š�sÓó€Jð ‡M�M”*˜S # vÓ:Ô;Ø
‡M�M”*˜S # qÓ)Ô*Ü�=Š=˜&Ð!Ð!r   c                 ó„   • [        U [        R                  5      (       a  U R                  5       $ [	        S U  5       5      $ )zH
Wraps hidden states in new Tensors, to detach them from their history.
c              3   ó8   #   • U  H  n[        U5      v •  M     g 7f©N)Úrepackage_hidden)Ú.0Úvs     r   Ú	<genexpr>Ú#repackage_hidden.<locals>.<genexpr>‚   s   é € Ð4²!¨QÔ% a×(Ð(²!ùs   ‚)Ú
isinstancer   ÚTensorÚdetachÚtuple)Úhs    r   rE   rE   {   s2   € ô �!”U—\‘\×"Ñ"Ø�x‰x‹zÐäÑ4±!Ó4Ó4Ð4r   c                 ó6  • UR                  S5      nU R                  USSS5      u  pEUR                  UR                  SS5      R	                  U5      5      nUR                  S5      R                  5       R                  5       nUR                  5       SU-  -  $ )z²
Computes top-k accuracy, from predicted and true labels.

:param scores: scores from the model
:param targets: true labels
:param k: k in top-k accuracy
:return: top-k accuracy
r   r   Téÿÿÿÿç      Y@)ÚsizeÚtopkÚeqÚviewÚ	expand_asÚfloatÚsumÚitem)ÚscoresÚtargetsÚkÚ
batch_sizeÚ_ÚindÚcorrectÚcorrect_totals           r   Úaccuracyrb   …   s„   € ð —‘˜a“€JØ�[‰[˜˜A˜t TÓ*�F€AØ�f‰f�W—\‘\ " aÓ(×2Ñ2°3Ó7Ó8€GØ—L‘L Ó$×*Ñ*Ó,×0Ñ0Ó2€MØ×ÑÓ 5¨:Ñ#5Ñ6Ð6r   c                 ó®   • U R                    HE  nUS    H9  nUR                  c  M  UR                  R                  R                  U* U5        M;     MG     g)z­
Clips gradients computed during backpropagation to avoid explosion of gradients.

:param optimizer: optimizer with the gradients to be clipped
:param grad_clip: clip value
ÚparamsN)Úparam_groupsÚgradÚdataÚclamp_)Ú	optimizerÚ	grad_clipÚgroupÚparams       r   Úclip_gradientrm   –   sF   € ð ×'Ô'ˆØ˜8”_ˆEØ�z‰zÓ%Ø—
‘
—‘×&Ñ&¨	 z°9Ö=ó %ò (r   c                   ó.   • \ rS rSrSrS rS rSS jrSrg)	ÚAverageMeteré£   zB
Keeps track of most recent, average, sum, and count of a metric.
c                 ó$   • U R                  5         g rD   )Úreset©Úselfs    r   Ú__init__ÚAverageMeter.__init__¨   s   € Ø�
‰
�r   c                 ó<   • SU l         SU l        SU l        SU l        g ©Nr   )ÚvalÚavgrX   Úcountrs   s    r   rr   ÚAverageMeter.reset«   s   € ØˆŒØˆŒØˆŒØˆ�
r   c                 ó¤   • Xl         U =R                  X-  -  sl        U =R                  U-  sl        U R                  U R                  -  U l        g rD   )ry   rX   r{   rz   )rt   ry   Úns      r   ÚupdateÚAverageMeter.update±   s8   € ØŒØ�Š�C‘GÑ�Ø�
Š
�a‰�
Ø—8‘8˜dŸj™jÑ(ˆ�r   )rz   r{   rX   ry   N©r   )	Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__ru   rr   r   Ú__static_attributes__© r   r   ro   ro   £   s   † ñòò÷)r   ro   r�   c                 ó  • UR                  5       S:X  a   [        R                  " / U R                  S9/$ [	        U5      nUR                  S5      nU R                  USSS5      u  pVUR                  5       nUR                  UR                  SS5      R                  U5      5      n/ nU HW  n	USU	 R                  S5      R                  5       R                  S5      n
UR                  U
R                  SU-  5      5        MY     U$ )z6Computes the precision@k for the specified values of kr   ©Údevicer   TrP   NrQ   )Únumelr   Úzerosr‹   r   rR   rS   ÚtrT   rU   rV   rW   rX   r=   Úmul_)Úoutputr   rS   Úmaxkr]   r^   Úpredr`   Úresr\   Ú	correct_ks              r   Úbip_accuracyr•   ¸   sß   € ð ‡|�|ƒ~˜ÓÜ—’˜B v§}¡}Ñ5Ð6Ð6Üˆt‹9€DØ—‘˜Q“€Jà�k‰k˜$  4¨Ó.�G€AØ�6‰6‹8€DØ�g‰g�f—k‘k ! RÓ(×2Ñ2°4Ó8Ó9€Gà
€CÛˆØ˜B˜Q�K×$Ñ$ RÓ(×.Ñ.Ó0×4Ñ4°QÓ7ˆ	Ø�
‰
�9—>‘> %¨*Ñ"4Ó5Ö6ñ ð €Jr   c                 ó�   • U R                  S5      u  pp4USU-  -
  USU-  -
  USU-  -   USU-  -   /n[        R                  " USS9$ )NrP   g      à?r   ©Úunbindr   Ústack)r   Úx_cÚy_cÚwrN   Úbs         r   Úbox_cxcywh_to_xyxyrž   Ë   sR   € Ø—X‘X˜b“\�N€CˆaØ
��a‘‰-˜3  q¡™=¨C°#¸±'©M¸SÀ3ÈÁ7¹]ÐL€AÜ�;Š;�q˜bÑ!Ð!r   c                 ó|   • U R                  S5      u  pp4X-   S-  X$-   S-  X1-
  XB-
  /n[        R                  " USS9$ )NrP   r/   r   r—   )r   Úx0Úy0Úx1Úy1r�   s         r   Úbox_xyxy_to_cxcywhr¤   Ñ   sB   € Ø—X‘X˜b“\�N€BˆBØ
‰'�Q‰˜™ A™¨©°2±7Ð<€AÜ�;Š;�q˜bÑ!Ð!r   c                 óV  • [        U 5      n[        U5      n[        R                  " U S S 2S S S24   US S 2S S24   5      n[        R                  " U S S 2S SS 24   US S 2SS 24   5      nXT-
  R	                  SS9nUS S 2S S 2S4   US S 2S S 2S4   -  nUS S 2S 4   U-   U-
  nXx-  n	X˜4$ )Nr/   r   ©Úminr   )r   r   r   r§   Úclamp)
Úboxes1Úboxes2Úarea1Úarea2ÚltÚrbÚwhÚinterÚunionÚious
             r   Úbox_iour³   Ø   s¾   € Ü�VÓ€EÜ�VÓ€Eä	�Š�6š!˜T 2 A 2˜+Ñ&¨ªq°"°1°"¨u©Ó	6€BÜ	�Š�6š!˜T 1¡2˜+Ñ&¨ªq°!±"¨u©Ó	6€Bà
‰'�‰˜QˆÐ	€BØŠq’!�Qˆw‰K˜"šQ¢ 1˜W™+Ñ%€Eà’!�T�'‰N˜UÑ" UÑ*€Eà
‰-€CØˆ:Ðr   c                 óÜ  • U SS2SS24   U SS2SS24   :¬  R                  5       (       d   eUSS2SS24   USS2SS24   :¬  R                  5       (       d   e[        X5      u  p#[        R                  " U SS2SSS24   USS2SS24   5      n[        R                  " U SS2SSS24   USS2SS24   5      nXT-
  R                  SS9nUSS2SS2S4   USS2SS2S4   -  nX'U-
  U-  -
  $ )z®
Generalized IoU from https://giou.stanford.edu/

The boxes should be in [x0, y0, x1, y1] format

Returns a [N, M] pairwise matrix, where N = len(boxes1)
and M = len(boxes2)
Nr/   r   r¦   r   )Úallr³   r   r§   r   r¨   )r©   rª   r²   r±   r­   r®   r¯   Úareas           r   Úgeneralized_box_iour·   è   s  € ð ’1�a‘b�5‰M˜V¢A r¨ r E™]Ñ*×/Ñ/×1Ñ1Ð1Ð1Ø’1�a‘b�5‰M˜V¢A r¨ r E™]Ñ*×/Ñ/×1Ñ1Ð1Ð1Ü˜Ó(�J€Cä	�Š�6š!˜T 2 A 2˜+Ñ&¨ªq°"°1°"¨u©Ó	6€BÜ	�Š�6š!˜T 1¡2˜+Ñ&¨ªq°!±"¨u©Ó	6€Bà
‰'�‰˜QˆÐ	€BØŠa’�Aˆg‰;˜šAšq !˜G™Ñ$€Dà˜‘, $Ñ&Ñ&Ð&r   c                   ó2   ^ • \ rS rSrSrU 4S jrS rSrU =r$ )ÚMLPr9   z4Very simple multi-layer perceptron (also called FFN)c                 ó¦   >• [         TU ]  5         X@l        U/US-
  -  n[        R                  " S [        U/U-   XS/-   5       5       5      U l        g )Nr   c              3   óR   #   • U  H  u  p[         R                  " X5      v •  M     g 7frD   )r:   ÚLinear)rF   r~   r\   s      r   rH   ÚMLP.__init__.<locals>.<genexpr>  s    é € ð $
Ú(N¡ ŒB�IŠI�a�OˆOÒ(Nùs   ‚%')Úsuperru   Ú
num_layersr:   Ú
ModuleListÚzipr?   )rt   Ú	input_dimÚ
hidden_dimÚ
output_dimr¿   rN   Ú	__class__s         €r   ru   ÚMLP.__init__  sR   ø€ Ü‰ÑÔØ$ŒØˆL˜J¨™NÑ+ˆÜ—m’mñ $
Ü(+¨Y¨K¸!©O¸QÀÑ=MÔ(Nó$
ó 
ˆ�r   c                 ó®   • [        U R                  5       H;  u  p#X R                  S-
  :  a  [        R                  " U" U5      5      OU" U5      nM=     U$ )Nr   )Ú	enumerater?   r¿   ÚFÚrelu)rt   r   ÚiÚlayers       r   ÚforwardÚMLP.forward  sB   € Ü! $§+¡+Ö.‰HˆAØ$%¯©¸!Ñ(;Ó$;”—’‘u˜Q“xÔ ÁÀqÃŠAñ /àˆr   )r?   r¿   )	r‚   rƒ   r„   r…   r†   ru   rÍ   r‡   Ú__classcell__)rÅ   s   @r   r¹   r¹      s   ø† Ù>õ
÷ð r   r¹   Úszr‹   Úreturnc                 ó*  • [         R                  " [         R                  " X 5      5      S:H  R                  SS5      nUR	                  5       R                  US:H  [	        S5      5      R                  US:H  [	        S5      5      R                  US9nU$ )z/Generate the attention mask for causal decodingr   r   z-infg        rŠ   )r   ÚtriuÚonesr3   rW   Úmasked_fillÚto)rÐ   r‹   Úmasks      r   Úgenerate_square_subsequent_maskrØ     st   € ä�JŠJ”u—z’z "Ó)Ó*¨aÑ/×:Ñ:¸1¸aÓ@€Dà�
‰
‹ß	‰�T˜Q‘Y¤ f£Ó	.ß	‰�T˜Q‘Y¤ c£
Ó	+ß�b�€bÐð	 	ð
 €Kr   c                   ó0   • \ rS rSrSrSSS\4S jrS rSrg	)
ÚEarlyStoppingi  z“Early stops the training if validation loss doesn't improve after a given patience.
Source from: https://github.com/Bjarten/early-stopping-pytorch
r/   Fr   c                 óˆ   • Xl         X l        SU l        SU l        SU l        [
        R                  U l        X0l        X@l	        g)a  
Args:
    patience (int): How long to wait after last time validation loss improved.
                    Default: 7
    verbose (bool): If True, prints a message for each validation loss improvement.
                    Default: False
    delta (float): Minimum change in the monitored quantity to qualify as an improvement.
                    Default: 0
    path (str): Path for the checkpoint to be saved to.
                    Default: 'checkpoint.pt'
    trace_func (function): trace print function.
                    Default: print
r   NF)
Ú	_patienceÚ_verboseÚ_counterÚ_best_scoreÚ_early_stopr1   ÚInfÚ_val_loss_minÚ_deltaÚ_trace_func)rt   ÚpatienceÚverboseÚdeltaÚ
trace_funcs        r   ru   ÚEarlyStopping.__init__!  s<   € ð "ŒØŒØˆŒØˆÔØ ˆÔÜŸV™VˆÔØŒØ%Õr   c                 óR  • U* nSnU R                   cG  X l         SnU R                  (       a&  SU R                  S SUS S3nU R                  U5        Xl        U$ X R                   U R                  -   :  ae  U =R
                  S-  sl        U R                  SU R
                   SU R                   35        U R
                  U R                  :¼  a	  SU l        S	nU$ X l         SnS
U l        U R                  (       a&  SU R                  S SUS S3nU R                  U5        Xl        U$ )NTzValidation loss decreased (z.6fz --> z).r   zEarlyStopping counter: z out of Fr   )rß   rÝ   râ   rä   rã   rÞ   rÜ   rà   )rt   Úval_lossÚscoreÚsave_checkpointÚverbs        r   Ú__call__ÚEarlyStopping.__call__8  s6  € Ø�	ˆØˆØ×ÑÑ#Ø$ÔØ"ˆOØ�}�}Ø4°T×5GÑ5GÈÐ4LÈEÐRZÐ[^ÐQ_Ð_aÐb�Ø× Ñ  Ô&Ø!)Ôð" Ðð! ×%Ñ%¨¯©Ñ3Ó3Ø�MŠM˜QÑ�MØ×ÑØ)¨$¯-©-¨¸ÀÇÁÐ@PÐQôð �}‰} §¡Ó.Ø#'�Ô Ø"'�ð Ðð  %ÔØ"ˆOØˆDŒMØ�}�}Ø4°T×5GÑ5GÈÐ4LÈEÐRZÐ[^ÐQ_Ð_aÐb�Ø× Ñ  Ô&Ø!)ÔØÐr   )rß   rÞ   rã   rà   rÜ   rä   râ   rÝ   N)	r‚   rƒ   r„   r…   r†   r!   ru   rï   r‡   rˆ   r   r   rÚ   rÚ     s   † ñð !"¨5¸Àeô &õ.r   rÚ   Úmc                 óð  • [        U 5      S:X  a  g[        [        U 5      5      n[        U[        5      =(       a    UR                  5       nU(       a4  [        U R                  5        Vs/ s H  n[        U5      PM     sn5      nO)[        U R                  5        Vs/ s H  o3PM     sn5      nU H7  nU(       a  U [	        U5         nOX   n[        SR                  X55      5        M9     gs  snf s  snf )z6
Print dict elements in separate lines sorted by keys
r   Nr   )r    ÚnextÚiterrJ   ÚstrÚ	isnumericÚsortedÚkeysÚintr!   r"   )rñ   Ú	first_keyÚ
is_numericr\   rø   rG   s         r   Ú
print_dictrü   U  sº   € ô ˆ1ƒv�ƒ{Øô ”T˜!“W“€IÜ˜I¤sÓ+×E°	×0CÑ0CÓ0E€JÞÜ q§v¡v¤xÓ0¢x !”s˜1–v¡xÑ0Ó1‰ä !§&¡&¤(Ó+¢(˜Q’q¡(Ñ+Ó,ˆãˆÞØ”#�a“&‘	‰Aà‘ˆAÜˆh�o‰o˜aÓ#Ö$ò ùò	 1ùâ+s   Á*C.ÂC3Úlstc           	      óØ   • [        U 5       H[  u  p[        U[        5      (       a'  [        SR	                  U[        U5      U5      5        MA  [        SR	                  X5      5        M]     g)z'
Print list elements in separate lines
z{}: ({}) - {}r   N)rÈ   rJ   Úlistr!   r"   r    )rý   rË   Úelms      r   Ú
print_listr  l  sM   € ô ˜C–.‰ˆÜ�cœ4× Ñ Ü�/×(Ñ(¨¬C°«H°cÓ:Ö;ä�(—/‘/ !Ó)Ö*ò	 !r   rx   )r�   )Úcpu)*Únumpyr1   r   Útorch.nnr:   Útorch.nn.functionalÚ
functionalrÉ   ÚPILr   Útorchvision.models.resnetr   r   Útorchvision.ops.boxesr   r   r   r-   r7   rA   rE   rb   rm   Úobjectro   Úno_gradr•   rž   r¤   r³   r·   ÚModuler¹   rù   rõ   rK   rØ   rÚ   Údictrü   rÿ   r  rˆ   r   r   Ú<module>r     sÑ   ðÛ Û Ý ß Ð Ý ß 9Ý *òò6ô( òFô("ò5ò7ò"
>ô)�6ô )ð* ‡‚ƒóó ðò$"ò"òò 'ô0ˆ"�)‰)ô ñ"¨ð °Sð ÀUÇ\Á\õ ÷6ñ 6ðr%�$ô %ð.+�Dõ +r   