
    Wi                     b    d dl Zd dlmZ d dlmZ d dlmZmZ d dlm	Z	 d dl
mZ  G d de      Zy)	    N)Image)tqdm)ListUnion)SentenceTransformer)BaseEmbedderc            	            e Zd ZdZ	 	 ddeeef   deeef   def fdZdde	e   de	e   dz  d	e
d
ej                  fdZdde	e   d	e
d
ej                  fdZdde	e   d	e
d
ej                  fdZd Zd Z xZS )MultiModalBackenda  Multimodal backend using Sentence-transformers.

    The sentence-transformers embedding model used for
    generating word, document, and image embeddings.

    Arguments:
        embedding_model: A sentence-transformers embedding model that
                         can either embed both images and text or only text.
                         If it only embeds text, then `image_model` needs
                         to be used to embed the images.
        image_model: A sentence-transformers embedding model that is used
                     to embed only images.
        batch_size: The sizes of image batches to pass

    Examples:
    To create a model, you can load in a string pointing to a
    sentence-transformers model:

    ```python
    from bertopic.backend import MultiModalBackend

    sentence_model = MultiModalBackend("clip-ViT-B-32")
    ```

    or  you can instantiate a model yourself:
    ```python
    from bertopic.backend import MultiModalBackend
    from sentence_transformers import SentenceTransformer

    embedding_model = SentenceTransformer("clip-ViT-B-32")
    sentence_model = MultiModalBackend(embedding_model)
    ```
    Nembedding_modelimage_model
batch_sizec                     t         |           || _        t        |t              r|| _        n,t        |t              rt	        |      | _        nt        d      d | _        |Dt        |t              r|| _        n,t        |t              rt	        |      | _        nt        d      	 | j
                  j                         j                  j                  | _        y # t        $ r | j
                  j                  | _        Y y  d | _        Y y xY w)NzPlease select a correct SentenceTransformers model: 
`from sentence_transformers import SentenceTransformer` 
`model = SentenceTransformer('clip-ViT-B-32')`)super__init__r   
isinstancer   r   str
ValueErrorr   _first_module	processor	tokenizerAttributeError)selfr   r   r   	__class__s       h/home/sietch6/trending-topics-pipeline/venv/lib/python3.12/site-packages/bertopic/backend/_multimodal.pyr   zMultiModalBackend.__init__-   s     	$ o':;#2D -#6#GD A   "+':;#. K-#6{#C  E 	"!11??AKKUUDN 	<!11;;DN	"!DNs   )3C $DD	documentsimagesverbosereturnc                     d}|d   | j                  |      }d}t        |t              r| j                  ||      }d}||t	        j
                  ||gd      }||S ||S ||S y)aG  Embed a list of n documents/words or images into an n-dimensional
        matrix of embeddings.

        Either documents, images, or both can be provided. If both are provided,
        then the embeddings are averaged.

        Arguments:
            documents: A list of documents or words to be embedded
            images: A list of image paths to be embedded
            verbose: Controls the verbosity of the process

        Returns:
            Document/words embeddings with shape (n, m) with `n` documents/words
            that each have an embeddings size of `m`
        Nr   )axis)embed_documentsr   listembed_imagesnpmean)r   r   r   r   doc_embeddingsimage_embeddingsaveraged_embeddingss          r   embedzMultiModalBackend.embedW   s    " Q<#!11)<N  fd##00A #%*:*F"$''>;K*LST"U*&&'!!)## *    c                     |D cg c]  }| j                  |       }}| j                  j                  ||      }|S c c}w )a  Embed a list of n documents/words into an n-dimensional
        matrix of embeddings.

        Arguments:
            documents: A list of documents or words to be embedded
            verbose: Controls the verbosity of the process

        Returns:
            Document/words embeddings with shape (n, m) with `n` documents/words
            that each have an embeddings size of `m`
        show_progress_bar)_truncate_documentr   encode)r   r   r   doctruncated_docs
embeddingss         r   r!   z!MultiModalBackend.embed_documents}   sL     CLL3$11#6LL))00SZ0[
 Ms   >wordsc                 @    | j                   j                  ||      }|S )am  Embed a list of n words into an n-dimensional
        matrix of embeddings.

        Arguments:
            words: A list of words to be embedded
            verbose: Controls the verbosity of the process

        Returns:
            Document/words embeddings with shape (n, m) with `n` documents/words
            that each have an embeddings size of `m`
        r,   )r   r/   )r   r3   r   r2   s       r   embed_wordszMultiModalBackend.embed_words   s%     ))00'0R
r*   c                    | j                   rZt        t        j                  t	        |      | j                   z              }g }t        t        |      |       D ]  }|| j                   z  }|| j                   z  | j                   z   }||| D cg c])  }t        |t              rt        j                  |      n|+ }	}| j                  | j                  j                  |	      }
n| j                  j                  |	d      }
|j                  |
j                                t        |d   t              s|	D ]  }|j!                            t        j"                  |      }|S |D cg c]  }t        j                  |       }	}| j                  | j                  j                  |	      }|S | j                  j                  |	d      }|S c c}w c c}w )N)disableFr,   r   )r   intr$   ceillenr   ranger   r   r   openr   r/   r   extendtolistclosearray)r   r   r   nr_iterationsr2   istart_index	end_indeximageimages_to_embedimg_embfilepaths               r   r#   zMultiModalBackend.embed_images   s   ??Fdoo(E FGM J%.GD &$//10DOOC	 Y__jktXu#OTE3)?EJJu%UJ# # ##/"..55oFG"2299/]b9cG!!'.."23 fQi-!0 &&&" *-J  EKKuzz(3KOK+!--44_E
  "1188\a8b
)# Ls   .G,Gc                     | j                   rZ| j                   j                  |      }t        |      dkD  r1|dd }| j                   j                  |      }| j	                  |      S |S )NM      L   )r   r/   r:   decoder.   )r   documenttokenstruncated_tokenss       r   r.   z$MultiModalBackend._truncate_document   sb    >>^^**84F6{R#)!B< >>001AB ..x88r*   )N    )NF)F)__name__
__module____qualname____doc__r   r   r   r8   r   r   boolr$   ndarrayr)   r!   r5   r#   r.   __classcell__)r   s   @r   r
   r
   
   s     J 8<	("s$778(" 3 334(" 	("T$$tCy $$$s)d2B $$TX $$egeoeo $$Lc T bjj  c T bjj @r*   r
   )numpyr$   PILr   r   typingr   r   sentence_transformersr   bertopic.backendr   r
    r*   r   <module>r_      s%        5 )~ ~r*   