
    Wi%                         d dl Zd dlmZ d dlmZ d dlmZ d dlmZ d dl	m
Z
 d dlmZ  G d d	e      Z G d
 de      Zy)    N)tqdm)List)Dataset)	normalize)Pipeline)BaseEmbedderc                        e Zd ZdZdef fdZddee   dede	j                  fdZded	e	j                  de	j                  fd
Z xZS )HFTransformerBackenda  Hugging Face transformers model.

    This uses the `transformers.pipelines.pipeline` to define and create
    a feature generation pipeline from which embeddings can be extracted.

    Arguments:
        embedding_model: A Hugging Face feature extraction pipeline

    Examples:
    To use a Hugging Face transformers model, load in a pipeline and point
    to any model found on their model hub (https://huggingface.co/models):

    ```python
    from bertopic.backend import HFTransformerBackend
    from transformers.pipelines import pipeline

    hf_model = pipeline("feature-extraction", model="distilbert-base-cased")
    embedding_model = HFTransformerBackend(hf_model)
    ```
    embedding_modelc                 f    t         |           t        |t              r|| _        y t        d      )NzPlease select a correct transformers pipeline. For example: pipeline('feature-extraction', model='distilbert-base-cased', device=0))super__init__
isinstancer   r   
ValueError)selfr   	__class__s     l/home/sietch6/trending-topics-pipeline/venv/lib/python3.12/site-packages/bertopic/backend/_hftransformers.pyr   zHFTransformerBackend.__init__"   s3    ox0#2D Z     	documentsverbosereturnc           
          t        |      }g }t        t        || j                  |dd            t	        |      |       D ]&  \  }}|j                  | j                  ||             ( t        j                  |      S )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`
        T)
truncationpadding)totaldisable)		MyDatasetr   zipr   lenappend_embednparray)r   r   r   dataset
embeddingsdocumentfeaturess          r   embedzHFTransformerBackend.embed-   s     I&
"&	4//DRV/WXg,K#
 	?Hh
 dkk(H=>	? xx
##r   r&   r'   c                    t        j                  |      }| j                  j                  |ddd      d   }t        j                  t        j
                  |d      |j                        }t        j                  ||z  d      }t        j                  |j                  d      d|j                  d      j                               }t        ||z        d	   }|S )
a&  Mean pooling.

        Arguments:
            document: The document for which to extract the attention mask
            features: The embeddings for each token

        Adopted from:
        https://huggingface.co/sentence-transformers/all-MiniLM-L12-v2#usage-huggingface-transformers
        Tr"   )r   r   return_tensorsattention_mask   g&.>)a_mina_maxr   )r"   r#   r   	tokenizerbroadcast_toexpand_dimsshapesumclipmaxr   )	r   r&   r'   token_embeddingsr+   input_mask_expandedsum_embeddingssum_mask	embeddings	            r   r!   zHFTransformerBackend._embedE   s     88H---77T[_pt7u
 !oobnn^R.PRbRhRhi 03F FJ77##A&%))!,002

 nx78;	r   )F)__name__
__module____qualname____doc__r   r   r   strboolr"   ndarrayr(   r!   __classcell__)r   s   @r   r
   r
      sX    *	 	$tCy $4 $BJJ $0s bjj RZZ r   r
   c                   "    e Zd ZdZd Zd Zd Zy)r   z5Dataset to pass to `transformers.pipelines.pipeline`.c                     || _         y Ndocs)r   rH   s     r   r   zMyDataset.__init__a   s	    	r   c                 ,    t        | j                        S rF   )r   rH   )r   s    r   __len__zMyDataset.__len__d   s    499~r   c                      | j                   |   S rF   rG   )r   idxs     r   __getitem__zMyDataset.__getitem__g   s    yy~r   N)r<   r=   r>   r?   r   rJ   rM    r   r   r   r   ^   s    ?r   r   )numpyr"   r   typingr   torch.utils.datar   sklearn.preprocessingr   transformers.pipelinesr   bertopic.backendr   r
   r   rN   r   r   <module>rU      s5       $ + + )O< Od
 
r   