
    ^(j1                        d dl mZ d dlZd dlZd dlmZ d dlZd dlZd dlZd dl	m
Z
 d dlZd dlmZmZ d dlmZ d"dZd Zd	 Zd#d
Zed$d       Zd Z ej,                         d        Z G d d      Z G d dej2                  j4                        Z G d dej2                  j4                        Zdddej:                  fdZddd ed      dej:                  fdZ dej:                  fdZ!dd ed      dej:                  fdZ"dej:                  fdZ# G d dejH                        Z% G d d       Z&ed%d!       Z'y)&    )contextmanagerN)Path)Image)nnoptim)datac                 b    | |   D cg c]  } ||j                  |             }}||iS c c}w )z4Apply passed in transforms for HuggingFace Datasets.)convert)examples	transform	image_keymodeimageimagess         A/Users/danicosta/Desktop/Flux2/ComfyUI/comfy/k_diffusion/utils.pyhf_datasets_augs_helperr      s<    :B9:MN:Mid+,:MFNv Os   ,c                     || j                   z
  }|dk  rt        d| j                    d| d      | dd|z  z      }|j                  j                  dk(  r|j	                         j                         S |S )zNAppends dimensions to the end of a tensor until it has target_dims dimensions.r   z
input has z dims but target_dims is z, which is less).Nmps)ndim
ValueErrordevicetypedetachclone)xtarget_dimsdims_to_appendexpandeds       r   append_dimsr       sz     166)N:affX-F{mSbcdd'N223H )1(<(<(E8??""$S8S    c                 B    t        d | j                         D              S )z7Returns the number of trainable parameters in a module.c              3   <   K   | ]  }|j                           y wr   )numel).0ps     r   	<genexpr>zn_params.<locals>.<genexpr>"   s     6"5Qqwwy"5s   )sum
parameters)modules    r   n_paramsr+       s    6&"3"3"5666r!   c                    t        |       } | j                  j                  dd       | j                         sSt        j
                  j                  |      5 }t        | d      5 }t        j                  ||       ddd       ddd       |Rt        j                  t        | d      j                               j                         }||k7  rt        d|  d| d      | S # 1 sw Y   gxY w# 1 sw Y   kxY w)	zLDownloads a file if it does not exist, optionally checking its SHA-256 hash.T)parentsexist_okwbNrbzhash of z (url: z) failed to validate)r   parentmkdirexistsurllibrequesturlopenopenshutilcopyfileobjhashlibsha256read	hexdigestOSError)pathurldigestresponseffile_digests         r   download_filerE   %   s    :DKKdT2;;=^^##C(Hd46F!x+ 7G(nnT$%5%:%:%<=GGI[ HTF'#6JKLLK 7G6F((s$   C.%C"<C."C+	'C..C7c              #   B  K   | j                         D cg c]  }|j                   }}	 | j                  |       t        | j                               D ]  \  }}||   |_         yc c}w # t        | j                               D ]  \  }}||   |_         w xY ww)zdA context manager that places a model into training mode and restores
    the previous mode on exit.N)modulestrainingtrain	enumerate)modelr   r*   modesis        r   
train_moderN   3   s      ,1==?;?V__?E;'kk$"5==?3IAv#AhFO 4	 < #5==?3IAv#AhFO 4s%   BA*BA/ 1B/-BBc                     t        | d      S )zfA context manager that places a model into evaluation mode and restores
    the previous mode on exit.F)rN   )rK   s    r   	eval_moderP   ?   s     eU##r!   c                 0   t        | j                               }t        |j                               }|j                         |j                         k(  sJ |j                         D ]-  \  }}||   j	                  |      j                  |d|z
         / t        | j                               }t        |j                               }|j                         |j                         k(  sJ |j                         D ]  \  }}	||   j                  |	        y)zIncorporates updated model parameters into an exponential moving averaged
    version of a model. It should be called after each optimizer step.   )alphaN)dictnamed_parameterskeysitemsmul_add_named_bufferscopy_)
rK   averaged_modeldecaymodel_paramsaveraged_paramsnameparammodel_buffersaveraged_buffersbufs
             r   
ema_updatere   E   s     ..01L>::<=O/"6"6"8888#))+e""5)..uAI.F , ,,./MN88:;#3#8#8#::::"((*	c$$S) +r!   c                   4    e Zd ZdZ	 	 ddZd Zd Zd Zd Zy)		EMAWarmupaY  Implements an EMA warmup using an inverse decay schedule.
    If inv_gamma=1 and power=1, implements a simple average. inv_gamma=1, power=2/3 are
    good values for models you plan to train for a million or more steps (reaches decay
    factor 0.999 at 31.6K steps, 0.9999 at 1M steps), inv_gamma=1, power=3/4 for models
    you plan to train for less (reaches decay factor 0.999 at 10K steps, 0.9999 at
    215.4k steps).
    Args:
        inv_gamma (float): Inverse multiplicative factor of EMA warmup. Default: 1.
        power (float): Exponential factor of EMA warmup. Default: 1.
        min_value (float): The minimum EMA decay rate. Default: 0.
        max_value (float): The maximum EMA decay rate. Default: 1.
        start_at (int): The epoch to start averaging at. Default: 0.
        last_epoch (int): The index of last epoch. Default: 0.
    c                 X    || _         || _        || _        || _        || _        || _        y r   )	inv_gammapower	min_value	max_valuestart_at
last_epoch)selfri   rj   rk   rl   rm   rn   s          r   __init__zEMAWarmup.__init__h   s,    "
"" $r!   c                 H    t        | j                  j                               S )z2Returns the state of the class as a :class:`dict`.)rT   __dict__rW   ro   s    r   
state_dictzEMAWarmup.state_dictq   s    DMM'')**r!   c                 :    | j                   j                  |       y)zLoads the class's state.
        Args:
            state_dict (dict): scaler state. Should be an object returned
                from a call to :meth:`state_dict`.
        N)rr   update)ro   rt   s     r   load_state_dictzEMAWarmup.load_state_dictu   s     	Z(r!   c                     t        d| j                  | j                  z
        }dd|| j                  z  z   | j                   z  z
  }|dk  rdS t        | j                  t        | j                  |            S )z Gets the current EMA decay rate.r   rR           )maxrn   rm   ri   rj   minrl   rk   )ro   epochvalues      r   	get_valuezEMAWarmup.get_value}   se    At67Q//TZZK??QYrSCDNNE8R$SSr!   c                 .    | xj                   dz  c_         y)zUpdates the step count.rR   N)rn   rs   s    r   stepzEMAWarmup.step   s    1r!   N)      ?r   ry   r   r   r   )	__name__
__module____qualname____doc__rp   rt   rw   r~   r    r!   r   rg   rg   X   s+     UV%+)Tr!   rg   c                   4     e Zd ZdZ	 	 d fd	Zd Zd Z xZS )	InverseLRaM  Implements an inverse decay learning rate schedule with an optional exponential
    warmup. When last_epoch=-1, sets initial lr as lr.
    inv_gamma is the number of steps/epochs required for the learning rate to decay to
    (1 / 2)**power of its original value.
    Args:
        optimizer (Optimizer): Wrapped optimizer.
        inv_gamma (float): Inverse multiplicative factor of learning rate decay. Default: 1.
        power (float): Exponential factor of learning rate decay. Default: 1.
        warmup (float): Exponential warmup factor (0 <= warmup < 1, 0 to disable)
            Default: 0.
        min_lr (float): The minimum learning rate. Default: 0.
        last_epoch (int): The index of last epoch. Default: -1.
        verbose (bool): If ``True``, prints a message to stdout for
            each update. Default: ``False``.
    c                     || _         || _        d|cxk  rdk  st        d       t        d      || _        || _        t
        |   |||       y Nry   rR   zInvalid value for warmup)ri   rj   r   warmupmin_lrsuperrp   )	ro   	optimizerri   rj   r   r   rn   verbose	__class__s	           r   rp   zInverseLR.__init__   Z    "
Va788  788J8r!   c                 d    | j                   st        j                  d       | j                         S NzTTo get the last learning rate computed by the scheduler, please use `get_last_lr()`._get_lr_called_within_stepwarningswarn_get_closed_form_lrrs   s    r   get_lrzInverseLR.get_lr   -    ..MM 8 9 ''))r!   c           	         d| j                   | j                  dz   z  z
  }d| j                  | j                  z  z   | j                   z  }| j                  D cg c]  }|t        | j                  ||z        z    c}S c c}w NrR   )r   rn   ri   rj   base_lrsrz   r   ro   r   lr_multbase_lrs       r   r   zInverseLR._get_closed_form_lr   s~    T[[T__q%899t77TZZKG#}}.,G T[['G*;<<,. 	. .s   #A>)r   r   ry   ry   Fr   r   r   r   rp   r   r   __classcell__r   s   @r   r   r      s!      MO(-9*.r!   r   c                   4     e Zd ZdZ	 	 d fd	Zd Zd Z xZS )ExponentialLRaE  Implements an exponential learning rate schedule with an optional exponential
    warmup. When last_epoch=-1, sets initial lr as lr. Decays the learning rate
    continuously by decay (default 0.5) every num_steps steps.
    Args:
        optimizer (Optimizer): Wrapped optimizer.
        num_steps (float): The number of steps to decay the learning rate by decay in.
        decay (float): The factor by which to decay the learning rate every num_steps
            steps. Default: 0.5.
        warmup (float): Exponential warmup factor (0 <= warmup < 1, 0 to disable)
            Default: 0.
        min_lr (float): The minimum learning rate. Default: 0.
        last_epoch (int): The index of last epoch. Default: -1.
        verbose (bool): If ``True``, prints a message to stdout for
            each update. Default: ``False``.
    c                     || _         || _        d|cxk  rdk  st        d       t        d      || _        || _        t
        |   |||       y r   )	num_stepsr]   r   r   r   r   rp   )	ro   r   r   r]   r   r   rn   r   r   s	           r   rp   zExponentialLR.__init__   r   r!   c                 d    | j                   st        j                  d       | j                         S r   r   rs   s    r   r   zExponentialLR.get_lr   r   r!   c           	         d| j                   | j                  dz   z  z
  }| j                  d| j                  z  z  | j                  z  }| j                  D cg c]  }|t        | j                  ||z        z    c}S c c}w r   )r   rn   r]   r   r   rz   r   r   s       r   r   z!ExponentialLR._get_closed_form_lr   s|    T[[T__q%899::!dnn"45$//I#}}.,G T[['G*;<<,. 	. .s   #A=)g      ?ry   ry   r   Fr   r   s   @r   r   r      s!      KM(-9*.r!   r   ry   r   cpuc                 Z    t        j                  | ||      |z  |z   j                         S )z-Draws samples from an lognormal distribution.r   dtype)torchrandnexp)shapelocscaler   r   s        r   rand_log_normalr      s(    KKfE:UBSHMMOOr!   infc                 ~   t        j                  ||t         j                        }t        j                  ||t         j                        }|j                         j	                  |      j                  |      j                         }|j                         j	                  |      j                  |      j                         }t        j                  | |t         j                        ||z
  z  |z   }	|	j                         j                  |      j                  |      j                         j                  |      S )zEDraws samples from an optionally truncated log-logistic distribution.r   )r   	as_tensorfloat64logsubdivsigmoidrandlogitmuladdr   to)
r   r   r   rk   rl   r   r   min_cdfmax_cdfus
             r   rand_log_logisticr      s    	&NI	&NImmo!!#&**5199;Gmmo!!#&**5199;G

5u}}=7ARSV]]A779==##C(,,.11%88r!   c                     t        j                  |      }t        j                  |      }t        j                  | ||      ||z
  z  |z   j	                         S )z/Draws samples from an log-uniform distribution.r   )mathr   r   r   r   )r   rk   rl   r   r   s        r   rand_log_uniformr      sJ    #I#IJJuV59Y=RSV__ddffr!   c                 L   t        j                  ||z        dz  t         j                  z  }t        j                  ||z        dz  t         j                  z  }t        j                  | ||      ||z
  z  |z   }t        j
                  |t         j                  z  dz        |z  S )zJDraws samples from a truncated v-diffusion training timestep distribution.   r   )r   atanpir   r   tan)	r   
sigma_datark   rl   r   r   r   r   r   s	            r   rand_v_diffusionr      s    ii	J./!3dgg=Gii	J./!3dgg=G

5u579JKgUA99Q[1_%
22r!   c                     t        j                  | ||      j                         }t        j                  | ||      }|| z  |z   }||z  |z   }	|||z   z  }
t        j                  ||
k  ||	      j                         S )z2Draws samples from a split lognormal distribution.r   )r   r   absr   wherer   )r   r   scale_1scale_2r   r   nr   n_leftn_rightratios              r   rand_split_log_normalr      s|    E&6::<A

5u5A'\CF'kCGw()E;;q5y&'26688r!   c                   >     e Zd ZdZh dZd fd	Zd Zd Zd Z xZ	S )FolderOfImageszURecursively finds all images in a directory. It does not support
    classes/targets.>	   .bmp.jpg.pgm.png.ppm.tif.jpeg.tiff.webpc                      t                    t        |       _        |t	        j
                         n| _        t         fd j                  j                  d      D               _	        y )Nc              3   p   K   | ]-  }|j                   j                         j                  v s*| / y wr   )suffixlowerIMG_EXTENSIONS)r%   r?   ro   s     r   r'   z*FolderOfImages.__init__.<locals>.<genexpr>  s/     p-ATT[[EVEVEX\`\o\oEoD-As   +66*)
r   rp   r   rootr   Identityr   sortedrglobpaths)ro   r   r   r   s   `  r   rp   zFolderOfImages.__init__  sL    J	*3*;pTYY__S-App
r!   c                 :    d| j                    dt        |        dS )NzFolderOfImages(root="z", len: ))r   lenrs   s    r   __repr__zFolderOfImages.__repr__  s    &tyyk#d)AFFr!   c                 ,    t        | j                        S r   )r   r   rs   s    r   __len__zFolderOfImages.__len__  s    4::r!   c                     | j                   |   }t        |d      5 }t        j                  |      j                  d      }d d d        | j	                        }|fS # 1 sw Y   xY w)Nr0   RGB)r   r7   r   r
   r   )ro   keyr?   rC   r   s        r   __getitem__zFolderOfImages.__getitem__  sV    zz#$JJqM))%0E u%v s   %AA&r   )
r   r   r   r   r   rp   r   r   r   r   r   s   @r   r   r     s&     aNqGr!   r   c                       e Zd Zd Zd Zy)	CSVLoggerc                    t        |      | _        || _        | j                  j                         rt	        | j                  d      | _        y t	        | j                  d      | _         | j                  | j                    y )Naw)r   filenamecolumnsr3   r7   filewrite)ro   r   r  s      r   rp   zCSVLogger.__init__  sZ    X==!T]]C0DIT]]C0DIDJJ%r!   c                 2    t        |d| j                  dd y )N,T)sepr  flush)printr  )ro   argss     r   r  zCSVLogger.write&  s    t499D9r!   N)r   r   r   rp   r  r   r!   r   r   r     s    &:r!   r   c              #     K   t         j                  j                  j                  }t         j                  j                  j
                  j                  }	 | | t         j                  j                  _        |)|t         j                  j                  j
                  _        d | |t         j                  j                  _        |*|t         j                  j                  j
                  _        yy# | |t         j                  j                  _        |*|t         j                  j                  j
                  _        w w xY ww)zGA context manager that sets whether TF32 is allowed on cuDNN or matmul.N)r   backendscudnn
allow_tf32cudamatmul)r  r  	cudnn_old
matmul_olds       r   	tf32_moder  *  s      $$//I$$++66J
?.3ENN  +4:ENN&&1.7ENN  +4>ENN&&1  .7ENN  +4>ENN&&1 s!   AEAC4 &AE4AEE)r   r   )T)NN)(
contextlibr   r:   r   pathlibr   r8   r4   r   PILr   r   r   r   torch.utilsr   r   r    r+   rE   rN   rP   no_gradre   rg   lr_scheduler_LRSchedulerr   r   float32r   floatr   r   r   r   Datasetr   r   r  r   r!   r   <module>r     sQ   %          T7
 ' '$ * *$- -`&.""// &.R&.E&&33 &.R  "E P
 "$2uU|\ainiviv 9 :?emm g (*R5<X]ejerer 3 @EEMM 9T\\ 4: : ? ?r!   