Ë
    l^(hª   ã                  ó€   — d dl mZ d dlmZ d dlZd dlmc mZ d dlm	Z	mZ d dl
mZmZ  G d„ dej                  «      Zy)é    )Úannotations)ÚIterableN)ÚTensorÚnn)ÚSentenceTransformerÚutilc                  ób   ‡ — e Zd Z	 	 	 	 d	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Zdd„Zdd„Zedd„«       Zˆ xZS )	ÚMegaBatchMarginLossc                ó¤   •— t         ‰| �  «        || _        || _        || _        || _        |r| j                  | _        y| j                  | _        y)aÞ  
        Given a large batch (like 500 or more examples) of (anchor_i, positive_i) pairs, find for each pair in the batch
        the hardest negative, i.e. find j != i such that cos_sim(anchor_i, positive_j) is maximal. Then create from this a
        triplet (anchor_i, positive_i, positive_j) where positive_j serves as the negative for this triplet.

        Then train as with the triplet loss.

        Args:
            model: SentenceTransformerModel
            positive_margin: Positive margin, cos(anchor, positive)
                should be > positive_margin
            negative_margin: Negative margin, cos(anchor, negative)
                should be < negative_margin
            use_mini_batched_version: As large batch sizes require a lot
                of memory, we can use a mini-batched version. We break
                down the large batch into smaller batches with fewer
                examples.
            mini_batch_size: Size for the mini-batches. Should be a
                divisor for the batch size in your data loader.

        References:
            - This loss function was inspired by the ParaNMT paper: https://www.aclweb.org/anthology/P18-1042/

        Requirements:
            1. (anchor, positive) pairs
            2. Large batches (500 or more examples)

        Inputs:
            +---------------------------------------+--------+
            | Texts                                 | Labels |
            +=======================================+========+
            | (anchor, positive) pairs              | none   |
            +---------------------------------------+--------+

        Recommendations:
            - Use ``BatchSamplers.NO_DUPLICATES`` (:class:`docs <sentence_transformers.training_args.BatchSamplers>`) to
              ensure that no in-batch negatives are duplicates of the anchor or positive samples.

        Example:
            ::

                from sentence_transformers import SentenceTransformer, SentenceTransformerTrainingArguments, SentenceTransformerTrainer, losses
                from datasets import Dataset

                train_batch_size = 250
                train_mini_batch_size = 32

                model = SentenceTransformer('all-MiniLM-L6-v2')
                train_dataset = Dataset.from_dict({
                    "anchor": [f"This is sentence number {i}" for i in range(500)],
                    "positive": [f"This is sentence number {i}" for i in range(1, 501)],
                })
                loss = losses.MegaBatchMarginLoss(model=model, mini_batch_size=train_mini_batch_size)

                args = SentenceTransformerTrainingArguments(
                    output_dir="output",
                    per_device_train_batch_size=train_batch_size,
                )
                trainer = SentenceTransformerTrainer(
                    model=model,
                    args=args,
                    train_dataset=train_dataset,
                    loss=loss,
                )
                trainer.train()
        N)	ÚsuperÚ__init__ÚmodelÚpositive_marginÚnegative_marginÚmini_batch_sizeÚforward_mini_batchedÚforward_non_mini_batchedÚforward)Úselfr   r   r   Úuse_mini_batched_versionr   Ú	__class__s         €ún/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/sentence_transformers/losses/MegaBatchMarginLoss.pyr   zMegaBatchMarginLoss.__init__   sL   ø€ ôT 	‰ÑÔØˆŒ
Ø.ˆÔØ.ˆÔØ.ˆÔÙ4L�t×0Ñ0ˆ�ÐRV×RoÑRoˆ�ó    c           
     ó,  — |\  }}t        |j                  «       «      }t        j                  «       5  | j                  j                  «        | j	                  |«      d   j                  «       }| j                  j                  «        d d d «       t        j                  t        «      t        |«      |j                  ¬«      }t        dt        |«      | j                  «      D �]  }|| j                  z   }	| j	                  |D �
ci c]  }
|
||
   ||	 “Œ c}
«      d   }|D �
ci c]  }
|
g “Œ }}
t        j                  «       5  t        j                  ||«      }|d|||	 z  z
  }t        j                  |d¬«      \  }}d d d «       D ]#  }|D ]  }
||
   j!                  ||
   |   «       Œ Œ% |D ]  }
t        j"                  ||
   «      ||
<   Œ | j	                  |D �
ci c]  }
|
||
   ||	 “Œ c}
«      d   }| j	                  |«      d   }|j$                  |j$                  k(  sJ ‚|j$                  |j$                  k(  sJ ‚t'        j(                  ||«      }t'        j(                  ||«      }t'        j*                  | j,                  |z
  «      t'        j*                  || j.                  z
  «      z   }|j1                  «       }|	t        «      k  s�Œò|j3                  «        �Œ S # 1 sw Y   �ŒexY wc c}
w c c}
w # 1 sw Y   �Œ„xY wc c}
w )NÚsentence_embedding)Údevicer   é   é   ©Údim)ÚlistÚkeysÚtorchÚno_gradr   ÚevalÚdetachÚtrainÚeyeÚlenr   Úranger   r   Úpytorch_cos_simÚmaxÚappendÚstackÚshapeÚFÚcosine_similarityÚrelur   r   ÚmeanÚbackward)r   Úsentence_featuresÚlabelsÚanchorÚpositiveÚfeature_namesÚall_positive_embÚdiagonal_matrixÚ	start_idxÚend_idxÚkeyÚ
anchor_embÚhard_negative_featuresÚ
cos_scoresÚnegative_scoresÚnegatives_maxÚnegatives_idsÚhard_negative_idÚpositive_embÚnegative_embÚ
pos_cosineÚ
neg_cosineÚlossess                          r   r   z(MegaBatchMarginLoss.forward_mini_batched^   s  € Ø,Ñˆ�Ü˜VŸ[™[›]Ó+ˆä�]‰]‹_ñ 	Ø�J‰J�O‰OÔØ#Ÿz™z¨(Ó3Ð4HÑI×PÑPÓRÐØ�J‰J×ÑÔ÷	ô
  Ÿ)™)¤CÐ(8Ó$9¼3Ð?OÓ;PÐYi×YpÑYpÔqˆô ˜q¤#Ð&6Ó"7¸×9MÑ9MÓNó (	"ˆIØ $×"6Ñ"6Ñ6ˆGØŸ™ÐTaÖ$bÈS S¨&°©+°iÀÐ*HÑ%HÒ$bÓcØ$ñˆJð :GÖ%G°# c¨2¡gÐ%GÐ"Ð%GÜ—‘“ñ QÜ!×1Ñ1°*Ð>NÓO�
à  _°Y¸wÐ%GÑ!GÑGð  ô 05¯y©y¸ÈaÔ/PÑ,�˜}÷Qð %2ò XÐ Ø(ò X�CØ*¨3Ñ/×6Ñ6°xÀ±}ÐEUÑ7VÕWñXðXð %ò W�Ü.3¯k©kÐ:PÐQTÑ:UÓ.VÐ& sÒ+ðWð  Ÿ:™:ÐXeÖ&fÐQT s¨H°S©M¸)ÀGÐ,LÑ'LÒ&fÓgØ$ñˆLð  Ÿ:™:Ð&<Ó=Ð>RÑSˆLà×#Ñ# |×'9Ñ'9Ò9Ð9Ð9Ø×#Ñ# |×'9Ñ'9Ò9Ð9Ð9ô ×,Ñ,¨Z¸ÓFˆJÜ×,Ñ,¨Z¸ÓFˆJÜ—V‘V˜D×0Ñ0°:Ñ=Ó>ÄÇÁÈ
ÐUY×UiÑUiÑHiÓAjÑjˆFØ—[‘[“]ˆFð œ˜Z›Ô(Ø—‘Ö!ðQ(	"ðT ˆ÷e	ñ 	üò %cùò &H÷Qñ Qüò 'gs*   ³AK-ÄK:
Ä&
K?Å<LÇ#L
Ë-K7ÌL	c                óê  — |D �cg c]  }| j                  |«      d   ‘Œ }}|\  }}t        j                  ||«      }t        j                  |«      }|dt        j
                  |j                  d|j                  iŽz  z
  }	t        j                  |	d¬«      \  }
}t        j                  | j                  |z
  «      t        j                  |
| j                  z
  «      z   }|j                  «       S c c}w )Nr   r   r   r   r   )r   r   r+   r#   Údiagonalr(   r/   r   r,   r0   r2   r   r   r3   )r   r5   r6   Úsentence_featureÚrepsÚembeddings_aÚembeddings_brA   Úpositive_scoresrB   rC   Ú_rJ   s                r   r   z,MegaBatchMarginLoss.forward_non_mini_batched—   sØ   € Ø[lÖmÐGW�—
‘
Ð+Ó,Ð-AÓBÐmˆÐmØ%)Ñ"ˆ�lä×)Ñ)¨,¸ÓEˆ
ÜŸ.™.¨Ó4ˆØ$Ø”—	‘	˜:×+Ñ+ÐF°J×4EÑ4EÑFÑFñ
ˆô !Ÿ9™9 _¸!Ô<Ñˆ�qÜ—‘˜×,Ñ,¨Ñ>Ó?Ä!Ç&Á&ÈÐY]×YmÑYmÑImÓBnÑnˆØ�{‰{‹}Ðùò ns   …C0c                 ó   — y)Naƒ  
@inproceedings{wieting-gimpel-2018-paranmt,
    title = "{P}ara{NMT}-50{M}: Pushing the Limits of Paraphrastic Sentence Embeddings with Millions of Machine Translations",
    author = "Wieting, John and Gimpel, Kevin",
    editor = "Gurevych, Iryna and Miyao, Yusuke",
    booktitle = "Proceedings of the 56th Annual Meeting of the Association for Computational Linguistics (Volume 1: Long Papers)",
    month = jul,
    year = "2018",
    address = "Melbourne, Australia",
    publisher = "Association for Computational Linguistics",
    url = "https://aclanthology.org/P18-1042",
    doi = "10.18653/v1/P18-1042",
    pages = "451--462",
}
© )r   s    r   ÚcitationzMegaBatchMarginLoss.citation¤   s   € ðr   )gš™™™™™é?g333333Ó?Té2   )r   r   r   Úfloatr   rW   r   Úboolr   ÚintÚreturnÚNone)r5   zIterable[dict[str, Tensor]]r6   r   rZ   r   )rZ   Ústr)	Ú__name__Ú
__module__Ú__qualname__r   r   r   ÚpropertyrU   Ú__classcell__)r   s   @r   r
   r
      s~   ø„ ð "%Ø!$Ø)-Ø!ðOpà"ðOpð ðOpð ð	Opð
 #'ðOpð ðOpð 
õOpób6órð òó ôr   r
   )Ú
__future__r   Úcollections.abcr   r#   Útorch.nn.functionalr   Ú
functionalr0   r   Úsentence_transformersr   r   ÚModuler
   rT   r   r   ú<module>rh      s,   ðÝ "å $ã ß Ð ß ç ;ôh˜"Ÿ)™)õ hr   