a
    d                     @   s   d Z ddlZddlZddlZddlmZ ddlmZmZ ddl	Z	ddl
ZddlZddlZdejd< e dd Ze d	d
 Zdd Zdd Zdd ZG dd deZe ZejdddZdeeee f eejdddZG dd dZdS )zp CLIP tokenizer

Copied from https://github.com/openai/CLIP. Originally MIT License, Copyright (c) 2021 OpenAI.
    N)	lru_cache)UnionListfalseZTOKENIZERS_PARALLELISMc                   C   s   t jt jt jtdS )Nzbpe_simple_vocab_16e6.txt.gz)ospathjoindirnameabspath__file__ r   r   \/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/open_clip/tokenizer.pydefault_bpe   s    r   c                  C   s   t ttdtdd t ttdtdd  t ttdtdd  } | dd }d	}td
D ],}|| vrf| | |d
|  |d7 }qfdd |D }tt| |S )a:  
    Returns list of utf-8 byte and a corresponding list of unicode strings.
    The reversible bpe codes work on unicode strings.
    This means you need a large # of unicode characters in your vocab if you want to avoid UNKs.
    When you're at something like a 10B token dataset you end up needing around 5K for decent coverage.
    This is a significant percentage of your normal, say, 32K bpe vocab.
    To avoid that, we want lookup tables between utf-8 bytes and unicode strings.
    And avoids mapping to whitespace/control characters the bpe code barfs on.
    !~      ¡   ¬   ®   ÿNr      c                 S   s   g | ]}t |qS r   )chr).0nr   r   r   
<listcomp>,       z$bytes_to_unicode.<locals>.<listcomp>)listrangeordappenddictzip)bscsr   br   r   r   bytes_to_unicode   s    N

r%   c                 C   s6   t  }| d }| dd D ]}|||f |}q|S )zReturn set of symbol pairs in a word.
    Word is represented as tuple of symbols (symbols being variable-length strings).
    r   r   N)setadd)wordpairsZ	prev_charcharr   r   r   	get_pairs0   s    r+   c                 C   s"   t | } tt| } |  S N)ftfyZfix_texthtmlunescapestriptextr   r   r   basic_clean<   s    
r3   c                 C   s   t dd| } |  } | S )Nz\s+ )resubr0   r1   r   r   r   whitespace_cleanB   s    r7   c                   @   s:   e Zd Ze dfedddZdd Zdd Zd	d
 ZdS )SimpleTokenizerN)bpe_pathc                    sH  t   _dd  j D  _t| dd}|dd }dd |D }t	t  
 }|d	d |D  }|D ]}|d
| qv|sddg}nddg| }|| tt|tt| _dd  j D  _tt|tt| _dd |D  _d|}t|d tj _t j _ fdd|D  _d S )Nc                 S   s   i | ]\}}||qS r   r   r   kvr   r   r   
<dictcomp>K   r   z,SimpleTokenizer.__init__.<locals>.<dictcomp>utf-8
r   i  c                 S   s   g | ]}t | qS r   )tuplesplit)r   merger   r   r   r   N   r   z,SimpleTokenizer.__init__.<locals>.<listcomp>c                 S   s   g | ]}|d  qS )</w>r   )r   r<   r   r   r   r   P   r    <start_of_text><end_of_text>c                 S   s   i | ]\}}||qS r   r   r:   r   r   r   r=   Y   r   c                 S   s   i | ]
}||qS r   r   r   tr   r   r   r=   [   r   |z:|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+c                    s   g | ]} j | qS r   encoderrG   selfr   r   r   `   r   )r%   byte_encoderitemsbyte_decodergzipopenreaddecoderA   r   valuesr   r   extendr    r!   r   lenrK   decoder	bpe_rankscacher5   compile
IGNORECASEpatZ
vocab_sizeZall_special_ids)rM   r9   Zspecial_tokensZmergesZvocabrB   Zspecialr   rL   r   __init__I   s*    


zSimpleTokenizer.__init__c           
         sv  | j v r j | S t|d d |d d f }t|}|sF|d S t| fddd}| jvrhq^|\}}g }d}|t|k r4z&|||}	||||	  |	}W n$   |||d   Y q4Y n0 || |kr|t|d k r||d  |kr|||  |d7 }qx|||  |d7 }qxt|}|}t|dkrTq^qFt|}qFd		|}| j |< |S )
NrC   c                    s    j | tdS )Ninf)rY   getfloat)pairrL   r   r   <lambda>l   r   z%SimpleTokenizer.bpe.<locals>.<lambda>)keyr   r      r4   )
rZ   r@   r+   minrY   rW   indexrV   r   r   )
rM   tokenr(   r)   ZbigramfirstsecondZnew_wordijr   rL   r   bpeb   sB    


2




zSimpleTokenizer.bpec                    sn   g }t t| }t j|D ]F}d fdd|dD }| fdd 	|
dD  q"|S )NrD   c                 3   s   | ]} j | V  qd S r,   )rN   )r   r$   rL   r   r   	<genexpr>   r   z)SimpleTokenizer.encode.<locals>.<genexpr>r>   c                 3   s   | ]} j | V  qd S r,   rJ   )r   Z	bpe_tokenrL   r   r   ro      r   r4   )r7   r3   lowerr5   findallr]   r   encoderV   rn   rA   )rM   r2   Z
bpe_tokensri   r   rL   r   rr      s    &zSimpleTokenizer.encodec                    sD   d  fdd|D }t fdd|D jddddd	}|S )
NrD   c                    s   g | ]} j | qS r   )rX   )r   ri   rL   r   r   r      r   z*SimpleTokenizer.decode.<locals>.<listcomp>c                    s   g | ]} j | qS r   )rP   )r   crL   r   r   r      r   r>   replace)errorsrC   r4   )r   	bytearrayrT   rt   )rM   tokensr2   r   rL   r   rT      s    (zSimpleTokenizer.decode)	__name__
__module____qualname__r   strr^   rn   rr   rT   r   r   r   r   r8   H   s   )r8   Z
output_idsc                 C   s   |    } t| S r,   )cpunumpy
_tokenizerrT   r|   r   r   r   rT      s    rT   M   textscontext_lengthreturnc                    s   t | tr| g} tjd tjd   fdd| D }tjt||tjd}t|D ]B\}}t||kr~|d| } |d< t	|||dt|f< qV|S )a  
    Returns the tokenized representation of given input string(s)

    Parameters
    ----------
    texts : Union[str, List[str]]
        An input string or a list of input strings to tokenize
    context_length : int
        The context length to use; all CLIP models use 77 as the context length

    Returns
    -------
    A two-dimensional tensor containing the resulting tokens, shape = [number of input strings, context_length]
    rE   rF   c                    s"   g | ]}gt |  g qS r   )r   rr   r   r2   Z	eot_tokenZ	sot_tokenr   r   r      r   ztokenize.<locals>.<listcomp>)ZdtypeNr_   )

isinstancer{   r   rK   torchzerosrW   long	enumerateZtensor)r   r   Z
all_tokensresultrl   rw   r   r   r   tokenize   s    


r   c                   @   sH   e Zd ZdZedddZdd Zdeeee f e	e
jdd	d
ZdS )HFTokenizerzHuggingFace tokenizer wrapper)tokenizer_namec                 C   s   ddl m} ||| _d S )Nr   )AutoTokenizer)Ztransformersr   Zfrom_pretrained	tokenizer)rM   r   r   r   r   r   r^      s    zHFTokenizer.__init__c                 C   s   | j | d S r,   )r   save_pretrained)rM   destr   r   r   r      s    zHFTokenizer.save_pretrainedr   r   c                 C   s8   t |tr|g}dd |D }| j|d|dddj}|S )Nc                 S   s   g | ]}t t|qS r   )r7   r3   r   r   r   r   r      r   z(HFTokenizer.__call__.<locals>.<listcomp>pt
max_lengthT)Zreturn_tensorsr   paddingZ
truncation)r   r{   r   	input_ids)rM   r   r   r   r   r   r   __call__   s    
zHFTokenizer.__call__N)r   )rx   ry   rz   __doc__r{   r^   r   r   r   intr   Tensorr   r   r   r   r   r      s   r   )r   )r   rQ   r.   r   	functoolsr   typingr   r   r-   regexr5   r   environr   r%   r+   r3   r7   objectr8   r   r   rT   r{   r   Z
LongTensorr   r   r   r   r   r   <module>   s,   


Q" 