a
    d                     @   s   d dl mZmZ d dlmZ d dlmZ d dlmZ d dl	m
Z
 G dd deZeG dd	 d	ZG d
d dejZG dd dZdS )    )	dataclassfield)Enum)DictN)AverageMeterc                   @   s   e Zd ZdZdZdZdZdS )TrainerStater            N)__name__
__module____qualname__ZSTARTINGZTRAININGZVALIDATEZ	TERMINATE r   r   W/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/kornia/x/utils.pyr      s   r   c                   @   s   e Zd ZU edddidZeed< edddidZeed< eddd	idZ	eed
< edddidZ
eed< edddidZeed< edddidZeed< dS )Configurationz./helpzThe input data directory.)defaultmetadata	data_pathr   z2The number of batches for the training dataloader.
batch_sizez)The number of epochs to run the training.
num_epochsgMbP?z.The learning rate to be used for the optimize.lrz./outputzThe output data directory.output_path)   r   zThe input image size.
image_sizeN)r   r   r   r   r   str__annotations__r   intr   r   floatr   r   tupler   r   r   r   r      s   
r   c                       s(   e Zd ZdZ fddZdd Z  ZS )LambdaaD  Module to create a lambda function as nn.Module.

    Args:
        fcn: a pointer to any function.

    Example:
        >>> import torch
        >>> import kornia as K
        >>> fcn = Lambda(lambda x: K.geometry.resize(x, (32, 16)))
        >>> fcn(torch.rand(1, 4, 64, 32)).shape
        torch.Size([1, 4, 32, 16])
    c                    s   t    || _d S N)super__init__fcn)selfr$   	__class__r   r   r#   9   s    
zLambda.__init__c                 C   s
   |  |S r!   )r$   )r%   xr   r   r   forward=   s    zLambda.forward)r   r   r   __doc__r#   r)   __classcell__r   r   r&   r   r    +   s   r    c                   @   s|   e Zd ZdZddddZedd Zeee	ddd	d
Z
eeef e	ddddZedddZeeef dddZdS )StatsTrackerz/Stats tracker for computing metrics on the fly.N)returnc                 C   s
   i | _ d S r!   _statsr%   r   r   r   r#   D   s    zStatsTracker.__init__c                 C   s   | j S r!   r.   r0   r   r   r   statsG   s    zStatsTracker.stats)keyvalr   r-   c                 C   s,   || j vrt | j |< | j | || dS )z'Update the stats by the key value pair.N)r/   r   update)r%   r2   r3   r   r   r   r   r4   K   s    
zStatsTracker.update)dicr   r-   c                 C   s$   |  D ]\}}| ||| qdS )zUpdate the stats by the dict.N)itemsr4   )r%   r5   r   kvr   r   r   update_from_dictQ   s    zStatsTracker.update_from_dictc                 C   s   d dd | j D S )N c                 S   s2   g | ]*\}}|   d |jdd|jddqS )z: z.2fr:   )upperr3   ).0r7   r8   r   r   r   
<listcomp>W       z)StatsTracker.__repr__.<locals>.<listcomp>)joinr/   r6   r0   r   r   r   __repr__V   s    zStatsTracker.__repr__c                 C   s   | j S )zReturn the dict format.r.   r0   r   r   r   as_dictY   s    zStatsTracker.as_dict)r   r   r   r*   r#   propertyr1   r   r   r   r4   r   r9   r@   r   rA   r   r   r   r   r,   A   s   
r,   )Zdataclassesr   r   enumr   typingr   Ztorch.nnnnZkornia.metrics.average_meterr   r   r   Moduler    r,   r   r   r   r   <module>   s   