a
    d                     @   s   d dl mZ d dlmZ d dlmZmZmZmZm	Z	 d dl
Z
d dlm  mZ dd Zdee ee	eef  ee e	ee
jf ed
ddZdee ee	eef  e	ee
jf edddZdS )    )partial)islice)CallableListOptionalSequenceUnionNc                 c   s(   t | }tt||}|sq$|V  qdS )zBatch data into lists of length *n*. The last batch may be shorter.
    NOTE based on more-itertools impl, to be replaced by python 3.12 itertools.batched impl
    N)iterlistr   )iterablenitbatch r   g/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/open_clip/zero_shot_classifier.pybatched	   s
    r   
   cpuF)
classnames	templatesnum_classes_per_batchdeviceuse_tqdmc                    s
  t trtdksJ t |tr0t|dks4J t d ttt|}|rddl}|du rjdn|d | d }	t|j|	|d}
nt}
fdd t H |r fdd|
t	||D }tj
|dd	}n |}W d   n1 s0    Y  |S )
a   Build zero-shot classifier weights by iterating over class names in batches
    Args:
        model: CLIP model instance
        tokenizer: CLIP tokenizer instance
        classnames: A sequence of class (label) names
        templates: A sequence of callables or format() friendly strings to produce templates per class name
        num_classes_per_batch: The number of classes to batch together in each forward, all if None
        device: Device to use.
        use_tqdm: Enable TQDM progress bar.
    r   N   )totalZ
unit_scalec                    sp   t | }fdd| D }| }tj|dd}||djdd}||jddd }|j}|S )Nc                    s,   g | ]$} D ]}r| |n||qqS r   format).0ctemplate)r   
use_formatr   r   
<listcomp>6       zFbuild_zero_shot_classifier.<locals>._process_batch.<locals>.<listcomp>dimr   T)r%   Zkeepdim)	lentoF	normalizeencode_textZreshapemeannormT)Zbatch_classnamesZnum_batch_classestextsclass_embeddings)r   modelnum_templatesr   	tokenizerr    r   r   _process_batch4   s    z2build_zero_shot_classifier.<locals>._process_batchc                    s   g | ]} |qS r   r   )r   r   )r3   r   r   r!   @   r"   z.build_zero_shot_classifier.<locals>.<listcomp>r$   )
isinstancer   r&   strtqdmr   r	   torchno_gradr   cat)r0   r2   r   r   r   r   r   Znum_classesr6   Znum_iter	iter_wrapZbatched_embedszeroshot_weightsr   )r3   r   r0   r1   r   r2   r    r   build_zero_shot_classifier   s"    

&r<   )r   r   r   r   c                    s  t |trt|dksJ t |tr0t|dks4J |rHddl}|j}nt}t |d tt  g }||D ]\  fdd|D }	||	|}	| 	|	}
t
j|
ddjdd}||  }|| qptj|dd|}W d   n1 s0    Y  |S )a   Build zero-shot classifier weights by iterating over class names 1 by 1
    Args:
        model: CLIP model instance
        tokenizer: CLIP tokenizer instance
        classnames: A sequence of class (label) names
        templates: A sequence of callables or format() friendly strings to produce templates per class name
        device: Device to use.
        use_tqdm: Enable TQDM progress bar.
    r   Nc                    s"   g | ]}r|  n| qS r   r   )r   r   	classnamer    r   r   r!   e   r"   z5build_zero_shot_classifier_legacy.<locals>.<listcomp>r#   r$   r   )r4   r   r&   r6   r	   r5   r7   r8   r'   r*   r(   r)   r+   r,   appendstack)r0   r2   r   r   r   r   r6   r:   r;   r.   r/   Zclass_embeddingr   r=   r   !build_zero_shot_classifier_legacyG   s$    

2rA   )r   r   F)r   F)	functoolsr   	itertoolsr   typingr   r   r   r   r   r7   Ztorch.nn.functionalnnZ
functionalr(   r   r5   intr   boolr<   rA   r   r   r   r   <module>   s.      7  