Ë
    l^(hÃ/  ã                  óÌ   — d dl mZ d dlZd dlmZmZ d dlmZmZ d dl	m
Z d dlmZ d dlmZ  ej                   e«      Z G d„ d	e«      Z G d
„ de«      Ze G d„ de«      «       Zy)é    )ÚannotationsN)Ú	dataclassÚfield)ÚOptionalÚUnion)ÚTrainingArguments)ÚParallelMode)ÚExplicitEnumc                  ó   — e Zd ZdZdZdZdZy)ÚBatchSamplersaÎ  
    Stores the acceptable string identifiers for batch samplers.

    The batch sampler is responsible for determining how samples are grouped into batches during training.
    Valid options are:

    - ``BatchSamplers.BATCH_SAMPLER``: **[default]** Uses :class:`~sentence_transformers.sampler.DefaultBatchSampler`, the default
      PyTorch batch sampler.
    - ``BatchSamplers.NO_DUPLICATES``: Uses :class:`~sentence_transformers.sampler.NoDuplicatesBatchSampler`,
      ensuring no duplicate samples in a batch. Recommended for losses that use in-batch negatives, such as:

        - :class:`~sentence_transformers.losses.MultipleNegativesRankingLoss`
        - :class:`~sentence_transformers.losses.CachedMultipleNegativesRankingLoss`
        - :class:`~sentence_transformers.losses.MultipleNegativesSymmetricRankingLoss`
        - :class:`~sentence_transformers.losses.CachedMultipleNegativesSymmetricRankingLoss`
        - :class:`~sentence_transformers.losses.MegaBatchMarginLoss`
        - :class:`~sentence_transformers.losses.GISTEmbedLoss`
        - :class:`~sentence_transformers.losses.CachedGISTEmbedLoss`
    - ``BatchSamplers.GROUP_BY_LABEL``: Uses :class:`~sentence_transformers.sampler.GroupByLabelBatchSampler`,
      ensuring that each batch has 2+ samples from the same label. Recommended for losses that require multiple
      samples from the same label, such as:

        - :class:`~sentence_transformers.losses.BatchAllTripletLoss`
        - :class:`~sentence_transformers.losses.BatchHardSoftMarginTripletLoss`
        - :class:`~sentence_transformers.losses.BatchHardTripletLoss`
        - :class:`~sentence_transformers.losses.BatchSemiHardTripletLoss`

    If you want to use a custom batch sampler, you can create a new Trainer class that inherits from
    :class:`~sentence_transformers.trainer.SentenceTransformerTrainer` and overrides the
    :meth:`~sentence_transformers.trainer.SentenceTransformerTrainer.get_batch_sampler` method. The
    method must return a class instance that supports ``__iter__`` and ``__len__`` methods. The former
    should yield a list of indices for each batch, and the latter should return the number of batches.

    Usage:
        ::

            from sentence_transformers import SentenceTransformer, SentenceTransformerTrainer, SentenceTransformerTrainingArguments
            from sentence_transformers.training_args import BatchSamplers
            from sentence_transformers.losses import MultipleNegativesRankingLoss
            from datasets import Dataset

            model = SentenceTransformer("microsoft/mpnet-base")
            train_dataset = Dataset.from_dict({
                "anchor": ["It's nice weather outside today.", "He drove to work."],
                "positive": ["It's so sunny.", "He took the car to the office."],
            })
            loss = MultipleNegativesRankingLoss(model)
            args = SentenceTransformerTrainingArguments(
                output_dir="checkpoints",
                batch_sampler=BatchSamplers.NO_DUPLICATES,
            )
            trainer = SentenceTransformerTrainer(
                model=model,
                args=args,
                train_dataset=train_dataset,
                loss=loss,
            )
            trainer.train()
    Úbatch_samplerÚno_duplicatesÚgroup_by_labelN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚBATCH_SAMPLERÚNO_DUPLICATESÚGROUP_BY_LABEL© ó    úa/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/sentence_transformers/training_args.pyr   r      s   „ ñ:ðx $€MØ#€MØ%�Nr   r   c                  ó   — e Zd ZdZdZdZy)ÚMultiDatasetBatchSamplersa•  
    Stores the acceptable string identifiers for multi-dataset batch samplers.

    The multi-dataset batch sampler is responsible for determining in what order batches are sampled from multiple
    datasets during training. Valid options are:

    - ``MultiDatasetBatchSamplers.ROUND_ROBIN``: Uses :class:`~sentence_transformers.sampler.RoundRobinBatchSampler`,
      which uses round-robin sampling from each dataset until one is exhausted.
      With this strategy, it's likely that not all samples from each dataset are used, but each dataset is sampled
      from equally.
    - ``MultiDatasetBatchSamplers.PROPORTIONAL``: **[default]** Uses :class:`~sentence_transformers.sampler.ProportionalBatchSampler`,
      which samples from each dataset in proportion to its size.
      With this strategy, all samples from each dataset are used and larger datasets are sampled from more frequently.

    Usage:
        ::

            from sentence_transformers import SentenceTransformer, SentenceTransformerTrainer, SentenceTransformerTrainingArguments
            from sentence_transformers.training_args import MultiDatasetBatchSamplers
            from sentence_transformers.losses import CoSENTLoss
            from datasets import Dataset, DatasetDict

            model = SentenceTransformer("microsoft/mpnet-base")
            train_general = Dataset.from_dict({
                "sentence_A": ["It's nice weather outside today.", "He drove to work."],
                "sentence_B": ["It's so sunny.", "He took the car to the bank."],
                "score": [0.9, 0.4],
            })
            train_medical = Dataset.from_dict({
                "sentence_A": ["The patient has a fever.", "The doctor prescribed medication.", "The patient is sweating."],
                "sentence_B": ["The patient feels hot.", "The medication was given to the patient.", "The patient is perspiring."],
                "score": [0.8, 0.6, 0.7],
            })
            train_legal = Dataset.from_dict({
                "sentence_A": ["This contract is legally binding.", "The parties agree to the terms and conditions."],
                "sentence_B": ["Both parties acknowledge their obligations.", "By signing this agreement, the parties enter into a legal relationship."],
                "score": [0.7, 0.8],
            })
            train_dataset = DatasetDict({
                "general": train_general,
                "medical": train_medical,
                "legal": train_legal,
            })

            loss = CoSENTLoss(model)
            args = SentenceTransformerTrainingArguments(
                output_dir="checkpoints",
                multi_dataset_batch_sampler=MultiDatasetBatchSamplers.PROPORTIONAL,
            )
            trainer = SentenceTransformerTrainer(
                model=model,
                args=args,
                train_dataset=train_dataset,
                loss=loss,
            )
            trainer.train()
    Úround_robinÚproportionalN)r   r   r   r   ÚROUND_ROBINÚPROPORTIONALr   r   r   r   r   P   s   „ ñ8ðt  €KØ!�Lr   r   c                  ó²   ‡ — e Zd ZU dZ edddi¬«      Zded<    eej                  ddi¬«      Z	d	ed
<    ee
j                  ddi¬«      Zded<   ˆ fd„Zˆ xZS )Ú$SentenceTransformerTrainingArgumentsa¥  
    SentenceTransformerTrainingArguments extends :class:`~transformers.TrainingArguments` with additional arguments
    specific to Sentence Transformers. See :class:`~transformers.TrainingArguments` for the complete list of
    available arguments.

    Args:
        output_dir (`str`):
            The output directory where the model checkpoints will be written.
        prompts (`Union[Dict[str, Dict[str, str]], Dict[str, str], str]`, *optional*):
            The prompts to use for each column in the training, evaluation and test datasets. Four formats are accepted:

            1. `str`: A single prompt to use for all columns in the datasets, regardless of whether the training/evaluation/test
               datasets are :class:`datasets.Dataset` or a :class:`datasets.DatasetDict`.
            2. `Dict[str, str]`: A dictionary mapping column names to prompts, regardless of whether the training/evaluation/test
               datasets are :class:`datasets.Dataset` or a :class:`datasets.DatasetDict`.
            3. `Dict[str, str]`: A dictionary mapping dataset names to prompts. This should only be used if your training/evaluation/test
               datasets are a :class:`datasets.DatasetDict` or a dictionary of :class:`datasets.Dataset`.
            4. `Dict[str, Dict[str, str]]`: A dictionary mapping dataset names to dictionaries mapping column names to
               prompts. This should only be used if your training/evaluation/test datasets are a
               :class:`datasets.DatasetDict` or a dictionary of :class:`datasets.Dataset`.

        batch_sampler (Union[:class:`~sentence_transformers.training_args.BatchSamplers`, `str`], *optional*):
            The batch sampler to use. See :class:`~sentence_transformers.training_args.BatchSamplers` for valid options.
            Defaults to ``BatchSamplers.BATCH_SAMPLER``.
        multi_dataset_batch_sampler (Union[:class:`~sentence_transformers.training_args.MultiDatasetBatchSamplers`, `str`], *optional*):
            The multi-dataset batch sampler to use. See :class:`~sentence_transformers.training_args.MultiDatasetBatchSamplers`
            for valid options. Defaults to ``MultiDatasetBatchSamplers.PROPORTIONAL``.
    NÚhelpzòThe prompts to use for each column in the datasets. Either 1) a single string prompt, 2) a mapping of column names to prompts, 3) a mapping of dataset names to prompts, or 4) a mapping of dataset names to a mapping of column names to prompts.)ÚdefaultÚmetadatazOptional[str]ÚpromptszThe batch sampler to use.zUnion[BatchSamplers, str]r   z'The multi-dataset batch sampler to use.z%Union[MultiDatasetBatchSamplers, str]Úmulti_dataset_batch_samplerc                óØ  •— t         ‰| �  «        t        | j                  «      | _        t	        | j
                  «      | _        d| _        d| _        | j                  t        j                  k(  r&| j                  dk7  rt        j                  d«       y y | j                  t        j                  k(  r9| j                  s,| j                  dk7  rt        j                  d«       d| _        y y y )NTFÚunusedzáCurrently using DataParallel (DP) for multi-gpu training, while DistributedDataParallel (DDP) is recommended for faster training. See https://sbert.net/docs/sentence_transformer/training/distributed.html for more information.z¶When using DistributedDataParallel (DDP), it is recommended to set `dataloader_drop_last=True` to avoid hanging issues with an uneven last batch. Setting `dataloader_drop_last=True`.)ÚsuperÚ__post_init__r   r   r   r&   Úprediction_loss_onlyÚddp_broadcast_buffersÚparallel_moder	   ÚNOT_DISTRIBUTEDÚ
output_dirÚloggerÚwarningÚDISTRIBUTEDÚdataloader_drop_last)ÚselfÚ	__class__s    €r   r*   z2SentenceTransformerTrainingArguments.__post_init__½   sÒ   ø€ Ü‰ÑÔä*¨4×+=Ñ+=Ó>ˆÔÜ+DÀT×EeÑEeÓ+fˆÔ(ð %)ˆÔ!ð &+ˆÔ"à×Ñ¤×!=Ñ!=Ò=ð �‰ (Ò*Ü—‘ðvõð +ð ×Ñ¤<×#;Ñ#;Ò;ÀD×D]ÒD]ð �‰ (Ò*Ü—‘ð;ôð )-ˆDÕ%ð E^Ð;r   )r   r   r   r   r   r%   Ú__annotations__r   r   r   r   r   r&   r*   Ú__classcell__)r5   s   @r   r!   r!   �   s„   ø… ññ: #Øàð dð
ô€Gˆ]ó ñ 05Ø×+Ñ+°vÐ?ZÐ6[ô0€MÐ,ó ñ JOØ)×6Ñ6À&ÐJsÐAtôJÐÐ!Fó ÷-ð -r   r!   )Ú
__future__r   ÚloggingÚdataclassesr   r   Útypingr   r   Útransformersr   ÚTransformersTrainingArgumentsÚtransformers.training_argsr	   Útransformers.utilsr
   Ú	getLoggerr   r0   r   r   r!   r   r   r   ú<module>rA      se   ðÝ "ã ß (ß "å KÝ 3Ý +à	ˆ×	Ñ	˜8Ó	$€ô?&�Lô ?&ôD<" ô <"ð~ ôL-Ð+Hó L-ó ñL-r   