ó
    >:j#  ã                  óÎ   • S SK Jr  S SKJrJr  S SKrS SKrS SKJ	r	  S SKJ
r
Jr  SS jrSS jrSS jrSS	 jr\SS
 j5       r\SS j5       rSS jrSS jrSS jrSS jrSS jrg)é    )Úannotations)ÚAnyÚoverloadN)Ú
coo_matrix)ÚTensorÚdevicec           	     óæ  • [        U [        5      (       a}  [        S U  5       5      (       aO  [        R                  " U  Vs/ s H-  oR                  5       R                  [        R                  S9PM/     sn5      $ [        R                  " U 5      n O+[        U [        5      (       d  [        R                  " U 5      n U R                  (       a  U R                  [        R                  S9$ U $ s  snf )zô
Converts the input `a` to a PyTorch tensor if it is not already a tensor.
Handles lists of sparse tensors by stacking them.

Args:
    a (Union[list, np.ndarray, Tensor]): The input array or tensor.

Returns:
    Tensor: The converted tensor.
c              3  óh   #   • U  H(  n[        U[        5      =(       a    UR                  v •  M*     g 7f©N)Ú
isinstancer   Ú	is_sparse)Ú.0Úxs     Ú^/home/mande/repo/quber/.venv/lib/python3.13/site-packages/sentence_transformers/util/tensor.pyÚ	<genexpr>Ú%_convert_to_tensor.<locals>.<genexpr>   s#   é € Ð@ºa¸Œz˜!œVÓ$×4¨¯©Ô4ºaùs   ‚02©Údtype)r   ÚlistÚallÚtorchÚstackÚcoalesceÚtoÚfloat32Útensorr   r   )Úar   s     r   Ú_convert_to_tensorr      s¡   € ô �!”T×ÑäÑ@¹aÓ@×@Ñ@ä—;’;ÉaÓPÊaÈ§
¡
£§¡´e·m±m Ó DÉaÑPÓQÐQä—’˜Q“‰AÜ˜œ6×"Ñ"Ü�LŠL˜‹OˆØ‡{‡{Ø�t‰tœ%Ÿ-™-ˆtÐ(Ð(Ø€Hùò  Qs   Á4C.c                óP   • U R                  5       S:X  a  U R                  S5      n U $ )z²
If the tensor `a` is 1-dimensional, it is unsqueezed to add a batch dimension.

Args:
    a (Tensor): The input tensor.

Returns:
    Tensor: The tensor with a batch dimension.
é   r   )ÚdimÚ	unsqueeze©r   s    r   Ú_convert_to_batchr$   $   s#   € ð 	‡u�uƒw�!ƒ|Ø�K‰K˜‹NˆØ€Hó    c                óf   • [        U 5      n U R                  5       S:X  a  U R                  S5      n U $ )a  
Converts the input data to a tensor with a batch dimension.
Handles lists of sparse tensors by stacking them.

Args:
    a (Union[list, np.ndarray, Tensor]): The input data to be converted.

Returns:
    Tensor: The converted tensor with a batch dimension.
r    r   )r   r!   r"   r#   s    r   Ú_convert_to_batch_tensorr'   3   s-   € ô 	˜1Ó€AØ‡u�uƒw�!ƒ|Ø�K‰K˜‹NˆØ€Hr%   c                óD  • U R                   (       d)  [        R                  R                  R	                  U SSS9$ U R                  5       n U R                  5       U R                  5       p![        R                  " U R                  S5      U R                  S9nUR                  SUS   US-  5        [        R                  " U5      R                  SUS   5      nUS:„  nUR                  5       nXT==   X4   -  ss'   [        R                  " XU R                  5       5      $ )zÉ
Normalizes the embeddings matrix, so that each sentence embedding has unit length.

Args:
    embeddings (Tensor): The input embeddings matrix.

Returns:
    Tensor: The normalized embeddings matrix.
é   r    )Úpr!   r   ©r   )r   r   ÚnnÚ
functionalÚ	normalizer   ÚindicesÚvaluesÚzerosÚsizer   Ú
index_add_ÚsqrtÚindex_selectÚcloneÚsparse_coo_tensor)Ú
embeddingsr/   r0   Ú	row_normsÚmaskÚnormalized_valuess         r   Únormalize_embeddingsr<   D   sñ   € ð ××Ü�x‰x×"Ñ"×,Ñ,¨Z¸1À!Ð,ÐDÐDà×$Ñ$Ó&€JØ ×(Ñ(Ó*¨J×,=Ñ,=Ó,?ˆVô —’˜JŸO™O¨AÓ.°z×7HÑ7HÑI€IØ×Ñ˜˜G A™J¨°©	Ô2Ü—
’
˜9Ó%×2Ñ2°1°g¸a±jÓA€Ið �q‰=€DØŸ™›ÐØÓ˜y™Ñ.Óä×"Ò" 7¸z¿¹Ó?PÓQÐQr%   c                ó   • g r   © ©r8   Útruncate_dims     r   Útruncate_embeddingsrA   a   s   € ØY\r%   c                ó   • g r   r>   r?   s     r   rA   rA   e   s   € Ø]`r%   c                ó   • U SSU24   $ )a\  
Truncates the embeddings matrix.

Args:
    embeddings (Union[np.ndarray, torch.Tensor]): Embeddings to truncate.
    truncate_dim (Optional[int]): The dimension to truncate sentence embeddings to. `None` does no truncation.

Example:
    >>> from sentence_transformers import SentenceTransformer
    >>> from sentence_transformers.util import truncate_embeddings
    >>> model = SentenceTransformer("tomaarsen/mpnet-base-nli-matryoshka")
    >>> embeddings = model.encode(["It's so nice outside!", "Today is a beautiful day.", "He drove to work earlier"])
    >>> embeddings.shape
    (3, 768)
    >>> model.similarity(embeddings, embeddings)
    tensor([[1.0000, 0.8100, 0.1426],
            [0.8100, 1.0000, 0.2121],
            [0.1426, 0.2121, 1.0000]])
    >>> truncated_embeddings = truncate_embeddings(embeddings, 128)
    >>> truncated_embeddings.shape
    >>> model.similarity(truncated_embeddings, truncated_embeddings)
    tensor([[1.0000, 0.8092, 0.1987],
            [0.8092, 1.0000, 0.2716],
            [0.1987, 0.2716, 1.0000]])

Returns:
    Union[np.ndarray, torch.Tensor]: Truncated embeddings.
.Nr>   r?   s     r   rA   rA   i   s   € ð: �c˜=˜L˜=Ð(Ñ)Ð)r%   c                ó$  • Uc  U $ [        U [        R                  5      (       a  [        R                  " U 5      n U R
                  u  p#U R                  n[        R                  " [        R                  " U 5      [        X5      SS9u  pV[        R                  " U [        R                  S9n[        R                  " X$S9R                  S5      R                  S[        X5      5      nSXxR                  5       UR                  5       4'   SX) '   U $ )a{  
Keeps only the top-k values (in absolute terms) for each embedding and creates a sparse tensor.

Args:
    embeddings (Union[np.ndarray, torch.Tensor]): Embeddings to sparsify by keeping only top_k values.
    max_active_dims (int): Number of values to keep as non-zeros per embedding.

Returns:
    torch.Tensor: A sparse tensor containing only the top-k values per embedding.
r    )Úkr!   r   r+   éÿÿÿÿTr   )r   ÚnpÚndarrayr   r   Úshaper   ÚtopkÚabsÚminÚ
zeros_likeÚboolÚaranger"   ÚexpandÚflatten)	r8   Úmax_active_dimsÚ
batch_sizer!   r   Ú_Útop_indicesr:   Úbatch_indicess	            r   Úselect_max_active_dimsrW   ‰   sã   € ð ÑØÐä�*œbŸj™j×)Ñ)Ü—\’\ *Ó-ˆ
à ×&Ñ&�O€JØ×Ñ€Fô —Z’Z¤§	¢	¨*Ó 5¼¸_Ó9RÐXYÑZ�N€Aô ×Ò˜J¬e¯j©jÑ9€DÜ—L’L Ñ;×EÑEÀaÓH×OÑOÐPRÔTWÐXgÓTmÓn€MØ;?€D×	Ñ	Ó	  +×"5Ñ"5Ó"7Ð	7Ñ8ð €JˆuÑàÐr%   c                ót   • U  H1  n[        X   [        5      (       d  M  X   R                  U5      X'   M3     U $ )aY  
Send a PyTorch batch (i.e., a dictionary of string keys to Tensors) to a device (e.g. "cpu", "cuda", "mps").

Args:
    batch (Dict[str, Tensor]): The batch to send to the device.
    target_device (torch.device): The target device (e.g. "cpu", "cuda", "mps").

Returns:
    Dict[str, Tensor]: The batch with tensors sent to the target device.
)r   r   r   )ÚbatchÚtarget_deviceÚkeys      r   Úbatch_to_devicer\   «   s6   € ó ˆÜ�e‘j¤&×)Ó)Ø™Ÿ™ }Ó5ˆE‹Jñ ð €Lr%   c                ó  • U R                  5       n U R                  5       R                  5       R                  5       nU R	                  5       R                  5       R                  5       n[        X!S   US   44U R                  S9$ )Nr   r    )rI   )r   r/   ÚcpuÚnumpyr0   r   rI   )r   r/   r0   s      r   Úto_scipy_coor`   ¼   sd   € Ø	�
‰
‹€AØ�i‰i‹k�o‰oÓ×%Ñ%Ó'€GØ�X‰X‹Z�^‰^Ó×#Ñ#Ó%€FÜ�v¨¡
¨G°A©JÐ7Ð8ÀÇÁÑHÐHr%   c                ól  • U R                   (       d  U R                  5       n U R                  5       n [        R                  " U R                  S5      U R                  [        R                  S9nU R                  5       S:X  a"  SXR                  5       R                  5       '   U$ U R                  5       S:X  a`  U R                  5       R                  5       S:”  a<  U R                  5       n[        R                  " US   SS9u  p4UR                  5       X'   U$ [        SU R                  5        S	35      e)
zû
Compute count vector from sparse embeddings indicating how many samples have non-zero values in each dimension.

Args:
    embeddings: Sparse tensor of shape (batch_size, vocab_size) or (vocab_size,)

Returns:
    Count vector of shape (vocab_size,)
rF   )r   r   r    r)   r   T)Úreturn_countszExpected 1D or 2D tensor, got ÚD)r   Ú	to_sparser   r   r1   r2   r   Úint32r!   r/   Úsqueezer0   ÚnumelÚuniqueÚintÚ
ValueError)r8   Úcount_vectorr/   Úunique_dimsÚcountss        r   Úcompute_count_vectorrn   Ã   s  € ð ××Ø×)Ñ)Ó+ˆ
ð ×$Ñ$Ó&€Jä—;’;˜zŸ™¨rÓ2¸:×;LÑ;LÔTY×T_ÑT_Ñ`€LØ‡~�~Ó˜1Óà78ˆ×'Ñ'Ó)×1Ñ1Ó3Ñ4ØÐØ	�‰Ó	˜QÓ	à×ÑÓ×$Ñ$Ó&¨Ó*Ø ×(Ñ(Ó*ˆGä"'§,¢,¨w°q©zÈÑ"NÑˆKØ(.¯
©
«ˆLÑ%àÐäÐ9¸*¿.¹.Ó:JÐ9KÈ1ÐMÓNÐNr%   )r   zlist | np.ndarray | TensorÚreturnr   )r   r   ro   r   )r8   r   ro   r   )r8   ú
np.ndarrayr@   ú
int | Nonero   rp   )r8   útorch.Tensorr@   rq   ro   rr   )r8   únp.ndarray | torch.Tensorr@   rq   ro   rs   )r8   rs   rR   rq   ro   rr   )rY   údict[str, Any]rZ   r   ro   rt   )r   r   ro   r   )r8   rr   ro   rr   )Ú
__future__r   Útypingr   r   r_   rG   r   Úscipy.sparser   r   r   r   r$   r'   r<   rA   rW   r\   r`   rn   r>   r%   r   Ú<module>rx      si   ðÝ "ç  ã Û Ý #ß  ôô2ôô"Rð: 
Û \ó 
Ø \ð 
Û `ó 
Ø `ô*ô@ôDô"IõOr%   