a
    d%3                     @   s  d Z ddlZddlZddlmZmZmZ ddlm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mZmZ ddlZddlmZmZmZ dd	lmZ ed
edef dZeed ed f Z e eeef eed dddZ!e eeef edddZ"d1e e#edddZ$e eee#eef  dddZ%e e#dddZ&dd e ee'ed!d"d#Z(e#ee#e#f d$d%d&Z)e#ee*e#f d'd(d)Z+ee*e,e#f ee*e,f d'd*d+Z-ee*e#f ee*e#f d'd,d-Z.eed.d/d0Z/dS )2z;Utilities for Argument Parsing within Lightning Components.    N)_ArgumentGroupArgumentParser	Namespace)literal_eval)suppress)wraps)	AnyCallablecastDictListTupleTypeTypeVarUnion)str_to_boolstr_to_bool_or_intstr_to_bool_or_str)_ADD_ARGPARSE_RETURN_T.)boundpl.LightningDataModule
pl.Trainer)r   r   )clsargskwargsreturnc                    sZ   t |tr| |}t| t| jj} fdd|D }|jf i | | f i |S )az  Create an instance from CLI arguments. Eventually use variables from OS environment which are defined as
    ``"PL_<CLASS-NAME>_<CLASS_ARUMENT_NAME>"``.

    Args:
        cls: Lightning class
        args: The parser or namespace to take arguments from. Only known arguments will be
            parsed and passed to the :class:`Trainer`.
        **kwargs: Additional keyword arguments that may override ones in the parser or namespace.
            These must be valid Trainer arguments.

    Examples:

        >>> from pytorch_lightning import Trainer
        >>> parser = ArgumentParser(add_help=False)
        >>> parser = Trainer.add_argparse_args(parser)
        >>> parser.add_argument('--my_custom_arg', default='something')  # doctest: +SKIP
        >>> args = Trainer.parse_argparser(parser.parse_args(""))
        >>> trainer = Trainer.from_argparse_args(args, logger=False)
    c                    s   i | ]}| v r| | qS  r   ).0nameparamsr   m/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/utilities/argparse.py
<dictcomp>?       z&from_argparse_args.<locals>.<dictcomp>)	
isinstancer   parse_argparservarsinspect	signature__init__
parametersupdate)r   r   r   Zvalid_kwargsZtrainer_kwargsr   r    r"   from_argparse_args    s    

r-   )r   
arg_parserr   c           	      C   s   t |tr| n|}dd t| D }i }t| D ]B\}}||v rr|du rr|| \}}t|v rrt |trrd}|||< q8tf i |S )z4Parse CLI arguments, required for custom bool types.c                 S   s   i | ]\}}}|||fqS r   r   )r   arg	arg_typesarg_defaultr   r   r"   r#   I   r$   z#parse_argparser.<locals>.<dictcomp>NT)r%   r   
parse_argsget_init_arguments_and_typesr'   itemsboolr   )	r   r.   r   Ztypes_defaultZmodified_argskvr0   r1   r   r   r"   r&   E   s    
r&    PL_%(cls_name)s_%(cls_argument)s)r   templater   c              	   C   s   t | }i }|D ]v\}}}|| j | d }tj|}|du s|dkstt t|}W d   n1 st0    Y  |||< qt	f i |S )a  Parse environment arguments if they are defined.

    Examples:

        >>> from pytorch_lightning import Trainer
        >>> parse_env_variables(Trainer)
        Namespace()
        >>> import os
        >>> os.environ["PL_TRAINER_GPUS"] = '42'
        >>> os.environ["PL_TRAINER_BLABLABLA"] = '1.23'
        >>> parse_env_variables(Trainer)
        Namespace(gpus=42)
        >>> del os.environ["PL_TRAINER_GPUS"]
    )cls_nameZcls_argumentN )
r3   __name__upperosenvirongetr   	Exceptionr   r   )r   r9   Zcls_arg_defaultsZenv_argsZarg_name_envvalr   r   r"   parse_env_variables^   s    
&
rE   )r   r   c              
   C   s   t | j}g }|D ]}|| j}|| j}z`t|jdkrRtdd |jD }n8dt	|v sjdt	|v rtdd |jD }n
t|j}W n t
tfy   |f}Y n0 ||||f q|S )aN  Scans the class signature and returns argument names, types and default values.

    Returns:
        List with tuples of 3 values:
        (argument name, set with argument types, argument default value).

    Examples:

        >>> from pytorch_lightning import Trainer
        >>> args = get_init_arguments_and_types(Trainer)

    _LiteralGenericAliasc                 S   s   h | ]}t |qS r   )type)r   ar   r   r"   	<setcomp>   r$   z/get_init_arguments_and_types.<locals>.<setcomp>ztyping.Literalztyping_extensions.Literalc                 S   s    h | ]}|j D ]}t|qqS r   )__args__rG   )r   Z
union_argsrH   r   r   r"   rI      r$   )r(   r)   r+   
annotationdefaultrG   r<   tuplerJ   strAttributeError	TypeErrorappend)r   Zcls_default_paramsZname_type_defaultr/   Zarg_typer1   r0   r   r   r"   r3   |   s    

r3   c                 C   s@   t | tsJ t| | jdr.d| j S | j d| j S )Nzpytorch_lightning.zpl..)r%   rG   repr
__module__
startswithr<   __qualname__)r   r   r   r"   _get_abbrev_qualified_cls_name   s    rW   T)use_argument_group)r   parent_parserrX   r   c             	      s  t |trtd|r*t| }||}nt|gdd}g dtttt	f}| | j
fD ].}t|}fdd|D }t|dkrV qqVt| j
jp| jpd}|D ]"\}	 }
t fd	d
|D   sqi }t	 v r*|jddd t dkrt}n2t v rt}n"t v rt}ndd  D d }n d }|	dksF|	dkrJt}t dkrxtt v rxtt v rxt}|	dkrt}|	dkrt}|jd|	 f|	|
|||	|
tjkd| q|r|S |S )aj  Extends existing argparse by default attributes for ``cls``.

    Args:
        cls: Lightning class
        parent_parser:
            The custom cli arguments parser, which will be extended by
            the class's default arguments.
        use_argument_group:
            By default, this is True, and uses ``add_argument_group`` to add
            a new group.
            If False, this will use old behavior.

    Returns:
        If use_argument_group is True, returns ``parent_parser`` to keep old
        workflows. If False, will return the new ArgumentParser object.

    Only arguments of the allowed types (str, float, int, bool) will
    extend the ``parent_parser``.

    Raises:
        RuntimeError:
            If ``parent_parser`` is not an ``ArgumentParser`` instance

    Examples:

        >>> # Option 1: Default usage.
        >>> import argparse
        >>> from pytorch_lightning import Trainer
        >>> parser = argparse.ArgumentParser()
        >>> parser = Trainer.add_argparse_args(parser)
        >>> args = parser.parse_args([])

        >>> # Option 2: Disable use_argument_group (old behavior).
        >>> import argparse
        >>> from pytorch_lightning import Trainer
        >>> parser = argparse.ArgumentParser()
        >>> parser = Trainer.add_argparse_args(parser, use_argument_group=False)
        >>> args = parser.parse_args([])
    z.Please only pass an `ArgumentParser` instance.F)parentsadd_help)selfr   r   c                    s   g | ]}|d   vr|qS r   r   )r   x)ignore_arg_namesr   r"   
<listcomp>   r$   z%add_argparse_args.<locals>.<listcomp>r   r;   c                 3   s   | ]}| v r|V  qd S )Nr   r   at)r0   r   r"   	<genexpr>   r$   z$add_argparse_args.<locals>.<genexpr>?T)nargsconst   c                 S   s   g | ]}|t ur|qS r   )r5   ra   r   r   r"   r`      r$   ZgpusZ	tpu_cores   Ztrack_grad_normZ	precisionz--)destrL   rG   helprequired)r%   r   RuntimeErrorrW   add_argument_groupr   rN   intfloatr5   r*   r3   len_parse_args_from_docstring__doc__rM   r,   r   r   r   _gpus_allowed_typeset_int_or_float_type_precision_allowed_typeadd_argumentr@   r(   _empty)r   rY   rX   Z
group_nameparserZallowed_typessymbolZargs_and_typesZ	args_helpr/   r1   Z
arg_kwargsZuse_typer   )r0   r_   r"   add_argparse_args   sf    -



*


r{   )	docstringr   c                 C   s   d }d}i }|  dD ]}| }|s(qt|t| }|drL|d }q|d u rXqq||k rf qq||kr|j ddd\}}| ||< q||kr||  d| 7  < q|S )	Nr;   
)zArgs:z
Arguments:zParameters:   :rg   )maxsplit )splitlstriprp   rU   )r|   Zarg_block_indentZcurrent_argparsedlinestrippedZline_indentZarg_descriptionr   r   r"   rq     s(    

rq   )r^   r   c                 C   s   d| v rt | S t| S )N,)rN   rn   r^   r   r   r"   rs   5  s    rs   c                 C   s   dt | v rt| S t| S )NrR   )rN   ro   rn   r   r   r   r"   ru   ;  s    ru   c                 C   s&   z
t | W S  ty    |  Y S 0 dS )ze
    >>> _precision_allowed_type("32")
    32
    >>> _precision_allowed_type("bf16")
    'bf16'
    N)rn   
ValueErrorr   r   r   r"   rv   A  s    
rv   )fnr   c                    s*   t  ttttd fdd}tt|S )N)r\   r   r   r   c                    sh   | j }|r0dd t|D }|tt|| tt|}tt| t|  } | fi |S )Nc                 S   s   g | ]}|d  qS r]   r   )r   r/   r   r   r"   r`   T  r$   zH_defaults_from_env_vars.<locals>.insert_env_defaults.<locals>.<listcomp>)		__class__r3   r,   dictzipr'   rE   listr4   )r\   r   r   r   Zcls_arg_namesZenv_variablesr   r   r"   insert_env_defaultsO  s    z4_defaults_from_env_vars.<locals>.insert_env_defaults)r   r   r
   r   )r   r   r   r   r"   _defaults_from_env_varsN  s    r   )r8   )0rr   r(   r>   argparser   r   r   astr   
contextlibr   	functoolsr   typingr   r	   r
   r   r   r   r   r   r   Zpytorch_lightningplZ#pytorch_lightning.utilities.parsingr   r   r   Z!pytorch_lightning.utilities.typesr   r   Z_ARGPARSE_CLSr-   r&   rN   rE   r3   rW   r5   r{   rq   rn   rs   ro   ru   rv   r   r   r   r   r"   <module>   sB   ,
%$u" 