§
    ™Štj©Z  ã                  óî   — d Z ddlmZ ddlZddlmZmZ ddlmZ ddl	m
Z
 ddlmZ ddlZddlZddlmZ dd	lmZmZ  G d
„ d¦  «        Zd/d„Zd0d„Zd1d„Z	 d2d3d„Zd4d„Zd5d „Zd6d,„Z G d-„ d.¦  «        ZdS )7ad  Shared GradCache machinery (https://huggingface.co/papers/2101.06983).

A "cached" loss trades compute for memory in three steps:

    (1) a quick embedding step without gradients/computation graphs, to get all the embeddings;
    (2) calculate the loss, backward up to the embeddings, and cache the gradients wrt. the embeddings;
    (3) a 2nd embedding step with gradients/computation graphs, connecting the cached gradients into
        the backward chain.

Only one mini-batch of activations is alive at a time in (1) and (3), which is what bounds the memory.
The second forward pass reproduces the first exactly by replaying its RNG state (:class:`RandContext`),
so dropout draws the same masks and the cached gradients belong to the embeddings they are applied to.

:class:`CachedLossMixin` implements all of this. A loss only has to provide ``calculate_loss``.
é    )ÚannotationsN)ÚIterableÚIterator)Únullcontext)Úpartial)ÚAny)ÚTensor)Úget_device_statesÚset_device_statesc                  ó*   — e Zd ZdZdd„Zdd„Zdd„ZdS )	ÚRandContextz·Snapshot the CPU/CUDA/MPS RNG at init and restore it on enter, so the cached second forward
    replays the first's randomness (e.g. dropout). Ref: https://github.com/luyug/GradCache.ÚreturnÚNonec                ó  — t          j        ¦   «         | _        t          d„ |D ¦   «         ¦  «        rt           j                             ¦   «         nd | _        t          d„ |D ¦   «         ¦  «        }t          |Ž \  | _        | _	        d S )Nc              3  ój   K  — | ].}t          |t          j        ¦  «        o|j        j        d k    V — Œ/dS ©ÚmpsN©Ú
isinstanceÚtorchr	   ÚdeviceÚtype©Ú.0Úts     úi/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/sentence_transformers/base/losses/gradcache.pyú	<genexpr>z'RandContext.__init__.<locals>.<genexpr>)   s<   è è € Ð[Ð[Èa•:˜a¥¤Ñ.Ô.ÐI°1´8´=ÀEÒ3IÐ[Ð[Ð[Ð[Ð[Ð[ó    c              3  ón   K  — | ]0}t          |t          j        ¦  «        r|j        j        d k    °,|V — Œ1dS r   r   r   s     r   r   z'RandContext.__init__.<locals>.<genexpr>,   sE   è è € ÐsÐs a½:ÀaÍÌÑ;VÔ;VÐsÐ[\Ô[cÔ[hÐlqÒ[qÐ[q Ð[qÐ[qÐ[qÐ[qÐsÐsr   )
r   Úget_rng_stateÚfwd_cpu_stateÚanyr   Úfwd_mps_stateÚtupler
   Úfwd_gpu_devicesÚfwd_gpu_states)ÚselfÚtensorsÚnon_mps_tensorss      r   Ú__init__zRandContext.__init__#   sŒ   € Ý"Ô0Ñ2Ô2ˆÔõ
 Ð[Ð[ÐSZÐ[Ñ[Ô[Ñ[Ô[ð�EŒI×#Ò#Ñ%Ô%Ð%àð 	Ôõ
  ÐsÐs¨7ÐsÑsÔsÑsÔsˆÝ4EÀÐ4WÑ1ˆÔ˜dÔ1Ð1Ð1r   c                ó�  — t           j                             | j        d¬¦  «        | _        | j                             ¦   «          t          j        | j        ¦  «         | j        �Gt           j	         
                    ¦   «         | _        t           j	                             | j        ¦  «         t          | j        | j        ¦  «         d S )NT)ÚdevicesÚenabled)r   ÚrandomÚfork_rngr%   Ú_forkÚ	__enter__Úset_rng_stater!   r#   r   r    Ú_mps_state_outsider   r&   )r'   s    r   r1   zRandContext.__enter__/   s    € Ý”\×*Ò*°4Ô3GÐQUÐ*ÑVÔVˆŒ
ØŒ
×ÒÑÔÐÝÔ˜DÔ.Ñ/Ô/Ð/ØÔÐ)õ ',¤i×&=Ò&=Ñ&?Ô&?ˆDÔ#ÝŒI×#Ò# DÔ$6Ñ7Ô7Ð7Ý˜$Ô.°Ô0CÑDÔDÐDÐDÐDr   c                ó¢   — | j         �$t          j                             | j        ¦  «         | j                             |||¦  «         d | _        d S ©N)r#   r   r   r2   r3   r0   Ú__exit__)r'   Úexc_typeÚexc_valÚexc_tbs       r   r6   zRandContext.__exit__:   sI   € ØÔÐ)ÝŒI×#Ò# DÔ$;Ñ<Ô<Ð<ØŒ
×Ò˜H g¨vÑ6Ô6Ð6ØˆŒ
ˆ
ˆ
r   N)r   r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r*   r1   r6   © r   r   r   r      sd   € € € € € ð_ð _ð
Xð 
Xð 
Xð 
Xð	Eð 	Eð 	Eð 	Eðð ð ð ð ð r   r   Úsentence_featureúdict[str, Any]r   Úintc                ó  — d| v rt          | d         ¦  «        dz
  S dD ];}|| v r5t          | |         t          j        ¦  «        r| |         j        d         c S Œ<t          d„ |                      ¦   «         D ¦   «         ¦  «        S )aH  Get the number of samples in sentence features, handling both padded and flattened inputs.

    With padded inputs, the batch size is the first dimension of any tensor.
    With flattened inputs (from ``DataCollatorWithFlattening``), the batch size is derived
    from ``cu_seq_lens_q`` which has shape ``(num_seqs + 1,)``.
    Úcu_seq_lens_qé   )Ú	input_idsÚattention_maskr   c              3  óz   K  — | ]6}t          |t          j        ¦  «        r|j        d k    ¯'|j        d          V — Œ7dS )r   N)r   r   r	   ÚndimÚshape)r   Úvalues     r   r   z"_get_batch_size.<locals>.<genexpr>O   sU   è è € ð ð Ø ÅÈEÕSXÔS_ÑA`ÔA`ðØejÔeoÐrsÒesÐesˆŒ�AŒÐesÐesÐesÐesðð r   )Úlenr   r   r	   rI   ÚnextÚvalues)r?   Úkeys     r   Ú_get_batch_sizerO   A   s±   € ð Ð*Ð*Ð*ÝÐ# OÔ4Ñ5Ô5¸Ñ9Ð9ð /ð 2ð 2ˆØÐ"Ð"Ð"¥zÐ2BÀ3Ô2GÍÌÑ'VÔ'VÐ"Ø# CÔ(Ô.¨qÔ1Ð1Ð1Ð1øÝð ð Ø$4×$;Ò$;Ñ$=Ô$=ðñ ô ñ ô ð r   ÚbeginÚendc                ó:	  — d| v�rt          | ¦  «        }t          ||¦  «        }i }dD �]y\  }}}|                      |¦  «        }|                      |¦  «        }	|�|	€Œ6|                      |¦  «        }
|
�z|
                     d¬¦  «        }|dk    rdn)t	          ||dz
                                ¦   «         ¦  «        }t	          ||dz
                                ¦   «         ¦  «        }||f||<   n|j        d         |k    r||}}nŒÞ||k     r‰|                     d¬¦  «        }|                     d¬¦  «        }|dk    rdn)t	          ||dz
                                ¦   «         ¦  «        }t	          ||dz
                                ¦   «         ¦  «        }nd\  }}||f||<   �Œ{d}d}|                      d¦  «        }t          |t          j
        ¦  «        r�|j        d	k    r’|j        d         }|||…                              d¬¦  «        }|                     ¦   «         rSt	          |                     ¦   «                              ¦   «                              ¦   «         ¦  «        }|dz   |k     r|dz   }i }|                      ¦   «         D ]\  }}t          |t          j
        ¦  «        s|||<   Œ%||v r||         \  }}|||…         ||<   ŒB|�.|j        d	k    r#|j        d         |k    r|||…d|…f         ||<   Œr|||…         ||<   Œ€|S | d         }t!          |¦  «        dz
  }t          ||¦  «        }t	          ||                              ¦   «         ¦  «        }t	          ||                              ¦   «         ¦  «        }t	          |d
                              ¦   «         ¦  «        }|||dz   …         ||         z
  }i }|                      ¦   «         D ]Ã\  }}|dv r|||<   Œ|dv rL|dd…         |dd
…         z
  } t	          |                      ¦   «                              ¦   «         ¦  «        ||<   Œ_|dk    r|d||…f         |z
  ||<   Œxt          |t          j
        ¦  «        r,|j        dk    r!|j        d
         |k    r|d||…f         ||<   Œ¾|||<   ŒÄ|S )a^  Create a mini-batch from sentence features, handling padded, flattened, and VLM inputs.

    With padded inputs, this slices along the batch dimension and drops trailing padding columns shared
    by the whole mini-batch, so short-sequence mini-batches aren't embedded at the full batch width
    (else ``mini_batch_num_tokens`` would not bound memory). Leading padding is kept to preserve positions.
    With flattened inputs (from ``DataCollatorWithFlattening``), this extracts the token ranges
    for sequences ``begin:end`` and rebuilds the metadata (``cu_seq_lens_q``, ``seq_idx``, etc.).

    VLMs like Qwen2-VL flatten per-sample visual tokens into a single tensor
    (e.g. ``pixel_values`` shape ``(total_visual_tokens, hidden_dim)``) with a grid tensor
    (e.g. ``image_grid_thw`` shape ``(num_items, 3)``) whose per-row product gives the token
    count per item.  ``num_images_per_sample`` / ``num_videos_per_sample`` (precomputed by
    ``Transformer.preprocess``) map grid rows to samples; when unavailable we fall back to
    assuming one grid row per sample when ``grid.shape[0] == batch_size``.
    rC   ))Úimage_grid_thwÚpixel_valuesÚnum_images_per_sample)Úvideo_grid_thwÚpixel_values_videosÚnum_videos_per_sampleNr   ©ÚdimrD   )r   r   rF   é   éÿÿÿÿ)rC   Úcu_seq_lens_k)Úmax_length_qÚmax_length_kÚseq_idx.)rO   ÚminÚgetÚcumsumrA   ÚitemrI   Úprodr   r   r	   rH   r"   ÚnonzeroÚmaxÚitemsrK   )!r?   rP   rQ   Ú
batch_sizeÚcustom_rangesÚgrid_keyÚ	pixel_keyÚ	count_keyÚgridrT   Únum_per_sampleÚcumsum_itemsÚ
grid_beginÚgrid_endÚtokens_per_itemÚtoken_cumsumÚtoken_beginÚ	token_endÚtoken_axis_endÚ	seq_widthrF   Úactive_columnsÚlast_activeÚresultrN   rJ   Úr_beginÚr_endrC   Únum_seqsÚtotal_tokensÚnew_cu_seq_lensÚmb_seq_lenss!                                    r   Ú_create_minibatchr‚   T   s  € ð  Ð.Ð.Ñ.Ý$Ð%5Ñ6Ô6ˆ
Ý�#�zÑ"Ô"ˆà46ˆð/
ð 	@ñ 	@Ñ*ˆH�i ð $×'Ò'¨Ñ1Ô1ˆDØ+×/Ò/°	Ñ:Ô:ˆLØˆ|˜|Ð3Øà-×1Ò1°)Ñ<Ô<ˆNØÐ)Ø-×4Ò4¸Ð4Ñ;Ô;�Ø"'¨1¢* *˜Q˜Qµ#°lÀ5È1Á9Ô6M×6RÒ6RÑ6TÔ6TÑ2UÔ2U�
Ý˜|¨C°!©GÔ4×9Ò9Ñ;Ô;Ñ<Ô<�Ø+5°xÐ*@�˜hÑ'Ð'Ø”˜A” *Ò,Ð,Ø',¨c˜H�
�
àà˜HÒ$Ð$Ø"&§)¢)° )Ñ"2Ô"2�Ø.×5Ò5¸!Ð5Ñ<Ô<�Ø#-°¢? ?˜a˜a½¸LÈÐVWÉÔ<X×<]Ò<]Ñ<_Ô<_Ñ8`Ô8`�Ý ¨X¸©\Ô :× ?Ò ?Ñ AÔ AÑBÔB�	�	à)-Ñ&�˜YØ(3°YÐ'?ˆM˜)Ñ$Ñ$ð &*ˆØ $ˆ	Ø)×-Ò-Ð.>Ñ?Ô?ˆÝ�n¥e¤lÑ3Ô3ð 	5¸Ô8KÈqÒ8PÐ8PØ&Ô,¨QÔ/ˆIØ+¨E°#¨IÔ6×:Ò:¸qÐ:ÑAÔAˆNØ×!Ò!Ñ#Ô#ð 5Ý! .×"8Ò"8Ñ":Ô":×">Ò">Ñ"@Ô"@×"EÒ"EÑ"GÔ"GÑHÔH�Ø ‘? YÒ.Ð.Ø%0°1¡_�Nà!#ˆØ*×0Ò0Ñ2Ô2ð 		/ð 		/‰JˆC�Ý˜e¥U¤\Ñ2Ô2ð /Ø#��s‘�Ø˜Ð%Ð%Ø!.¨sÔ!3‘�˜Ø# G¨E MÔ2��s‘�ØÐ+°´
¸a²°ÀEÄKÐPQÄNÐV_ÒD_ÐD_Ø# E¨# I¨°¨Ð$>Ô?��s‘�à# E¨# IÔ.��s‘�Øˆà$ _Ô5€MÝ�=Ñ!Ô! AÑ%€HÝ
ˆc�8Ñ
Ô
€Cå�m EÔ*×/Ò/Ñ1Ô1Ñ2Ô2€KÝ�M #Ô&×+Ò+Ñ-Ô-Ñ.Ô.€IÝ�} RÔ(×-Ò-Ñ/Ô/Ñ0Ô0€Là# E¨C°!©G OÔ4°}ÀUÔ7KÑK€Oà€FØ&×,Ò,Ñ.Ô.ð  ð  ‰
ˆˆUØÐ4Ð4Ð4Ø)ˆF�3‰KˆKØÐ4Ð4Ð4Ø)¨!¨"¨"Ô-°ÀÀÀÔ0DÑDˆKÝ˜kŸošoÑ/Ô/×4Ò4Ñ6Ô6Ñ7Ô7ˆF�3‰KˆKØ�IÒÐØ  [°Ð%:Ð :Ô;¸eÑCˆF�3‰KˆKÝ˜�uœ|Ñ,Ô,ð 	 °´¸q²°ÀUÄ[ÐQSÄ_ÐXdÒEdÐEdð    [°Ð%:Ð :Ô;ˆF�3‰KˆKàˆF�3‰KˆKØ€Mr   Úlossr   Úboolc                ó$   — t          | dd¦  «        S )a  Whether ``loss`` defers its backward pass to a hook on the loss tensor it returns.

    Such a loss re-embeds each mini-batch during the *backward* pass, by which time a decorator that
    patched ``SentenceTransformer.forward`` for the duration of the forward pass has been removed
    again. ``MatryoshkaLoss`` and ``AdaptiveLayerLoss`` both work by patching that forward, so they
    have to treat these losses specially: MatryoshkaLoss decorates ``calculate_loss`` instead, and
    AdaptiveLayerLoss warns that the combination is unsupported.

    Losses report this by setting ``uses_gradient_cache``. :class:`CachedLossMixin` sets it to True,
    and a loss that can turn the caching off at construction time (``MegaBatchMarginLoss``) overrides
    it per instance.
    Úuses_gradient_cacheF)Úgetattr)rƒ   s    r   r†   r†   ½   s   € õ �4Ð.°Ñ6Ô6Ð6r   Úmini_batch_sizeÚmini_batch_num_tokensú
int | Noneúlist[tuple[int, int]]c                ó0  ‡‡— t          | ¦  «        Š|€ˆˆfd„t          d‰‰¦  «        D ¦   «         S d| v r#| d         dd…                              ¦   «         }nVd| v rC| d                              d¬¦  «                             d¬¦  «                             ¦   «         }nt          d¦  «        ‚g }d}|‰k     r]|dk    r||dz
           nd}t          j        |||z   ¦  «        }t          ||dz   ¦  «        }| 	                    ||f¦  «         |}|‰k     °]|S )	a,  Compute the ``(begin, end)`` sequence ranges that split a batch into mini-batches.

    If ``mini_batch_num_tokens`` is None, every range spans ``mini_batch_size`` sequences.
    Otherwise, each range greedily packs as many sequences as possible while keeping the total
    number of non-padding tokens at or below ``mini_batch_num_tokens``. A single sequence whose
    length exceeds the budget forms its own mini-batch. Per-sequence token counts are read from
    ``cu_seq_lens_q`` for flattened inputs, or from the attention mask for padded inputs.
    Nc                ó:   •— g | ]}|t          |‰z   ‰¦  «        f‘ŒS r>   )ra   )r   rP   ri   rˆ   s     €€r   ú
<listcomp>z%_minibatch_ranges.<locals>.<listcomp>Ü   s-   ø€ ÐuÐuÐuÀe��˜E OÑ3°ZÑ@Ô@ÐAÐuÐuÐur   r   rC   rD   rF   rY   zÈmini_batch_num_tokens requires per-sequence token counts, but the tokenized inputs contain neither 'cu_seq_lens_q' (flattened inputs) nor 'attention_mask' (padded inputs). Use mini_batch_size instead.)
rO   ÚrangeÚtolistÚsumrc   Ú
ValueErrorÚbisectÚbisect_rightrg   Úappend)	r?   rˆ   r‰   Úcumulative_num_tokensÚrangesrP   Úprevious_num_tokensrQ   ri   s	    `      @r   Ú_minibatch_rangesr™   Í   sg  øø€ õ !Ð!1Ñ2Ô2€JØÐ$ØuÐuÐuÐuÐuÍuÐUVÐXbÐdsÑOtÔOtÐuÑuÔuÐuàÐ*Ð*Ð*à 0°Ô AÀ!À"À"Ô E× LÒ LÑ NÔ NÐÐØ	Ð-Ð	-Ð	-Ø 0Ð1AÔ B× FÒ FÈ1Ð FÑ MÔ M× TÒ TÐYZÐ TÑ [Ô [× bÒ bÑ dÔ dÐÐåð+ñ
ô 
ð 	
ð %'€FØ€EØ
�*Ò
Ð
ØBGÈ!Â)À)Ð3°E¸A±IÔ>Ð>ÐQRÐÝÔ!Ð"7Ð9LÐOdÑ9dÑeÔeˆå�#�u˜q‘yÑ!Ô!ˆØ�Š�u˜c�lÑ#Ô#Ð#Øˆð �*Ò
Ð
ð €Mr   r   c                ó8   — | �| dk    rt          d¦  «        ‚d S d S )Nr   z9mini_batch_num_tokens must be a positive integer or None.)r’   )r‰   s    r   Ú_validate_mini_batch_num_tokensr›   ö   s0   € ØÐ(Ð-BÀaÒ-GÐ-GÝÐTÑUÔUÐUð )Ð(Ð-GÐ-Gr   Úmodelc                óÚ   ‡— ddl m}mŠ t          | d         |¦  «        r)d„ | d         j                             ¦   «         D ¦   «         n| d         g}t          ˆfd„|D ¦   «         ¦  «        S )a-  Whether the model embeds its inputs with a StaticEmbedding, directly or behind a Router.

    StaticEmbedding features are an EmbeddingBag (``input_ids``, ``offsets``) with no batch dimension,
    so they cannot be sliced into mini-batches, meaning losses that mini-batch must reject such models.
    r   )ÚRouterÚStaticEmbeddingc                ó   — g | ]
}|d          ‘ŒS )r   r>   )r   Úroutes     r   rŽ   z.has_static_embedding_input.<locals>.<listcomp>  s   € Ð=Ð=Ð=�eˆˆqŒÐ=Ð=Ð=r   c              3  ó8   •K  — | ]}t          |‰¦  «        V — Œd S r5   )r   )r   ÚmodulerŸ   s     €r   r   z-has_static_embedding_input.<locals>.<genexpr>  s-   øè è € ÐOÐO°v�z˜& /Ñ2Ô2ÐOÐOÐOÐOÐOÐOr   )Ú2sentence_transformers.sentence_transformer.modulesrž   rŸ   r   Úsub_modulesrM   r"   )rœ   rž   Úinput_modulesrŸ   s      @r   Úhas_static_embedding_inputr§   û   s•   ø€ ð [ÐZÐZÐZÐZÐZÐZÐZõ BLÈEÐRSÌHÐV\ÑA]ÔA]ÐmÐ=Ð=˜u QœxÔ3×:Ò:Ñ<Ô<Ð=Ñ=Ô=Ð=ÐdiÐjkÔdlÐcmð õ ÐOÐOÐOÐOÀÐOÑOÔOÑOÔOÐOr   Úgrad_outputr	   Úsentence_featuresúIterable[dict[str, Tensor]]Úloss_objÚcacheúlist[list[Tensor]]Úrandom_statesúlist[list[RandContext]]r—   ú"list[list[tuple[int, int]] | None]c                óì  — t          j        ¦   «         5  t          ||||¦  «        D ]³\  }}}}	t          |                     |dd||	¬¦  «        |¦  «        D ]ƒ\  ^}
}}|
j        sŒt          j        |
                     ¦   «                              ¦   «         |                     ¦   «                              ¦   «         ¦  «        | z  }|                     ¦   «          Œ„Œ´	 ddd¦  «         dS # 1 swxY w Y   dS )a%  A backward hook to backpropagate the cached gradients mini-batch by mini-batch.

    ``loss_obj`` only needs an ``embed_minibatch_iter(sentence_feature, with_grad, copy_random_state,
    random_states, ranges)`` iterator whose items *start* with the embeddings tensor. Extra elements
    are ignored (e.g. ``CachedGISTEmbedLoss`` also yields the guide model's embeddings).
    :class:`CachedLossMixin` provides the standard implementation, and the cross-encoder
    ``CachedMultipleNegativesRankingLoss`` satisfies the contract with its own adapter.

    ``cache``, ``random_states`` and ``ranges`` belong to one specific forward pass and are passed in
    rather than read off ``loss_obj``, so that a second forward pass before the first backward pass
    cannot make this hook back-propagate the wrong batch's gradients. Replaying the forward pass's
    mini-batch ``ranges`` also matters because modules may modify the features in place between the
    two passes (e.g. ``Pooling`` with ``include_prompt=False`` zeroes prompt tokens in the attention
    mask), which would change recomputed token-budget boundaries.

    Every mini-batch is scaled by ``grad_output``, which is whatever the outer backward pass hands us,
    so the fp16 gradient scaler and the gradient accumulation division reach all of them.
    TF)r?   Ú	with_gradÚcopy_random_stater®   r—   N)	r   Úenable_gradÚzipÚembed_minibatch_iterÚrequires_gradÚdotÚflattenÚfloatÚbackward)r¨   r©   r«   r¬   r®   r—   r?   ÚgradÚrandom_stateÚcolumn_rangesÚreps_mbÚ_Úgrad_mbÚ	surrogates                 r   Ú_backward_hookrÃ   
  sk  € õ4 
Ô	Ñ	Ô	ð %ð %ÝCFØ˜u m°VñD
ô D
ð 	%ð 	%Ñ?Ð˜d L°-õ +.Ø×-Ò-Ø%5Ø"Ø&+Ø".Ø(ð .ñ ô ð ñ	+ô 	+ð %ð %Ñ&��˜1˜wð Ô,ð ð õ "œI g§o¢oÑ&7Ô&7×&=Ò&=Ñ&?Ô&?ÀÇÂÑARÔAR×AXÒAXÑAZÔAZÑ[Ô[Ð^iÑi�	Ø×"Ò"Ñ$Ô$Ð$Ð$ð#%ð	%ð%ð %ð %ñ %ô %ð %ð %ð %ð %ð %ð %ð %øøøð %ð %ð %ð %ð %ð %s   ”CC)Ã)C-Ã0C-c                  óŒ   — e Zd ZU dZded<   ded<   dZded<   d	Zd
ed<   dZdZd
ed<   	 d*d	dœd+d„Z		 d*d,d „Z
	 	 d-d.d&„Zd*d/d)„ZdS )0ÚCachedLossMixina.  The GradCache forward pass, shared by the losses that cache the gradients wrt. their embeddings.

    Subclasses must be an ``nn.Module`` holding a ``model``, must set ``mini_batch_size``, and must
    implement :meth:`calculate_loss`. They then call :meth:`forward_cached` from their ``forward``.
    r   rœ   rA   rˆ   NrŠ   r‰   Fr„   Úshow_progress_barTr†   ©Úwith_backwardÚrepsr­   ÚlabelsúTensor | NonerÈ   r   r	   c               ó   — t           ‚)al  Compute the loss over the whole batch, from the per-mini-batch embeddings.

        When ``with_backward`` is set, back-propagate the loss (chunk by chunk, if the implementation
        chunks it) and return the detached total, so that no part of the loss graph outlives its own
        backward pass. Losses that don't need the labels simply ignore them.
        )ÚNotImplementedError)r'   rÉ   rÊ   rÈ   s       r   Úcalculate_losszCachedLossMixin.calculate_lossO  s
   € õ "Ð!r   r?   údict[str, Tensor]rP   rQ   r²   r³   r½   úRandContext | Noneú!tuple[Tensor, RandContext | None]c                óf  — |rt           nt          j        }|€t          ¦   «         n|}t          |||¦  «        }	|5   |¦   «         5  |rt	          |	                     ¦   «         Ž nd}|                      |	¦  «        d         }
ddd¦  «         n# 1 swxY w Y   ddd¦  «         n# 1 swxY w Y   |
|fS )zEmbed a mini-batch of inputs.NÚsentence_embedding)r   r   Úno_gradr‚   r   rM   rœ   )r'   r?   rP   rQ   r²   r³   r½   Úgrad_contextÚrandom_state_contextÚsentence_feature_minibatchrÉ   s              r   Úembed_minibatchzCachedLossMixin.embed_minibatchZ  s`  € ð '0ÐB•{�{µU´]ˆØ0<Ð0D�{™}œ}˜}È,ÐÝ%6Ð7GÈÐPSÑ%TÔ%TÐ"Ø!ð 	Tð 	TØ�‘”ð Tð TØTeÐo�{Ð,F×,MÒ,MÑ,OÔ,OÐPÐPÐko�Ø—z’zÐ"<Ñ=Ô=Ð>RÔS�ðTð Tð Tñ Tô Tð Tð Tð Tð Tð Tð Tøøøð Tð Tð Tð Tð	Tð 	Tð 	Tñ 	Tô 	Tð 	Tð 	Tð 	Tð 	Tð 	Tð 	Tøøøð 	Tð 	Tð 	Tð 	Tð �\Ð!Ð!s5   »B$Á;BÂB$ÂB	ÂB$ÂB	ÂB$Â$B(Â+B(r®   úlist[RandContext] | Noner—   úlist[tuple[int, int]] | Noneú+Iterator[tuple[Tensor, RandContext | None]]c           
   #  óø   K  — |€t          || j        | j        ¦  «        }t          t	          j        |d| j         ¬¦  «        ¦  «        D ]/\  }\  }}|                      ||||||€dn||         ¬¦  «        V — Œ0dS )zUDo a forward pass on every mini-batch of the input features and yield the embeddings.NzEmbed mini-batches)ÚdescÚdisable)r?   rP   rQ   r²   r³   r½   )r™   rˆ   r‰   Ú	enumerateÚtqdmrÆ   rØ   )	r'   r?   r²   r³   r®   r—   ÚirP   rQ   s	            r   r¶   z$CachedLossMixin.embed_minibatch_iterm  sÁ   è è € ð ˆ>Ý&Ð'7¸Ô9MÈtÔOiÑjÔjˆFÝ(ÝŒIØØ)Ø Ô2Ð2ðñ ô ñ 
ô  
ð 	ð 	‰OˆA‰|��sð ×&Ò&Ø!1ØØØ#Ø"3Ø%2Ð%:˜T˜TÀÈaÔ@Pð 'ñ ô ð ð ð ð ð	ð 	r   r©   rª   c           
     óŒ  ‡ — t          |¦  «        }t          j        ¦   «         }ˆ fd„|D ¦   «         }g }g }t          ||¦  «        D ] \  }}g }	g }
‰                      |d||¬¦  «        D ]S\  }}|	                     |                     ¦   «                              ¦   «         ¦  «         |
                     |¦  «         ŒT|                     |	¦  «         |                     |
¦  «         Œ¡|s‰                      ||¦  «        S ‰                      ||d¬¦  «        }|                     ¦   «                              ¦   «         }d„ |D ¦   «         }d„ t          |¦  «        D ¦   «         }|r3t          d‰ j        j        › d	d
                     |¦  «        › d�¦  «        ‚|                     t          t           |‰ |||¬¦  «        ¦  «         |S )zDRun the three-step GradCache forward pass. See the module docstring.c                óF   •— g | ]}t          |‰j        ‰j        ¦  «        ‘ŒS r>   )r™   rˆ   r‰   )r   r?   r'   s     €r   rŽ   z2CachedLossMixin.forward_cached.<locals>.<listcomp>�  s<   ø€ ð 
ð 
ð 
à õ Ð.°Ô0DÀdÔF`ÑaÔað
ð 
ð 
r   F)r?   r²   r³   r—   TrÇ   c                ó&   — g | ]}d „ |D ¦   «         ‘ŒS )c                ó   — g | ]	}|j         ‘Œ
S r>   )r¼   )r   Úreps     r   rŽ   z=CachedLossMixin.forward_cached.<locals>.<listcomp>.<listcomp>°  s   € Ð.Ð.Ð.˜s�#”(Ð.Ð.Ð.r   r>   )r   Úrep_mbss     r   rŽ   z2CachedLossMixin.forward_cached.<locals>.<listcomp>°  s'   € ÐCÐCÐC°7Ð.Ð. gÐ.Ñ.Ô.ÐCÐCÐCr   c                ód   — g | ]-\  }}t          d „ |D ¦   «         ¦  «        ¯t          |¦  «        ‘Œ.S )c              3  ó   K  — | ]}|d u V — Œ	d S r5   r>   )r   Úgs     r   r   z<CachedLossMixin.forward_cached.<locals>.<listcomp>.<genexpr>±  s*   è è € ÐSpÐSpÐbcÐTUÐY]ÐT]ÐSpÐSpÐSpÐSpÐSpÐSpr   )r"   Ústr)r   ÚindexÚgrad_mbss      r   rŽ   z2CachedLossMixin.forward_cached.<locals>.<listcomp>±  s@   € ÐqÐqÐq©¨°ÕPSÐSpÐSpÐgoÐSpÑSpÔSpÑPpÔPpÐq�#˜e™*œ*ÐqÐqÐqr   zThe loss computation of z did not use input column(s) z, z : their embeddings received no gradient. Every input column is embedded (twice, with gradient caching), so remove the unused column(s) from the dataset instead.)r©   r«   r¬   r®   r—   )Úlistr   Úis_grad_enabledrµ   r¶   r•   ÚdetachÚrequires_grad_rÎ   rß   r’   Ú	__class__r:   ÚjoinÚregister_hookr   rÃ   )r'   r©   rÊ   Úgrad_enabledr—   rÉ   r®   r?   r¾   Úreps_mbsÚrandom_state_mbsr¿   r½   rƒ   r¬   Úunused_columnss   `               r   Úforward_cachedzCachedLossMixin.forward_cachedˆ  s8  ø€ å Ð!2Ñ3Ô3ÐÝÔ,Ñ.Ô.ˆð
ð 
ð 
ð 
à$5ð
ñ 
ô 
ˆð ˆØˆÝ/2Ð3DÀfÑ/MÔ/Mð 	3ð 	3Ñ+Ð˜mØˆHØ!ÐØ)-×)BÒ)BØ!1Øð #/Ø$ð *Cñ *ô *ð 	6ð 	6Ñ%�˜ð —’ §¢Ñ 0Ô 0× ?Ò ?Ñ AÔ AÑBÔBÐBØ ×'Ò'¨Ñ5Ô5Ð5Ð5Ø�KŠK˜Ñ!Ô!Ð!Ø× Ò Ð!1Ñ2Ô2Ð2Ð2àð 	5à×&Ò& t¨VÑ4Ô4Ð4ð ×"Ò" 4¨¸tÐ"ÑDÔDˆØ�{Š{‰}Œ}×+Ò+Ñ-Ô-ˆØCÐC¸dÐCÑCÔCˆØqÐq½IÀeÑ<LÔ<LÐqÑqÔqˆØð 	õ ð#¨4¬>Ô+Bð #ð #Ø—9’9˜^Ñ,Ô,ð#ð #ð #ñô ð ð 	×ÒÝÝØ"3ØØØ+Øðñ ô ñ		
ô 		
ð 		
ð ˆr   r5   )rÉ   r­   rÊ   rË   rÈ   r„   r   r	   )r?   rÏ   rP   rA   rQ   rA   r²   r„   r³   r„   r½   rÐ   r   rÑ   )NN)r?   rÏ   r²   r„   r³   r„   r®   rÙ   r—   rÚ   r   rÛ   )r©   rª   rÊ   rË   r   r	   )r:   r;   r<   r=   Ú__annotations__r‰   rÆ   Úrequires_media_countsr†   rÎ   rØ   r¶   rù   r>   r   r   rÅ   rÅ   <  s  € € € € € € ðð ð €J€J�JØÐÐÑØ(,ÐÐ,Ð,Ð,Ñ,Ø#ÐÐ#Ð#Ð#Ñ#ð !Ðð !%ÐÐ$Ð$Ð$Ñ$ð AEð	"Ø_dð	"ð 	"ð 	"ð 	"ð 	"ð 	"ð$ ,0ð"ð "ð "ð "ð "ð0 37Ø/3ðð ð ð ð ð6@ð @ð @ð @ð @ð @ð @r   rÅ   )r?   r@   r   rA   )r?   r@   rP   rA   rQ   rA   r   r@   )rƒ   r   r   r„   r5   )r?   r@   rˆ   rA   r‰   rŠ   r   r‹   )r‰   rŠ   r   r   )rœ   r   r   r„   )r¨   r	   r©   rª   r«   r   r¬   r­   r®   r¯   r—   r°   r   r   )r=   Ú
__future__r   r“   Úcollections.abcr   r   Ú
contextlibr   Ú	functoolsr   Útypingr   r   rà   r	   Útorch.utils.checkpointr
   r   r   rO   r‚   r†   r™   r›   r§   rÃ   rÅ   r>   r   r   ú<module>r     sÀ  ððð ð  #Ð "Ð "Ð "Ð "Ð "à €€€Ø .Ð .Ð .Ð .Ð .Ð .Ð .Ð .Ø "Ð "Ð "Ð "Ð "Ð "Ø Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð à €€€Ø €€€Ø Ð Ð Ð Ð Ð Ø GÐ GÐ GÐ GÐ GÐ GÐ GÐ Gðð ð ð ð ñ ô ð ðDð ð ð ð&fð fð fð fðR7ð 7ð 7ð 7ð& )-ð&ð &ð &ð &ð &ðRVð Vð Vð Vð
Pð Pð Pð Pð/%ð /%ð /%ð /%ðdLð Lð Lð Lð Lñ Lô Lð Lð Lð Lr   