a
    d;                  
   @   s  d Z ddlZddlmZ ddlmZmZmZ ddlZddl	m
Z
 ddlmZ ddlmZ ddlmZ ddlmZ dd	lmZ dd
lmZ ddlmZ ddlmZ ddlmZ ddlmZ ddlm Z  ddl!m"Z" ddl#m$Z$ ddl%m&Z& ddl'm(Z( ddl)m*Z* ddl+m,Z, e-e.Z/dhZ0d2e
j1dddZ2dd Z3d3dd Z4d4e
j1e5ee6 e5ee5 e7d$d%d&Z8d'd( Z9d5d*d+Z:d6e6ee5 e5e5ee7 e7ee5 ee d/d0d1Z;dS )7z\ Optimizer Factory w/ Custom Weight Decay
Hacked together by / Copyright 2021 Ross Wightman
    N)islice)OptionalCallableTuple)group_parameters   )	AdaBelief)	Adafactor)
Adahessian)AdamP)Adan)Lamb)Lars)Lion)	Lookahead)MADGRAD)Nadam)
NvNovoGrad)RAdam)	RMSpropTF)SGDPlionh㈵> )modelc                 C   sp   t |}g }g }|  D ]B\}}|js(q|jdksD|dsD||v rP|| q|| q|dd||dgS )Nr   z.bias        )paramsweight_decay)setnamed_parametersrequires_gradndimendswithappend)r   r   no_weight_decay_listdecayno_decaynameparamr   r   a/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/optim/optim_factory.pyparam_groups_weight_decay(   s    r*   c                    s   t   t  fdddS )Nc                      s   t t S N)tupler   r   itsizer   r)   <lambda>@       z_group.<locals>.<lambda>r   )iterr-   r   r-   r)   _group>   s    r3      c                    s   dd }t | di dd }g }g }|  D ]&\}}|||rH||n|| q,t|}	|d urp|	|   }tt||}t| dd t|D }
|
 fdd|D  |
S )Nc                    s:   |sdS t |ttfr,t fdd|D S  |S d S )NTc                    s   g | ]}  |qS r   )
startswith).0Zhpinr   r)   
<listcomp>H   r1   z0_layer_map.<locals>._in_head.<locals>.<listcomp>)
isinstancer,   listanyr5   )r8   hpr   r7   r)   _in_headD   s
    z_layer_map.<locals>._in_headZpretrained_cfg
classifierc                 S   s    i | ]\}}|D ]
}||qqS r   r   )r6   ilr8   r   r   r)   
<dictcomp>Y   r1   z_layer_map.<locals>.<dictcomp>c                    s   i | ]
}| qS r   r   )r6   r8   Znum_trunk_groupsr   r)   rB   Z   r1   )	getattrgetr   r#   lenr;   r3   	enumerateupdate)r   Zlayers_per_groupZ
num_groupsr>   Zhead_prefixZnames_trunkZ
names_headr8   _Znum_trunk_layers	layer_mapr   rC   r)   
_layer_mapC   s     rK   皙?      ?F)r   r   r$   layer_decayend_layer_decayverbosec                    sJ  t |}i }i }t| dr2t| | jdddd}nt| }t| d }	|	d t fddt|	D }
| 	 D ]\}}|j
sqv|jdks||v rd	}d
}nd}|}||}d||f }||vr|
| }||g d||< ||g d||< || d | || d | qv|r>ddl}td|j|dd  t| S )z
    Parameter groups for layer-wise lr decay & weight decay
    Based on BEiT: https://github.com/microsoft/unilm/blob/master/beit/optim_factory.py#L58
    group_matcherF)ZcoarseT)reverser   c                 3   s   | ]} |  V  qd S r+   r   )r6   r@   rN   Z	layer_maxr   r)   	<genexpr>v   r1   z+param_groups_layer_decay.<locals>.<genexpr>r&   r   r%   zlayer_%d_%s)lr_scaler   param_names)rU   r   r   rV   r   r   Nzparameter groups: 
%s   )indent)r   hasattrr   rQ   rK   maxvaluesr;   ranger   r    r!   rE   r#   json_loggerinfodumps)r   r   r$   rN   rO   rP   Zparam_group_namesZparam_groupsrJ   Z
num_layersZlayer_scalesr'   r(   Zg_decayZ
this_decayZlayer_idZ
group_nameZ
this_scaler]   r   rS   r)   param_groups_layer_decay^   sF    


ra   c                 C   s   t | j| j| j| jd}t| dddur2| j|d< t| dddurL| j|d< t| dddurf| j|d< t| dddur|	| j
 t| d	ddur| j|d
< |S )z cfg/argparse to kwargs helper
    Convert optimizer args in argparse args or cfg like object to keyword args for updated create fn.
    )optlrr   momentumopt_epsNeps	opt_betasbetasrN   opt_argsopt_foreachforeach)dictrb   rc   r   rd   rD   re   rg   rN   rH   ri   rj   )cfgkwargsr   r   r)   optimizer_kwargs   s"    



ro   Tc                 C   s   t |fi t| dd|iS )zk Legacy optimizer factory for backwards compatibility.
    NOTE: Use create_optimizer_v2 for new code.
    )rm   filter_bias_and_bn)create_optimizer_v2ro   )argsr   rp   r   r   r)   create_optimizer   s    rs   sgdr   ?)rb   rc   r   rd   rk   rp   rN   param_group_fnc	                 K   s  t | tjrri }
t| dr"|  }
|r0|| }qv|durNt| |||
d}d}qv|rh|rht| ||
}d}qv|  }n| }| }|	d}|d }|
drz dd	lm}m}m}m} d
}W n ty   d}Y n0 |rtj sJ d|
dr8zddl}d
}W n ty   d}Y n0 |r0tj s8J dtf d|i|	}|dur`|d| |du r|tv r|dd
 n||d< |dks|dkr|dd tj|f|d
d|}n|dkr|dd tj|f|dd|}n|dkrt|f|d
d|}n|dkr<tj|fi |}n|dkr\tj|fi |}nr|dkrt|fdd
d|}nN|dkrztj|fi |}W n$ t y   t|fi |}Y n0 n|dkrt!|fi |}n|dkrtj"|fi |}n|d kr*t#|fd!di|}n|d"krLt#|fd!d
i|}n|d#krltj$|fi |}nb|d$kr|dd% tj%|fi |}n6|d&krt&|fi |}n|d'krt'|fd(di|}n|d)krt'|fd(d
i|}n|d*krt(|fi |}n|d+kr:t(|fd,d
i|}n|d-kr^t)|f|d
d.|}np|d/krt)|fd|i|}nN|d0krt)|f|d
d
d1|}n(|d2krt)|f|d
d|}n|d3krt*|fd|i|}n|d4krt*|f|d
d5|}n|d6ks$|d7kr8t+|fi |}n|d8kr^tj,|fd9|d:|}np|d;krt-|fd9|d:|}nL|d<krt.|fi |}n.|d=krt/|fi |}n|d>kr|dd ||f|d
d|}n|d?kr|dd ||f|dd|}n|d@kr@||fdAdi|}n|dBkrb||fdAd
i|}nl|dCkr||fi |}nN|dDkr|dEdF ||fi |}n$|dGkr|dd |jj|f|d
d|}n|dHkr|dd |jj0|f|d
d|}n|dIkrD|dd |jj|fd|i|}n|dJkrv|dd |jj0|fd|i|}nX|dKkr|jj|fi |}n6|dLkr|jj1|fi |}n|dMkr|jj|fi |}n|dNkr|jj2|fi |}n|dOkr|jj3|fi |}n|dPkr:|jj4|fi |}n|dQkrZ|jj5|fi |}nt|dRkrz|jj4|fi |}nT|dSkr|jj.|fi |}n4|dTkr|jj6|fi |}ndrdUsJ t7t8|dVkr|d dWkrt9|}|S )Xa   Create an optimizer.

    TODO currently the model is passed in and all parameters are selected for optimization.
    For more general use an interface that allows selection of parameters to optimize and lr groups, one of:
      * a filter fn interface that further breaks params into groups in a weight_decay compatible fashion
      * expose the parameters interface and leave it up to caller

    Args:
        model_or_params (nn.Module): model containing parameters to optimize
        opt: name of optimizer to create
        lr: initial learning rate
        weight_decay: weight decay to apply in optimizer
        momentum:  momentum for momentum based optimizers (others may use betas via kwargs)
        foreach: Enable / disable foreach (multi-tensor) operation if True / False. Choose safe default if None
        filter_bias_and_bn:  filter out bias, bn and other 1d params from weight decay
        **kwargs: extra optimizer specific kwargs to pass through

    Returns:
        Optimizer
    no_weight_decayN)r   rN   r$   r   rI   Zfusedr   )FusedNovoGrad	FusedAdam	FusedLAMBFusedSGDTFz+APEX and CUDA required for fused optimizersbnbz1bitsandbytes and CUDA required for bnb optimizersr   rc   rk   rt   nesterovrf   )rd   r~   rd   sgdpZadamZadamwadampg{Gz?)Zwd_ratior~   nadamradamZadamax	adabeliefZrectifyZ
radabeliefZadadeltaZadagradg:0yE>	adafactorZadanpZno_proxZadanwlambZlambc
trust_clipZlarc)rd   r   larsZnlarc)rd   r   r~   ZnlarsmadgradZmadgradw)rd   Zdecoupled_decayZnovograd
nvnovogradZrmspropru   )alphard   Z	rmsproptfr   
adahessianZfusedsgdZfusedmomentumZ	fusedadamZadam_w_modeZ
fusedadamwZ	fusedlambZfusednovogradrh   )gffffff?g\(\?ZbnbsgdZ
bnbsgd8bitZbnbmomentumZbnbmomentum8bitZbnbadamZbnbadam8bitZbnbadamwZbnbadamw8bitZbnblambZbnblamb8bitZbnblarsZbnblarsb8bitZbnblionZbnblion8bitzInvalid optimizerr   	lookahead):r:   nnModulerY   rw   ra   r*   
parameterslowersplitr5   Zapex.optimizersry   rz   r{   r|   ImportErrortorchZcudaZis_availableZbitsandbytesrl   
setdefault_DEFAULT_FOREACHpopoptimZSGDr   ZAdamZAdamWr   r   AttributeErrorr   ZAdamaxr   ZAdadeltaZAdagradr	   r   r   r   r   r   ZRMSpropr   r   r
   ZSGD8bitZAdam8bitZ	AdamW8bitZLAMBZLAMB8bitZLARSZLion8bit
ValueErrorrF   r   )Zmodel_or_paramsrb   rc   r   rd   rk   rp   rN   rv   rn   rw   r   Z	opt_lowerZ	opt_splitry   rz   r{   r|   Zhas_apexr}   Zhas_bnbri   Z	optimizerr   r   r)   rq      s@    
























































rq   )r   r   )r4   N)rL   r   rM   NF)T)rt   Nr   ru   NTNN)<__doc__logging	itertoolsr   typingr   r   r   r   Ztorch.nnr   Ztorch.optimr   Ztimm.modelsr   r   r   r   r	   r   r
   r   r   Zadanr   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   Z
rmsprop_tfr   r   r   	getLogger__name__r^   r   r   r*   r3   rK   floatstrboolra   ro   rs   rq   r   r   r   r)   <module>   s   
  
     @
        