ó
    "Eñiš#  ã                   óþ   • S SK r S SKJr  S SKJrJrJr  S SKJrJ	r	  S SK
r
S SKJr  S SKJr  S SKJr  S SKJr  S S	KJrJr  S
S/r\	" SSS9r\" S5       " S S\\   5      5       rS r\" S5       " S S
\5      5       rg)é    N)Ú
namedtuple)ÚCallableÚIteratorÚSized)ÚAnyÚTypeVar)Údefault_collate)Úfunctional_datapipe)Údataframe_wrapper)ÚIterDataPipe)Ú_check_unpickable_fnÚvalidate_input_colÚCollatorIterDataPipeÚMapperIterDataPipeÚ_T_coT)Ú	covariantÚmapc                   ó‚   ^ • \ rS rSr% Sr\\S'   \\S'     SS\S\SS4U 4S jjjrS r	S\
\   4S	 jrS\4S
 jrSrU =r$ )r   é   ao  
Applies a function over each item from the source DataPipe (functional name: ``map``).

The function can be any regular Python function or partial object. Lambda
function is not recommended as it is not supported by pickle.

Args:
    datapipe: Source Iterable DataPipe
    fn: Function being applied over each item
    input_col: Index or indices of data which ``fn`` is applied, such as:

        - ``None`` as default to apply ``fn`` to the data directly.
        - Integer(s) is used for list/tuple.
        - Key(s) is used for dict.

    output_col: Index of data where result of ``fn`` is placed. ``output_col`` can be specified
        only when ``input_col`` is not ``None``

        - ``None`` as default to replace the index that ``input_col`` specified; For ``input_col`` with
          multiple indices, the left-most one is used, and other indices will be removed.
        - Integer is used for list/tuple. ``-1`` represents to append result at the end.
        - Key is used for dict. New key is acceptable.

Example:
    >>> # xdoctest: +SKIP
    >>> from torchdata.datapipes.iter import IterableWrapper, Mapper
    >>> def add_one(x):
    ...     return x + 1
    >>> dp = IterableWrapper(range(10))
    >>> # Invocation via functional form is preferred
    ... map_dp_1 = dp.map(add_one)
    >>> list(map_dp_1)
    [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]
    >>> # We discourage the usage of `lambda` functions as they are not serializable with `pickle`
    >>> # Use `functools.partial` or explicitly define the function instead
    >>> map_dp_2 = Mapper(dp, lambda x: x + 1)
    >>> list(map_dp_2)
    [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]
ÚdatapipeÚfnNÚreturnc                 óR  >• [         R                  R                  S5        [        TU ]  5         Xl        [        U5        X l        X0l        Uc  Ub  [        S5      e[        U[        [        45      (       a  [        U5      S:”  a  [        S5      eUS   nX@l        [        X#5        g )Nzpython.data_pipes.mapz3`output_col` must be None when `input_col` is None.é   z3`output_col` must be a single-element list or tupler   )ÚtorchÚ_CÚ_log_api_usage_onceÚsuperÚ__init__r   r   r   Ú	input_colÚ
ValueErrorÚ
isinstanceÚlistÚtupleÚlenÚ
output_colr   )Úselfr   r   r    r&   Ú	__class__s        €Úe/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/utils/data/datapipes/iter/callable.pyr   ÚMapperIterDataPipe.__init__H   s�   ø€ ô 	�‰×$Ñ$Ð%<Ô=Ü‰ÑÔØ Œä˜RÔ ØŒà"ŒØÑ Ñ!7ÜÐRÓSÐSÜ�j¤4¬ -×0Ñ0Ü�:‹ Ó"Ü Ð!VÓWÐWØ# A™ˆJØ$ŒÜ˜2Õ)ó    c                 ó<  ^• U R                   c  U R                  c  U R                  T5      $ U R                   c  U R                  T5      nOr[        U R                   [        [
        45      (       a/  [        U4S jU R                    5       5      nU R                  " U6 nOU R                  TU R                      5      n[        T[
        5      (       a  Sn[	        T5      mOSnU R                  ci  [        U R                   [        [
        45      (       a4  UTU R                   S   '   [        U R                   SS  SS9 H  nTU	 M     OAUTU R                   '   O1U R                  S:X  a  TR                  U5        OUTU R                  '   U(       a  [        T5      $ T$ )Nc              3   ó.   >#   • U  H
  nTU   v •  M     g 7f©N© )Ú.0ÚcolÚdatas     €r)   Ú	<genexpr>Ú/MapperIterDataPipe._apply_fn.<locals>.<genexpr>g   s   øé € Ð=ªn s˜˜cžªnùs   ƒTFr   r   )Úreverseéÿÿÿÿ)r    r&   r   r"   r#   r$   ÚsortedÚappend)r'   r2   ÚresÚargsÚt_flagÚidxs    `    r)   Ú	_apply_fnÚMapperIterDataPipe._apply_fn`   sI  ø€ Ø�>‰>Ñ! d§o¡oÑ&=Ø—7‘7˜4“=Ð à�>‰>Ñ!Ø—'‘'˜$“-‰CÜ˜Ÿ™¬¬u¨×6Ñ6ÜÔ=¨d¯nªnÓ=Ó=ˆDØ—'’'˜4�.‰Cà—'‘'˜$˜tŸ~™~Ñ.Ó/ˆCô �dœE×"Ñ"ØˆFÜ˜“:‰DàˆFà�?‰?Ñ"Ü˜$Ÿ.™.¬4´¨-×8Ñ8Ø*-��T—^‘^ AÑ&Ñ'Ü! $§.¡.°°Ð"4¸dÔC�CØ˜Sš	ò Dð (+��T—^‘^Ò$à�‰ "Ó$Ø—‘˜CÕ à(+��T—_‘_Ñ%ö %Œu�T‹{Ð.¨$Ð.r+   c              #   óX   #   • U R                    H  nU R                  U5      v •  M     g 7fr.   )r   r=   )r'   r2   s     r)   Ú__iter__ÚMapperIterDataPipe.__iter__ƒ   s"   é € Ø—M”MˆDØ—.‘. Ó&Ô&ò "ùs   ‚(*c                 ó¬   • [        U R                  [        5      (       a  [        U R                  5      $ [	        [        U 5      R                   S35      e)Nz# instance doesn't have valid length)r"   r   r   r%   Ú	TypeErrorÚtypeÚ__name__)r'   s    r)   Ú__len__ÚMapperIterDataPipe.__len__‡   s@   € ä�d—m‘m¤U×+Ñ+Ü�t—}‘}Ó%Ð%Üœ4 ›:×.Ñ.Ð/Ð/RÐSÓTÐTr+   )r   r   r    r&   )NN)rE   Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Ú__annotations__r   r   r=   r   r   r@   ÚintrF   Ú__static_attributes__Ú__classcell__©r(   s   @r)   r   r      sn   ø‡ ñ&ðP ÓØƒLð Øñ*àð*ð ð*ð 
÷*ð *ò0!/ðF'˜( 5™/ô 'ðU˜÷ Uò Ur+   c                 ó(  • [        UR                  5      S:”  a  [        S5      eUS   n[        R                  " U5      n/ n/ nU  H  nXc;  d  M
  [        S5      e   U H€  nX`;   a"  [        X   5      (       d  [        S5      eX   nO! SS KJn  UR                  R                  5       nUR                  [        U5      5        U" X&   5      n
UR                  U
5        M‚     [        SU5      nU" U6 nU$ ! [         a  n	[        S5      U	eS n	A	ff = f)Nr   z%Only supports one DataFrame per batchr   zConversion keys mismatchz5Collate (DF)DataPipe requires callable as dict valuesz?unable to import default collation function from the TorchArrowÚCollateResult)r%   ÚitemsÚRuntimeErrorÚ
df_wrapperÚget_columnsÚcallableÚtorcharrow.pytorchÚpytorchÚrecÚDefaultÚ	Exceptionr8   Ústrr   )Ú
conversionÚitemÚdfÚcolumns_nameÚtuple_namesÚtuple_valuesÚnameÚcollation_fnÚtapÚeÚvalueÚtpl_clsr$   s                r)   Ú_collate_helperrj   Ž   s  € ä
ˆ4�:‰:ƒ˜ÓäÐBÓCÐCØ	ˆa‰€BÜ×)Ò)¨"Ó-€LØ€KØ€LãˆØÕ#ÜÐ9Ó:Ð:ñ ó ˆØÓÜ˜JÑ,×-Ñ-Ü"ØKóð ð &Ñ+‰LðÝ0à"Ÿw™wŸ™Ó0�ð 	×Ñœ3˜t›9Ô%Ù˜R™XÓ&ˆØ×Ñ˜EÖ"ñ) ô0 ˜¨+Ó6€GÙ�\Ð"€EØ€Løô ó Ü"ØUóàðûðús   Â
 C6Ã6
DÄ DÄDÚcollatec            	       óz   ^ • \ rS rSrSr\S4S\S\S\4   \	\
\-  \\-  4   -  S-  S\S-  SS4U 4S	 jjjrS
rU =r$ )r   é¹   aÅ  
Collates samples from DataPipe to Tensor(s) by a custom collate function (functional name: ``collate``).

By default, it uses :func:`torch.utils.data.default_collate`.

.. note::
    While writing a custom collate function, you can import :func:`torch.utils.data.default_collate` for the
    default behavior and `functools.partial` to specify any additional arguments.

Args:
    datapipe: Iterable DataPipe being collated
    collate_fn: Customized collate function to collect and combine data or a batch of data.
        Default function collates to Tensor(s) based on data type.

Example:
    >>> # xdoctest: +SKIP
    >>> # Convert integer data to float Tensor
    >>> class MyIterDataPipe(torch.utils.data.IterDataPipe):
    ...     def __init__(self, start, end):
    ...         super(MyIterDataPipe).__init__()
    ...         assert end > start, "this example only works with end >= start"
    ...         self.start = start
    ...         self.end = end
    ...
    ...     def __iter__(self):
    ...         return iter(range(self.start, self.end))
    ...
    ...     def __len__(self):
    ...         return self.end - self.start
    >>> ds = MyIterDataPipe(start=3, end=7)
    >>> print(list(ds))
    [3, 4, 5, 6]
    >>> def collate_fn(batch):
    ...     return torch.tensor(batch, dtype=torch.float)
    >>> collated_ds = CollateIterDataPipe(ds, collate_fn=collate_fn)
    >>> print(list(collated_ds))
    [tensor(3.), tensor(4.), tensor(5.), tensor(6.)]
Nr   r^   .Ú
collate_fnr   c                 ó´   >• Ub  [         TU ]  XS9  g [        U5      (       a  [         TU ]  XS9  g [        R                  " [
        U5      n[         TU ]  XS9  g )N)r   )r   r   rW   Ú	functoolsÚpartialrj   )r'   r   r^   rn   r(   s       €r)   r   ÚCollatorIterDataPipe.__init__â   sZ   ø€ ð Ñ!Ü‰GÑ˜XÐÒ5ä˜
×#Ñ#Ü‘Ñ  Ð Ò9ô '×.Ò.¬À
ÓK�
Ü‘Ñ  Ð Ò9r+   r/   )rE   rH   rI   rJ   rK   r	   r   r   r   Údictr]   r   rN   rO   rP   s   @r)   r   r   ¹   sp   ø† ñ%ðX !Ø&*ñ:àð:ð ˜S #˜XÑ&Ø
ˆs�S‰y˜( S™.Ð(Ñ
)ñ*à
ñð:ð ˜t‘Oð:ð 
÷:ö :r+   )rp   Úcollectionsr   Úcollections.abcr   r   r   Útypingr   r   r   Útorch.utils.data._utils.collater	   Ú%torch.utils.data.datapipes._decoratorr
   Ú$torch.utils.data.datapipes.dataframer   rU   Ú#torch.utils.data.datapipes.datapiper   Ú'torch.utils.data.datapipes.utils.commonr   r   Ú__all__r   r   rj   r   r/   r+   r)   Ú<module>r}      s˜   ðã Ý "ß 5Ñ 5ß ã Ý ;Ý EÝ PÝ <÷ð Øð€ñ 	� 4Ñ(€ñ �UÓôoU˜ eÑ,ó oUó ðoUòd(ñV �YÓô::Ð-ó ::ó  ñ::r+   