a
    dH                     @   s   d Z ddlZddlmZ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 ddlmZ ddlZddlmZ ddlmZ dd	lmZ eeZeed
ddZ G dd deZ!G dd de!Z"dS )z]
Finetuning Callback
^^^^^^^^^^^^^^^^^^^^
Freeze and unfreeze models for finetuning purposes
    N)AnyCallableDict	GeneratorIterableListOptionalUnion)Module
ModuleDict)
_BatchNorm)	Optimizer)Callback)MisconfigurationException)rank_zero_warn)epochreturnc                 C   s   dS )Ng       @ )r   r   r   o/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/callbacks/finetuning.pymultiplicative#   s    r   c                	   @   s  e Zd ZdZddddZeeef dddZeeef ddd	d
Z	ddddddZ
eeeeeeef  f ee dddZed7eeeeeef  f eeedddZeeeeeeef  f ddddZeeddddZed8eeeeeef  f eddddZeeeedd d!Zed9eeeeeef  f eee eedd#d$d%Zddedd&d'd(Zeeeeef  eeeeef  d)d*d+Zdeeeeeef  dd,d-d.Z ddddd/d0Z!deeedd1d2d3Z"ddd4d5d6Z#dS ):BaseFinetuninga  
    This class implements the base logic for writing your own Finetuning Callback.

    Override ``freeze_before_training`` and ``finetune_function`` methods with your own logic.

    ``freeze_before_training``: This method is called before ``configure_optimizers``
        and should be used to freeze any modules parameters.

    ``finetune_function``: This method is called on every train epoch start and should be used to
        ``unfreeze`` any parameters. Those parameters needs to be added in a new ``param_group``
        within the optimizer.

    .. note:: Make sure to filter the parameters based on ``requires_grad``.

    Example::

        >>> from torch.optim import Adam
        >>> class MyModel(pl.LightningModule):
        ...     def configure_optimizer(self):
        ...         # Make sure to filter the parameters based on `requires_grad`
        ...         return Adam(filter(lambda p: p.requires_grad, self.parameters()))
        ...
        >>> class FeatureExtractorFreezeUnfreeze(BaseFinetuning):
        ...     def __init__(self, unfreeze_at_epoch=10):
        ...         super().__init__()
        ...         self._unfreeze_at_epoch = unfreeze_at_epoch
        ...
        ...     def freeze_before_training(self, pl_module):
        ...         # freeze any module you want
        ...         # Here, we are freezing `feature_extractor`
        ...         self.freeze(pl_module.feature_extractor)
        ...
        ...     def finetune_function(self, pl_module, current_epoch, optimizer, optimizer_idx):
        ...         # When `current_epoch` is 10, feature_extractor will start training.
        ...         if current_epoch == self._unfreeze_at_epoch:
        ...             self.unfreeze_and_add_param_group(
        ...                 modules=pl_module.feature_extractor,
        ...                 optimizer=optimizer,
        ...                 train_bn=True,
        ...             )
    Nr   c                 C   s   i | _ d| _d S NF)_internal_optimizer_metadata_restartingselfr   r   r   __init__R   s    zBaseFinetuning.__init__c                 C   s
   d| j iS )Ninternal_optimizer_metadata)r   r   r   r   r   
state_dictV   s    zBaseFinetuning.state_dictr   r   c                 C   s$   d| _ d|v r|d | _n|| _d S )NTr   )r   r   r   r   r   r   r   load_state_dict[   s    zBaseFinetuning.load_state_dict
pl.Trainerpl.LightningModuletrainer	pl_moduler   c                 C   sH   | j rDt| }t|jD ] \}}| | j| |}||_qd| _ d S r   )r   dictnamed_parameters	enumerate
optimizers_apply_mapping_to_param_groupsr   param_groups)r   r&   r'   r)   opt_idx	optimizerr-   r   r   r   on_fit_startc   s    
zBaseFinetuning.on_fit_start)modulesr   c                 C   sZ   t | tr|  } t | trDg }| D ]}|t| q$t|}n|  }dd |D S )aG  This function is used to flatten a module or an iterable of modules into a list of its leaf modules
        (modules with no children) and parent modules that have parameters directly themselves.

        Args:
            modules: A given module or an iterable of modules

        Returns:
            List of modules
        c                 S   s"   g | ]}t | r|jr|qS r   )listchildren_parameters).0mr   r   r   
<listcomp>       z2BaseFinetuning.flatten_modules.<locals>.<listcomp>)	
isinstancer   valuesr   extendr   flatten_modulesiterr1   )r1   Z_flatten_modulesr6   Z_modulesr   r   r   r<   n   s    


zBaseFinetuning.flatten_modulesT)r1   train_bnrequires_gradr   c                 c   sJ   t | } | D ]6}t|tr"|s"q|jddD ]}|j|kr.|V  q.qdS )al  Yields the `requires_grad` parameters of a given module or list of modules.

        Args:
            modules: A given module or an iterable of modules
            train_bn: Whether not to train the BatchNorm module
            requires_grad: Whether to create a generator for trainable or non-trainable parameters.
        Returns:
            Generator
        FZrecurseN)r   r<   r9   r   
parametersr?   )r1   r>   r?   modparamr   r   r   filter_params   s    

zBaseFinetuning.filter_paramsc                 C   s@   t | } | D ],}t|tr"d|_|jddD ]
}d|_q.qdS )zUnfreezes the parameters of the provided modules.

        Args:
            modules: A given module or an iterable of modules
        TFr@   N)r   r<   r9   r   track_running_statsrA   r?   )r1   modulerC   r   r   r   make_trainable   s    

zBaseFinetuning.make_trainable)rF   r   c                 C   s,   t | trd| _| jddD ]
}d|_qdS )ziFreezes the parameters of the provided module.

        Args:
            module: A given module
        Fr@   N)r9   r   rE   rA   r?   )rF   rC   r   r   r   freeze_module   s    
zBaseFinetuning.freeze_module)r1   r>   r   c                 C   s<   t | } | D ](}t|tr,|r,t | qt | qdS )zFreezes the parameters of the provided modules.

        Args:
            modules: A given module or an iterable of modules
            train_bn: If True, leave the BatchNorm layers in training mode

        Returns:
            None
        N)r   r<   r9   r   rG   rH   )r1   r>   rB   r   r   r   freeze   s
    
zBaseFinetuning.freeze)r/   paramsr   c                    s\   g }g }|D ]2 t  fdd| jD s4|  q|  q|rXtdt|  d |S )ab  This function is used to exclude any parameter which already exists in this optimizer.

        Args:
            optimizer: Optimizer used for parameter exclusion
            params: Iterable of parameters used to check against the provided optimizer

        Returns:
            List of parameters not contained in this optimizer param groups
        c                 3   s(   | ] }|d  D ]}t | V  qqdS )rJ   N)torchequal)r5   groupprC   r   r   	<genexpr>   r8   z5BaseFinetuning.filter_on_optimizer.<locals>.<genexpr>zThe provided params to be frozen already exist within another group of this optimizer. Those parameters will be skipped.
HINT: Did you init your optimizer in `configure_optimizer` as such:
 z<(filter(lambda p: p.requires_grad, self.parameters()), ...) )anyr-   appendr   type)r/   rJ   Z
out_paramsZremoved_paramsr   rO   r   filter_on_optimizer   s    z"BaseFinetuning.filter_on_optimizer      $@)r1   r/   lrinitial_denom_lrr>   r   c                 C   sp   t |  |du r |jd d nt|}|du r4|nd}t j| |dd}t ||}|rl|||| d dS )a  Unfreezes a module and adds its parameters to an optimizer.

        Args:
            modules: A module or iterable of modules to unfreeze.
                Their parameters will be added to an optimizer as a new param group.
            optimizer: The provided optimizer will receive new parameters and will add them to
                `add_param_group`
            lr: Learning rate for the new param group.
            initial_denom_lr: If no lr is provided, the learning from the first param group will be used
                and divided by `initial_denom_lr`.
            train_bn: Whether to train the BatchNormalization layers.
        Nr   rV   g      ?T)r>   r?   )rJ   rV   )r   rG   r-   floatrD   rT   Zadd_param_group)r1   r/   rV   rW   r>   Z	params_lrZdenom_lrrJ   r   r   r   unfreeze_and_add_param_group   s    
z+BaseFinetuning.unfreeze_and_add_param_group)r&   r'   stager   c                 C   s   |  | d S N)freeze_before_training)r   r&   r'   rZ   r   r   r   setup  s    zBaseFinetuning.setup)r-   mappingr   c                    sH   g }| D ]:}dd |  D } fdd|d D |d< || q|S )Nc                 S   s   i | ]\}}|d kr||qS )rJ   r   )r5   kvr   r   r   
<dictcomp>  r8   zABaseFinetuning._apply_mapping_to_param_groups.<locals>.<dictcomp>c                    s   g | ]} | qS r   r   )r5   rN   r^   r   r   r7     r8   zABaseFinetuning._apply_mapping_to_param_groups.<locals>.<listcomp>rJ   )itemsrR   )r-   r^   outputgZgroup_stater   rb   r   r,     s    z-BaseFinetuning._apply_mapping_to_param_groups)r'   r.   num_param_groupscurrent_param_groupsr   c                 C   s`   dd |  D }|| jvr0| ||| j|< n,|t|kr\| j| | ||d  | d S )Nc                 S   s   i | ]\}}||qS r   r   )r5   nrN   r   r   r   ra     r8   z)BaseFinetuning._store.<locals>.<dictcomp>)r)   r   r,   lenr;   )r   r'   r.   rf   rg   r^   r   r   r   _store  s    

zBaseFinetuning._storec                 C   s\   ddl m} ||j|jdD ]:\}}t|j}| ||j|| |j}| |||| qdS )Called when the epoch begins.r   )_get_active_optimizersN)	Z!pytorch_lightning.loops.utilitiesrl   r+   Zoptimizer_frequenciesri   r-   finetune_functionZcurrent_epochrj   )r   r&   r'   rl   r.   r/   rf   rg   r   r   r   on_train_epoch_start#  s    
z#BaseFinetuning.on_train_epoch_startr'   r   r/   r.   r   c                 C   s   t dS )z$Override to add your unfreeze logic.NNotImplementedError)r   r'   r   r/   r.   r   r   r   rm   .  s    z BaseFinetuning.finetune_functionr'   r   c                 C   s   t dS )z"Override to add your freeze logic.Nrp   r   r'   r   r   r   r\   4  s    z%BaseFinetuning.freeze_before_training)TT)T)NrU   T)$__name__
__module____qualname____doc__r   r   strr   r   r"   r0   staticmethodr	   r
   r   r   r<   boolr   rD   rG   rH   rI   r   rT   r   rX   rY   r]   r(   r,   intrj   rn   rm   r\   r   r   r   r   r   '   s^   ** &*   ,r   c                       s   e Zd ZdZdedddddddf	eeeee e	ee	e	edd	
 fd
dZ
eeef dddZeeef dd fddZdddd fddZdddddZdeeeddddZ  ZS )BackboneFinetuninga  Finetune a backbone model based on a learning rate user-defined scheduling.

    When the backbone learning rate reaches the current model learning rate
    and ``should_align`` is set to True, it will align with it for the rest of the training.

    Args:
        unfreeze_backbone_at_epoch: Epoch at which the backbone will be unfreezed.
        lambda_func: Scheduling function for increasing backbone learning rate.
        backbone_initial_ratio_lr:
            Used to scale down the backbone learning rate compared to rest of model
        backbone_initial_lr: Optional, Initial learning rate for the backbone.
            By default, we will use ``current_learning /  backbone_initial_ratio_lr``
        should_align: Whether to align with current learning rate when backbone learning
            reaches it.
        initial_denom_lr: When unfreezing the backbone, the initial learning rate will
            ``current_learning_rate /  initial_denom_lr``.
        train_bn: Whether to make Batch Normalization trainable.
        verbose: Display current learning rate for model and backbone
        rounding: Precision for displaying learning rate

    Example::

        >>> from pytorch_lightning import Trainer
        >>> from pytorch_lightning.callbacks import BackboneFinetuning
        >>> multiplicative = lambda epoch: 1.5
        >>> backbone_finetuning = BackboneFinetuning(200, multiplicative)
        >>> trainer = Trainer(callbacks=[backbone_finetuning])

    
   g?NTrU   F   )
unfreeze_backbone_at_epochlambda_funcbackbone_initial_ratio_lrbackbone_initial_lrshould_alignrW   r>   verboseroundingr   c
           
         sJ   t    || _|| _|| _|| _|| _|| _|| _|| _	|	| _
d | _d S r[   )superr   r   r   r   r   r   rW   r>   r   r   previous_backbone_lr)
r   r   r   r   r   r   rW   r>   r   r   	__class__r   r   r   X  s    
zBackboneFinetuning.__init__r   c                 C   s   | j | jdS )N)r   r   )r   r   r   r   r   r   r   q  s    zBackboneFinetuning.state_dictr    c                    s   |d | _ t | d S )Nr   )r   r   r"   r!   r   r   r   r"   w  s    
z"BackboneFinetuning.load_state_dictr#   r$   r%   c                    s0   t |dr$t|jtr$t ||S tddS )z
        Raises:
            MisconfigurationException:
                If LightningModule has no nn.Module `backbone` attribute.
        backbonez@The LightningModule should have a nn.Module `backbone` attributeN)hasattrr9   r   r
   r   r0   r   )r   r&   r'   r   r   r   r0   {  s    zBackboneFinetuning.on_fit_startrr   c                 C   s   |  |j d S r[   )rI   r   rs   r   r   r   r\     s    z)BackboneFinetuning.freeze_before_trainingro   c                 C   s
  || j kr|jd d }| jdur(| jn|| j }|| _| j|j||| j| jd | j	r~t
dt|| j dt|| j  n|| j kr|jd d }| |d | j }| jr||kr|n|}||jd d< || _| j	rt
dt|| j dt|| j  dS )	rk   r   rV   N)r>   rW   zCurrent lr: z, Backbone lr:    )r   r-   r   r   r   rY   r   r>   rW   r   loginforoundr   r   r   )r   r'   r   r/   r.   Z
current_lrZinitial_backbone_lrZnext_current_backbone_lrr   r   r   rm     sJ    


z$BackboneFinetuning.finetune_function)rt   ru   rv   rw   r   r{   r   rX   r   rz   r   r   rx   r   r   r"   r0   r\   r   rm   __classcell__r   r   r   r   r|   9  s:    
r|   )#rw   loggingtypingr   r   r   r   r   r   r   r	   rK   Ztorch.nnr
   r   Ztorch.nn.modules.batchnormr   Ztorch.optim.optimizerr   Zpytorch_lightningplZ$pytorch_lightning.callbacks.callbackr   Z&pytorch_lightning.utilities.exceptionsr   Z%pytorch_lightning.utilities.rank_zeror   	getLoggerrt   r   r{   rX   r   r   r|   r   r   r   r   <module>   s    (
  