a
    dG                     @   s
  d Z ddlZddlZddlmZmZ ddlmZ ddlm	Z	m
Z
 ddlmZ ddlmZ ddlmZmZmZmZ ddlZd	d
lmZmZmZmZ e rddlmZ G dd deZdd Zdd Z dd Z!dd Z"dd Z#dd Z$dd Z%dd Z&dd Z'dd  Z(d!d" Z)d#d$ Z*d%d& Z+d'd( Z,d)d* Z-d+d, Z.d-d. Z/d/d0 Z0G d1d2 d2eZ1G d3d4 d4e2eZ3G d5d6 d6e3Z4G d7d8 d8e3Z5G d9d: d:Z6d;d< Z7d=d> Z8dVee2e2dAdBdCZ9e
dWe:dEdFdGZ;dXdHdIZ<dJdK Z=dYdLdMZ>dNdO Z?dPdQ Z@dRdS ZAdTdU ZBdS )Zz
Generic utilities
    N)OrderedDictUserDict)MutableMapping)	ExitStackcontextmanager)fields)Enum)AnyContextManagerListTuple   )is_flax_availableis_tf_availableis_torch_availableis_torch_fx_proxyc                   @   s   e Zd ZdZdddZdS )cached_propertyz
    Descriptor that mimics @property but caches output in member variable.

    From tensorflow_datasets

    Built-in in functools from Python 3.8.
    Nc                 C   sX   |d u r| S | j d u rtdd| j j }t||d }|d u rT|  |}t||| |S )Nzunreadable attributeZ	__cached_)fgetAttributeError__name__getattrsetattr)selfobjZobjtypeattrcached r   c/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/transformers/utils/generic.py__get__-   s    

zcached_property.__get__)N)r   
__module____qualname____doc__r   r   r   r   r   r   $   s   r   c                 C   s2   |   } | dv rdS | dv r dS td| dS )zConvert a string representation of truth to true (1) or false (0).

    True values are 'y', 'yes', 't', 'true', 'on', and '1'; false values are 'n', 'no', 'f', 'false', 'off', and '0'.
    Raises ValueError if 'val' is anything else.
    >   on1tytrueyesr   >   falsenofn0offr   zinvalid truth value N)lower
ValueError)valr   r   r   	strtobool<   s    r1   c                 C   s   t | rdS t r*ddl}t| |jr*dS t rHddl}t| |jrHdS t rzddlm	} ddl
m} t| |j|frzdS t| tjS )zl
    Tests if `x` is a `torch.Tensor`, `tf.Tensor`, `jaxlib.xla_extension.DeviceArray` or `np.ndarray`.
    Tr   N)Tracer)r   r   torch
isinstanceTensorr   
tensorflowr   	jax.numpynumpyZjax.corer2   ndarraynp)xr3   tfjnpr2   r   r   r   	is_tensorJ   s     r>   c                 C   s   t | tjS N)r4   r:   r9   r;   r   r   r   	_is_numpye   s    rA   c                 C   s   t | S )z/
    Tests if `x` is a numpy array or not.
    )rA   r@   r   r   r   is_numpy_arrayi   s    rB   c                 C   s   dd l }t| |jS Nr   )r3   r4   r5   r;   r3   r   r   r   	_is_torchp   s    rE   c                 C   s   t  s
dS t| S )z]
    Tests if `x` is a torch tensor or not. Safe to call even if torch is not installed.
    F)r   rE   r@   r   r   r   is_torch_tensorv   s    rF   c                 C   s   dd l }t| |jS rC   )r3   r4   ZdevicerD   r   r   r   _is_torch_device}   s    rG   c                 C   s   t  s
dS t| S )z]
    Tests if `x` is a torch device or not. Safe to call even if torch is not installed.
    F)r   rG   r@   r   r   r   is_torch_device   s    rH   c                 C   s8   dd l }t| tr,t|| r(t|| } ndS t| |jS )Nr   F)r3   r4   strhasattrr   ZdtyperD   r   r   r   _is_torch_dtype   s    

rK   c                 C   s   t  s
dS t| S )z\
    Tests if `x` is a torch dtype or not. Safe to call even if torch is not installed.
    F)r   rK   r@   r   r   r   is_torch_dtype   s    rL   c                 C   s   dd l }t| |jS rC   )r6   r4   r5   r;   r<   r   r   r   _is_tensorflow   s    rN   c                 C   s   t  s
dS t| S )zg
    Tests if `x` is a tensorflow tensor or not. Safe to call even if tensorflow is not installed.
    F)r   rN   r@   r   r   r   is_tf_tensor   s    rO   c                 C   s*   dd l }t|dr|| S t| |jkS )Nr   is_symbolic_tensor)r6   rJ   rP   typer5   rM   r   r   r   _is_tf_symbolic_tensor   s    

rR   c                 C   s   t  s
dS t| S )z
    Tests if `x` is a tensorflow symbolic tensor or not (ie. not eager). Safe to call even if tensorflow is not
    installed.
    F)r   rR   r@   r   r   r   is_tf_symbolic_tensor   s    rS   c                 C   s   dd l m} t| |jS rC   )r7   r8   r4   r9   )r;   r=   r   r   r   _is_jax   s    rT   c                 C   s   t  s
dS t| S )zY
    Tests if `x` is a Jax tensor or not. Safe to call even if jax is not installed.
    F)r   rT   r@   r   r   r   is_jax_tensor   s    rU   c                 C   s   t | ttfr dd |  D S t | ttfr<dd | D S t| rP|   S t	| rh| 
   S t| r~t|  S t | tjtjfr|  S | S dS )zc
    Convert a TensorFlow tensor, PyTorch tensor, Numpy array or python list to a python list.
    c                 S   s   i | ]\}}|t |qS r   	to_py_obj.0kvr   r   r   
<dictcomp>       zto_py_obj.<locals>.<dictcomp>c                 S   s   g | ]}t |qS r   rV   )rY   or   r   r   
<listcomp>   r]   zto_py_obj.<locals>.<listcomp>N)r4   dictr   itemslisttuplerO   r8   tolistrF   detachcpurU   r:   asarrayr9   numberr   r   r   r   rW      s    rW   c                 C   sz   t | ttfr dd |  D S t | ttfr8t| S t| rH| 	 S t
| r`|   	 S t| rrt| S | S dS )zc
    Convert a TensorFlow tensor, PyTorch tensor, Numpy array or python list to a Numpy array.
    c                 S   s   i | ]\}}|t |qS r   )to_numpyrX   r   r   r   r\      r]   zto_numpy.<locals>.<dictcomp>N)r4   r`   r   ra   rb   rc   r:   arrayrO   r8   rF   re   rf   rU   rg   ri   r   r   r   rj      s    

rj   c                       sn   e Zd ZdZdd Zdd Zdd Zdd	 Zd
d Zdd Z	 fddZ
 fddZee dddZ  ZS )ModelOutputa  
    Base class for all model outputs as dataclass. Has a `__getitem__` that allows indexing by integer or slice (like a
    tuple) or strings (like a dictionary) that will ignore the `None` attributes. Otherwise behaves like a regular
    python dictionary.

    <Tip warning={true}>

    You can't unpack a `ModelOutput` directly. Use the [`~utils.ModelOutput.to_tuple`] method to convert it to a tuple
    before.

    </Tip>
    c           
         s  t  }t|s"t jj dtdd |dd  D sNt jj dt |d j}t fdd|dd  D }|rt|st	|t
r| }d}n(zt|}d}W n ty   d	}Y n0 |rtt|D ]\}}t	|ttfrt|d
krt	|d ts@|dkr*| |d j< ntd| d qt |d |d  |d d ur|d  |d < qn|d ur| |d j< n,|D ]&}t |j}	|	d ur|	 |j< qd S )Nz has no fields.c                 s   s   | ]}|j d u V  qd S r?   )defaultrY   fieldr   r   r   	<genexpr>  r]   z,ModelOutput.__post_init__.<locals>.<genexpr>r   z. should not have more than one required field.r   c                 3   s   | ]}t  |jd u V  qd S r?   )r   namern   r   r   r   rp     r]   TF   zCannot set key/value for z&. It needs to be a tuple (key, value).)r   lenr/   	__class__r   allr   rq   r>   r4   r`   ra   iter	TypeError	enumeraterb   rc   rI   r   )
r   Zclass_fieldsZfirst_fieldZother_fields_are_noneiteratorZfirst_field_iteratoridxelementro   r[   r   rr   r   __post_init__   sN    






zModelOutput.__post_init__c                 O   s   t d| jj dd S )Nz$You cannot use ``__delitem__`` on a 
 instance.	Exceptionru   r   r   argskwargsr   r   r   __delitem__/  s    zModelOutput.__delitem__c                 O   s   t d| jj dd S )Nz#You cannot use ``setdefault`` on a r~   r   r   r   r   r   
setdefault2  s    zModelOutput.setdefaultc                 O   s   t d| jj dd S )NzYou cannot use ``pop`` on a r~   r   r   r   r   r   pop5  s    zModelOutput.popc                 O   s   t d| jj dd S )NzYou cannot use ``update`` on a r~   r   r   r   r   r   update8  s    zModelOutput.updatec                 C   s.   t |trt|  }|| S |  | S d S r?   )r4   rI   r`   ra   to_tuple)r   rZ   Z
inner_dictr   r   r   __getitem__;  s    
zModelOutput.__getitem__c                    s4   ||   v r"|d ur"t || t || d S r?   )keyssuper__setitem____setattr__)r   rq   valueru   r   r   r   B  s    zModelOutput.__setattr__c                    s    t  || t  || d S r?   )r   r   r   )r   keyr   r   r   r   r   H  s    zModelOutput.__setitem__)returnc                    s   t  fdd  D S )za
        Convert self to a tuple containing all the attributes/keys that are not `None`.
        c                 3   s   | ]} | V  qd S r?   r   )rY   rZ   rr   r   r   rp   R  r]   z'ModelOutput.to_tuple.<locals>.<genexpr>)rc   r   rr   r   rr   r   r   N  s    zModelOutput.to_tuple)r   r   r    r!   r}   r   r   r   r   r   r   r   r   r	   r   __classcell__r   r   r   r   rl      s   4rl   c                   @   s   e Zd ZdZedd ZdS )ExplicitEnumzC
    Enum with more explicit error message for missing values.
    c                 C   s(   t | d| j dt| j  d S )Nz is not a valid z, please select one of )r/   r   rb   _value2member_map_r   )clsr   r   r   r   	_missing_Z  s    zExplicitEnum._missing_N)r   r   r    r!   classmethodr   r   r   r   r   r   U  s   r   c                   @   s   e Zd ZdZdZdZdZdS )PaddingStrategyz
    Possible values for the `padding` argument in [`PreTrainedTokenizerBase.__call__`]. Useful for tab-completion in an
    IDE.
    longest
max_lengthZ
do_not_padN)r   r   r    r!   ZLONGESTZ
MAX_LENGTHZ
DO_NOT_PADr   r   r   r   r   a  s   r   c                   @   s    e Zd ZdZdZdZdZdZdS )
TensorTypez
    Possible values for the `return_tensors` argument in [`PreTrainedTokenizerBase.__call__`]. Useful for
    tab-completion in an IDE.
    ptr<   r:   jaxN)r   r   r    r!   ZPYTORCHZ
TENSORFLOWZNUMPYZJAXr   r   r   r   r   l  s
   r   c                   @   s2   e Zd ZdZee dddZdd Zdd Zd	S )
ContextManagersz
    Wrapper for `contextlib.ExitStack` which enters a collection of context managers. Adaptation of `ContextManagers`
    in the `fastcore` library.
    )context_managersc                 C   s   || _ t | _d S r?   )r   r   stack)r   r   r   r   r   __init__~  s    zContextManagers.__init__c                 C   s   | j D ]}| j| qd S r?   )r   r   enter_context)r   Zcontext_managerr   r   r   	__enter__  s    
zContextManagers.__enter__c                 O   s   | j j|i | d S r?   )r   __exit__r   r   r   r   r     s    zContextManagers.__exit__N)	r   r   r    r!   r   r
   r   r   r   r   r   r   r   r   x  s   r   c                 C   sn   t | }|dkrt| j}n"|dkr4t| j}nt| j}|jD ]"}|dkrF|j| jdu rF dS qFdS )zr
    Check if a given model can return loss.

    Args:
        model_class (`type`): The class of the model.
    r<   r   Zreturn_lossTF)infer_frameworkinspect	signaturecallforward__call__
parametersrm   )model_class	frameworkr   pr   r   r   can_return_loss  s    
r   c                 C   sr   | j }t| }|dkr$t| j}n"|dkr:t| j}nt| j}d|v r^dd |jD S dd |jD S dS )zq
    Find the labels used by a given model.

    Args:
        model_class (`type`): The class of the model.
    r<   r   ZQuestionAnsweringc                 S   s    g | ]}d |v s|dv r|qS )label)Zstart_positionsZend_positionsr   rY   r   r   r   r   r_     r]   zfind_labels.<locals>.<listcomp>c                 S   s   g | ]}d |v r|qS )r   r   r   r   r   r   r_     r]   N)r   r   r   r   r   r   r   r   )r   Z
model_namer   r   r   r   r   find_labels  s    r    .)d
parent_key	delimiterc                 C   s   ddd}t || ||S )z/Flatten a nested dict into a single level dict.r   r   c                 s   sd   |   D ]V\}}|r(t|| t| n|}|rTt|trTt|||d  E d H  q||fV  qd S )N)r   )ra   rI   r4   r   flatten_dict)r   r   r   rZ   r[   r   r   r   r   _flatten_dict  s
    z#flatten_dict.<locals>._flatten_dict)r   r   )r`   )r   r   r   r   r   r   r   r     s    
r   F)use_temp_dirc                 c   s>   |r4t  }|V  W d    q:1 s(0    Y  n| V  d S r?   )tempfileTemporaryDirectory)Zworking_dirr   Ztmp_dirr   r   r   working_or_temp_dir  s    
&r   c                 C   s   t | rtj| |dS t| r6|du r,| jS | j| S t| rTddl}|j| |dS t| rjt	j| |dS t
dt|  ddS )z
    Framework-agnostic version of `numpy.transpose` that will work on torch/TensorFlow/Jax tensors as well as NumPy
    arrays.
    )axesNr   )permz"Type not supported for transpose: r   )rB   r:   	transposerF   TZpermuterO   r6   rU   r=   r/   rQ   )rk   r   r<   r   r   r   r     s    r   c                 C   sn   t | rt| |S t| r&| j| S t| rBddl}|| |S t| rVt| |S tdt	|  ddS )z
    Framework-agnostic version of `numpy.reshape` that will work on torch/TensorFlow/Jax tensors as well as NumPy
    arrays.
    r   Nz Type not supported for reshape: r   )
rB   r:   reshaperF   rO   r6   rU   r=   r/   rQ   )rk   Znewshaper<   r   r   r   r     s    
r   c                 C   s   t | rtj| |dS t| r:|du r.|  S | j|dS t| rXddl}|j| |dS t| rntj| |dS tdt	|  ddS )z
    Framework-agnostic version of `numpy.squeeze` that will work on torch/TensorFlow/Jax tensors as well as NumPy
    arrays.
    axisNZdimr   z Type not supported for squeeze: r   )
rB   r:   squeezerF   rO   r6   rU   r=   r/   rQ   rk   r   r<   r   r   r   r     s    r   c                 C   st   t | rt| |S t| r(| j|dS t| rFddl}|j| |dS t| r\tj| |dS t	dt
|  ddS )z
    Framework-agnostic version of `numpy.expand_dims` that will work on torch/TensorFlow/Jax tensors as well as NumPy
    arrays.
    r   r   Nr   $Type not supported for expand_dims: r   )rB   r:   expand_dimsrF   Z	unsqueezerO   r6   rU   r=   r/   rQ   r   r   r   r   r     s    r   c                 C   sb   t | rt| S t| r"|  S t| r<ddl}|| S t| rJ| jS tdt	|  ddS )z|
    Framework-agnostic version of `numpy.size` that will work on torch/TensorFlow/Jax tensors as well as NumPy arrays.
    r   Nr   r   )
rB   r:   sizerF   ZnumelrO   r6   rU   r/   rQ   )rk   r<   r   r   r   tensor_size  s    

r   c                    s^   |   D ]P\}}t|ttfr6 fdd|D | |< q|durd|vr  d| | |< q| S )zB
    Adds the information of the repo_id to a given auto map.
    c                    s.   g | ]&}|d ur&d|vr&  d| n|qS )N--r   )rY   r[   repo_idr   r   r_   1  r]   z.add_model_info_to_auto_map.<locals>.<listcomp>Nr   )ra   r4   rc   rb   )Zauto_mapr   r   r   r   r   r   add_model_info_to_auto_map+  s    r   c                 C   s   t | D ]l}|j}|j}|ds6|ds6|dkr< dS |dsN|dkrT dS |dsp|d	sp|d
kr
 dS q
td|  ddS )z
    Infers the framework of a given model without using isinstance(), because we cannot guarantee that the relevant
    classes are imported or available.
    r6   ZkerasZTFPreTrainedModelr<   r3   ZPreTrainedModelr   Zflaxr   ZFlaxPreTrainedModelz%Could not infer framework from class r   N)r   getmror   r   
startswithrx   )r   Z
base_classmodulerq   r   r   r   r   8  s    r   )r   r   )F)N)N)Cr!   r   r   collectionsr   r   collections.abcr   
contextlibr   r   Zdataclassesr   enumr   typingr	   r
   r   r   r8   r:   Zimport_utilsr   r   r   r   r7   r=   propertyr   r1   r>   rA   rB   rE   rF   rG   rH   rK   rL   rN   rO   rR   rS   rT   rU   rW   rj   rl   rI   r   r   r   r   r   r   r   boolr   r   r   r   r   r   r   r   r   r   r   r   <module>   s`   	h

