
    ^jG              ]          d Z ddlZddlmZ ddlmZ ddlmZmZm	Z	m
Z
mZmZ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mZ dd
lmZ ddlmZ  G d d      Zddddddddddddddddddddddddddeedddddddddddej:                   ej<                  d      dddf,deeeeef   ef      dee
e      dee
e       deed f   d!ed"ed#e!d$ee   d%e!d&e d'e"d(ed)e!d*ee"   d+eee e f      d,eee e f      d-e d.e d/e d0ee    d1e d2e d3ee"   d4ed5ed6e"d7ee d f   d8ee d f   d9ee    d:ee"   d;ee   d<ed=e!d>ed?ed@edAedBe!dCe!dDejF                  dEee"ej<                  f   dFe!dGe"dHe!dIeejH                  jJ                  jL                  ef   fZdJZ'y)Ka  NaFlex data loader for dynamic sequence length training.

This module provides a specialized data loader for Vision Transformer models that supports:
- Dynamic sequence length sampling during training for improved efficiency
- Variable patch size training with probabilistic selection
- Patch-level random erasing augmentation
- Efficient GPU prefetching with normalization

Hacked together by / Copyright 2025, Ross Wightman, Hugging Face
    N)suppress)partial)CallableDictIteratorListOptionalTupleUnion   )IMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)_worker_initadapt_to_chs)NaFlexMapDatasetWrapperNaFlexCollator)PatchRandomErasing)create_transformc                   v   e Zd ZdZeed ej                  d      dddddd	f
d
ej                  j                  j                  deedf   deedf   dedej                  deej                     dedededededdfdZdeeeeej*                  f   ej*                  f      fdZdefdZed        Zed        Zy)NaFlexPrefetchLoaderz;Data prefetcher for NaFlex format which normalizes patches.   cudaN        constr   r   Tloadermean.stdchannelsdevice	img_dtypere_probre_modere_countre_num_splitspatchify_channels_lastreturnc                    || _         || _        |xs t        j                  | _        t        ||      }t        ||      }|| _        || _        |rdd|fnd|df}t        j                  |D cg c]  }|dz  	 c}|| j                        j                  |      | _
        t        j                  |D cg c]  }|dz  	 c}|| j                        j                  |      | _        |dkD  rt        |||	|
|      | _        nd| _        |j                  dk(  xr t        j                  j!                         | _        |j                  dk(  xr t        j$                  j!                         | _        yc c}w c c}w )	a  Initialize NaFlexPrefetchLoader.

        Args:
            loader: DataLoader to prefetch from.
            mean: Mean values for normalization.
            std: Standard deviation values for normalization.
            channels: Number of image channels.
            device: Device to move tensors to.
            img_dtype: Data type for image tensors.
            re_prob: Random erasing probability.
            re_mode: Random erasing mode.
            re_count: Maximum number of erasing rectangles.
            re_num_splits: Number of augmentation splits.
            patchify_channels_last: Per-patch flat layout produced by the
                upstream Patchify. ``True`` (default): channel index is the
                innermost (P-P-C). ``False``: channel index is the outermost
                (C-P-P). Controls how ``patches`` tensors are viewed for
                normalization and erasing.
        r      )r   dtyper   )
erase_probmode	max_count
num_splitsr   Nr   npu)r   r   torchfloat32r    r   r   r%   tensorviewr   r   r   random_erasingtyper   is_availableis_cudar.   is_npu)selfr   r   r   r   r   r    r!   r"   r#   r$   r%   normalization_shapexs                 b/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/data/naflex_loader.py__init__zNaFlexPrefetchLoader.__init__   sU   B "3emm D(+3) &<#2Hq!X.qRZ\]N^LL"#QW#F$..JJN$ObJc 		<<!"QW"6IIMNaIb 	 R<"4""(#D #'D {{f,J1H1H1JkkU*Guyy/E/E/G# $"s   -E34E8c           
   #   p  K   d}| j                   rPt        j                  j                  | j                        }t        t        j                  j                  |      }nd| j                  rPt        j                  j                  | j                        }t        t        j                  j                  |      }nd}t        }| j                  D ]N  \  }} |       5  |j                         D ]W  \  }}t        |t        j                        s!|dk(  r| j                  nd}||   j                  | j                  d|      ||<   Y |j                  | j                  d      }|d   }	|	j                   }
|	j"                  dk(  rP|
\  }}}| j$                  r|	j'                  ||d	| j(                        }n|	j'                  ||| j(                  d	      }n|	j"                  d
k(  r|
dd \  }}| j$                  rK|
d	   | j(                  k(  sJ d| j(                   d|
d	           |	j'                  ||d	| j(                        }nd|
d   | j(                  k(  sJ d| j(                   d|
d           |	j'                  ||| j(                  d	      }nt+        d|	j"                   d      |j-                  | j.                        j1                  | j2                        }| j4                  | j$                  s |j7                  dd	      j9                         }| j5                  ||d   |j;                  dd            }| j$                  s |j7                  dd	      j9                         }|j'                  |
      |d<   ddd       |sf nd}|| j                   r:t        j                  j=                  | j                        j?                  |       nE| j                  r9t        j                  j=                  | j                        j?                  |       |}|}Q f y# 1 sw Y   xY ww)zIterate through the loader with prefetching and normalization.

        Yields:
            Tuple of (input_dict, targets) with normalized patches.
        T)r   )streamNpatches)r   non_blockingr)   )r   r@   r         z	Expected z channels, got z&Unexpected patches tensor dimensions: z. Expected 3 or 5.patch_coordpatch_valid)rE   rF   F) r6   r/   r   Streamr   r   r>   r7   r.   r   r   items
isinstanceTensorr    toshapendimr%   r2   r   
ValueErrorsubr   divr   r3   	transpose
contiguousgetcurrent_streamwait_stream)r8   firstr>   stream_contextnext_input_dictnext_targetkvr)   patches_tensororiginal_shape
batch_sizenum_patches_r?   
input_dicttargets                    r;   __iter__zNaFlexPrefetchLoader.__iter__`   s     <<ZZ&&dkk&:F$UZZ%6%6vFN[[YY%%T[[%9F$UYY%5%5fENF%N,0KK L	!(O[! =J+113 DAq!!U\\223y.d-<Q-?-B-B#';;)-"' .C .* *nnDKKdnS
 "1!;!/!5!5!&&!+1?.JQ22"0"5"5j+rSWS`S`"a #1"5"5j+t}}^`"a#((A-.<Ra.@+J22-b1T]]B ['onUWFXEYZ[B"0"5"5j+rSWS`S`"a  .a0DMMA Z'onUVFWEXYZA"0"5"5j+t}}^`"a$'MnNaNaMbbt%uvv "++dii044TXX>&&2  66")"3"3B";"F"F"H"11$3M$B$3$7$7t$L 2 G
  66")"3"3B";"F"F"H .5\\.-I	*{=J~  &((!<<JJ--T[[-AMMfU[[II,,DKK,@LLVT(J FYL	!\ &  [=J =Js&   CP61P*I+P*<B.P6*P3	/P6c                 ,    t        | j                        S )zhGet length of underlying loader.

        Returns:
            Number of batches in the loader.
        )lenr   r8   s    r;   __len__zNaFlexPrefetchLoader.__len__   s     4;;    c                 .    | j                   j                  S )zrGet sampler from underlying loader.

        Returns:
            Sampler from the underlying DataLoader.
        )r   samplerrf   s    r;   rj   zNaFlexPrefetchLoader.sampler        {{"""rh   c                 .    | j                   j                  S )zrGet dataset from underlying loader.

        Returns:
            Dataset from the underlying DataLoader.
        )r   datasetrf   s    r;   rm   zNaFlexPrefetchLoader.dataset   rk   rh   )__name__
__module____qualname____doc__r   r   r/   r   utilsdata
DataLoaderr
   floatintr	   r)   strboolr<   r   r   rJ   rc   rg   propertyrj   rm    rh   r;   r   r      sY   E
 '<%9#/5<<#7/3"!"+/@HKK$$//@H s
#@H ucz"	@H
 @H LL@H  ,@H @H @H @H @H %)@H @HD_!(5c5<<.?)@%,,)N#OP _!B    # # # #rh   r   )      @  i  i   r}       Fr   r   g      ?g?bilinear   *   Tr   all
patch_sizepatch_size_choicespatch_size_choice_probstrain_seq_lens.max_seq_lenr^   is_trainingmixup_fnno_augr!   r"   r#   re_splittrain_crop_modescaleratiohflipvflipcolor_jittercolor_jitter_probgrayscale_probgaussian_blur_probauto_augmentnum_aug_repeatsnum_aug_splitsinterpolationr   r   crop_pct	crop_modecrop_border_pixelsnum_workersdistributedrank
world_sizeseedepochuse_prefetcher
pin_memoryr    r   persistent_workersworker_seedingr%   r&   c-                    |r|dk(  sJ d       t        t        fi ddd|	d|d|d|d	|d
|d|d|d|d|d|d|d|d|d|d|d|d|
d|d|d|&ddd|,}-t        |      }.||.z  }/t        | t        j
                  j                  j                        rJ d       t        | |-|||||/||$|!|"|#d|%|,      }0t        j
                  j                  j                  |0dd| d|'t        t        |+       |*!      }1|&rt        |1|||(|)|
|||,"	      }1|1S t        d||||&d||d|,#
      | _        t        |$      }2d}3|!r<t        | t        j
                  j                  j                        sdd%lm}4  |4|       }3t        j
                  j                  j                  | |d| |3|2|'d&      }1|&rt        |1|||(|)|,'      }1|1S )(ux
  Create a data loader with dynamic sequence length sampling for training.

    Args:
        dataset: Dataset to load from.
        patch_size: Single patch size to use.
        patch_size_choices: List of patch sizes for variable patch size training.
        patch_size_choice_probs: Probabilities for each patch size choice.
        train_seq_lens: Training sequence lengths for dynamic batching.
        max_seq_len: Fixed sequence length for validation.
        batch_size: Batch size for validation and max training sequence length.
        is_training: Whether this is for training (enables dynamic batching).
        mixup_fn: Optional mixup function.
        no_aug: Disable augmentation.
        re_prob: Random erasing probability.
        re_mode: Random erasing mode.
        re_count: Maximum number of erasing rectangles.
        re_split: Random erasing split flag.
        train_crop_mode: Training crop mode.
        scale: Scale range for random resize crop.
        ratio: Aspect ratio range for random resize crop.
        hflip: Horizontal flip probability.
        vflip: Vertical flip probability.
        color_jitter: Color jitter factor.
        color_jitter_prob: Color jitter probability.
        grayscale_prob: Grayscale conversion probability.
        gaussian_blur_prob: Gaussian blur probability.
        auto_augment: AutoAugment policy.
        num_aug_repeats: Number of augmentation repeats.
        num_aug_splits: Number of augmentation splits.
        interpolation: Interpolation method.
        mean: Normalization mean values.
        std: Normalization standard deviation values.
        crop_pct: Crop percentage for validation.
        crop_mode: Crop mode.
        crop_border_pixels: Crop border pixels.
        num_workers: Number of data loading workers.
        distributed: Whether using distributed training.
        rank: Process rank for distributed training.
        world_size: Total number of processes.
        seed: Random seed.
        epoch: Starting epoch.
        use_prefetcher: Whether to use prefetching.
        pin_memory: Whether to pin memory.
        img_dtype: Image data type.
        device: Device to move tensors to.
        persistent_workers: Whether to use persistent workers.
        worker_seeding: Worker seeding mode.
        patchify_channels_last: Per-patch flat layout. ``True`` (default):
            channel index varies fastest (NaFlex default). ``False``: C-P-P
            (channels first — Gemma4 / HF native layout). Forwarded to the
            upstream ``Patchify`` and to the prefetcher's normalization layout.

    Returns:
        DataLoader or NaFlexPrefetchLoader instance.
    r   z=Augmentation repeats not currently supported in NaFlex loaderr   Tr   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r!   r"   r#   r   naflexr%   Fz IterableDataset Wrapper is a WIP)transform_factoryr   r   r   seq_lensmax_tokens_per_batchr   r   r   r   r   shuffler   r%   N)r   )r^   r   r   rj   r   worker_init_fnr   )r   r   r    r   r!   r"   r#   r%   )
r   r   r   r   r   r   r   r   patchifyr%   )r   )OrderedDistributedSampler)r^   r   r   rj   
collate_fnr   	drop_last)r   r   r    r   r%   )r   r   maxrI   r/   rr   rs   IterableDatasetr   rt   r   r   	transformr   timm.data.distributed_samplerr   )5rm   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   r   r   r   r   r   r   r   r   r   r   r   r   r    r   r   r   r%   r   max_train_seq_lenr   naflex_datasetr   r   rj   r   s5                                                        r;   create_naflex_loaderr      s   R !#d%dd##

 
 ,	

 
 
 
 
 &
 0
 *
  2
 &
 (
 
  !
" #
$  %
&  2'
( )
* +
, -
. */
0 1
2 $:3
8  /),==gu{{//??@<<<50/!1$;#!5#!#9
& !!,,#!"<O1 - 	
 )#!'=
Fv M[ -')!##9
 $<
 z'5;;3C3C3S3STO/8G!!,,!#!! - 	
 )#'=F Mrh   )(rq   math
contextlibr   	functoolsr   typingr   r   r   r   r	   r
   r   r/   	constantsr   r   r   r   r   r   r   r   naflex_random_erasingr   transforms_factoryr   r   r0   r   rv   ru   rx   rw   r)   rr   rs   rt   r   rz   rh   r;   <module>r      s  	    I I I  B . C 5 0~# ~#F =A269=*D!'+)-/3/3!-1 "$&&* '"7!5$(#',0!#!&+75<<+?#'#'+_iU5c?C#789i %T#Y/i "*$u+!6	i
 c3hi i i i 8$i i i i i i  "#!i" eUl+,#i$ eUl+,%i& 'i( )i* +i, $E?-i. /i0 "1i2 sm3i4 5i6 7i8 9i: E3J;i< 5#:=i> 5/?i@ C=AiB %SMCiF GiH IiJ KiL MiN OiP QiR SiT UiV ;;WiX c5<<'(YiZ ![i\ ]i^ !%_i` 
u{{**,@@	Aairh   