
    WiU>                        d dl Zd dlZd dlmZ d dlZd dlm	Z	m
Z
 	 	 	 	 	 	 	 	 	 	 	 	 dde	e   dej                  de	e   dz  dej                  dej                  de
eef   dz  d	ed
ededede
eef   dedededej"                  fdZy)    N)ListUniondocshierarchical_topicstopics
embeddingsreduced_embeddingssamplehide_annotationshide_document_hover	nr_levelslevel_scalecustom_labelstitlewidthheightreturnc                    | j                   }||dkD  rd}g }t        |      D ]  }t        j                  t        j                  |      |k(        d   }t        |      dk  rt        |      nt        t        |      |z        }|j                  t        j                  j                  ||d              t        j                  |      }t        j                  dt        j                  |      |   i      }|D cg c]  }||   	 c}|d<   |D cg c]  }||   	 c}|d<   |3|.|,| j                  |j                  j                         d	
      }n:|}n7|||   }n/|-|+| j                  |j                  j                         d	
      }|/	 ddlm}  |dddd      j#                        }|j$                  }n||||   }n|||}dddf   |d<   |dddf   |d<   |j*                  j                         }|
dk(  s|
dk(  rt        j,                  t        j.                  t1        j2                  dd      t1        j2                  t        |      dz
  d      |	            j5                  t              j7                         }|j9                          |D cg c]  }||   	 }}nX|
dk(  s|
dk(  rCt        j:                  t=        t        |            |	      D cg c]
  }||d       c}ddd   }nt?        d      tA        |      D ]x  \  }}|jB                  jE                         D ci c]  }|| }}|jF                  |j*                  |k  ddf   } | jH                  j5                  t              | _$        | jK                  d      } | jM                         D ](  }!|!d   jN                  D ]  }|!d   jH                  ||<    * |D "cg c]  }"d }#}"tQ        |#      rUtA        |jS                               D ]-  \  }\  }$}%|%|jU                         v r|$|%k7  r	||%   ||$<   )d|#|<   / tQ        |#      rU|jB                  jW                  |      |d|dz    <   |d|dz       j5                  t              |d|dz    <   { g }&i }'t=        |jH                  j5                  t              jY                               D ]  }||jH                  j5                  t              j[                         k  r| j]                  |      sFt_        |t`              r:| ddjc                  te        tg        | jh                  |   |          dd       z   }(nj| jj                  |r| jj                  || jl                  z      }(n?| ddjc                  | j]                  |      D )"cg c]
  \  })}"|)dd   c}"})dd       z   }(|(dd! |(dd! d"|'|<   |&jo                  |(       | d|jF                  |jH                  ta        |      k(  d#f   jp                  d   z   }(djc                  |(js                  d      dd D *cg c]  }*|*dd  	 c}*      }+|(dd! |+dd! d"|'|<   |&jo                  |(        g },t=        t        |            D ]  }-g }.| jl                  r|.jo                  tu        jv                  |jF                  |d|-dz       dk(  df   |jF                  |d|-dz       dk(  df   d$d%d&|s|jF                  |d|-dz       dk(  df   nddty        d'd(d)*      +             |rf|jF                  |jB                  j{                  |      ddf   } t}        | d|-dz       jE                         D cg c]  }t        |       c}      }/n9t}        |d|-dz       jE                         D cg c]  }t        |       c}      }/|/D ]  }|dk7  s
|r<|jF                  |d|-dz       |k(  |jB                  j{                  |      z  ddf   } n|jF                  |d|-dz       |k(  ddf   } |sd| jF                  t        |       ddf<   d,| d&<   | j~                  j                         | jF                  t        |       dz
  df<   | j                  j                         | jF                  t        |       dz
  df<   |'t        |         d-   | jF                  t        |       dz
  d&f<   |.jo                  tu        jv                  | j~                  | j                  |s| j                  nd|s| j                  ndd&|'t        |         d.   d$ty        d(d)/      0              |,jo                  |.        |,D .cg c]  }.t        |.       }0}.d|0d   fg}1tA        |0dd       D ]%  \  }}2|1|   d   }3|2|3z   }4|1jo                  |3|4f       ' tu        j                         }5|,D ]  }.|.D ]  }6|5j                  |6         t=        t        |5j                              D ]  }||0d   k\  sd|5j                  |   _F        ! g }7tA        |1      D ]t  \  }}ty        d1ta        |      d2dgt        |5j                        z  ig3      }8t=        |d   |d   z
        D ]  }d|8d4   d   d2   ||d   z   <    |7jo                  |8       v ty        d5d6id7d i|78      g}9|j~                  j[                         t        |j~                  j[                         d9z        z
  |j~                  jY                         t        |j~                  jY                         d9z        z   f}:|j                  j[                         t        |j                  j[                         d9z        z
  |j                  jY                         t        |j                  jY                         d9z        z   f};|5j                  d:t        |:      dz  |;d   t        |:      dz  |;d   ty        d'd;      <       |5j                  d:|:d   t        |;      dz  |:d   t        |;      dz  ty        d=d;      <       |5j                  |:d   t        |;      dz  d>dd?       |5j                  |;d   t        |:      dz  d@ddA       |5j                  |9dB| d)dCdDty        dEdFG      dH||I       |5j                  dJ       |5j                  dJ       |5S c c}w c c}w # t&        t(        f$ r t)        d      w xY wc c}w c c}w c c}w c c}"w c c}"})w c c}*w c c}w c c}w c c}.w )Ka  Visualize documents and their topics in 2D at different levels of hierarchy.

    Arguments:
        topic_model: A fitted BERTopic instance.
        docs: The documents you used when calling either `fit` or `fit_transform`
        hierarchical_topics: A dataframe that contains a hierarchy of topics
                             represented by their parents and their children
        topics: A selection of topics to visualize.
                Not to be confused with the topics that you get from `.fit_transform`.
                For example, if you want to visualize only topics 1 through 5:
                `topics = [1, 2, 3, 4, 5]`.
        embeddings: The embeddings of all documents in `docs`.
        reduced_embeddings: The 2D reduced embeddings of all documents in `docs`.
        sample: The percentage of documents in each topic that you would like to keep.
                Value can be between 0 and 1. Setting this value to, for example,
                0.1 (10% of documents in each topic) makes it easier to visualize
                millions of documents as a subset is chosen.
        hide_annotations: Hide the names of the traces on top of each cluster.
        hide_document_hover: Hide the content of the documents when hovering over
                             specific points. Helps to speed up generation of visualizations.
        nr_levels: The number of levels to be visualized in the hierarchy. First, the distances
                   in `hierarchical_topics.Distance` are split in `nr_levels` lists of distances.
                   Then, for each list of distances, the merged topics are selected that have a
                   distance less or equal to the maximum distance of the selected list of distances.
                   NOTE: To get all possible merged steps, make sure that `nr_levels` is equal to
                   the length of `hierarchical_topics`.
        level_scale: Whether to apply a linear or logarithmic (log) scale levels of the distance
                     vector. Linear scaling will perform an equal number of merges at each level
                     while logarithmic scaling will perform more mergers in earlier levels to
                     provide more resolution at higher levels (this can be used for when the number
                     of topics is large).
        custom_labels: If bool, whether to use custom topic labels that were defined using
                       `topic_model.set_topic_labels`.
                       If `str`, it uses labels from other aspects, e.g., "Aspect1".
                       NOTE: Custom labels are only generated for the original
                       un-merged topics.
        title: Title of the plot.
        width: The width of the figure.
        height: The height of the figure.

    Examples:
    To visualize the topics simply run:

    ```python
    topic_model.visualize_hierarchical_documents(docs, hierarchical_topics)
    ```

    Do note that this re-calculates the embeddings and reduces them to 2D.
    The advised and preferred pipeline for using this function is as follows:

    ```python
    from sklearn.datasets import fetch_20newsgroups
    from sentence_transformers import SentenceTransformer
    from bertopic import BERTopic
    from umap import UMAP

    # Prepare embeddings
    docs = fetch_20newsgroups(subset='all',  remove=('headers', 'footers', 'quotes'))['data']
    sentence_model = SentenceTransformer("all-MiniLM-L6-v2")
    embeddings = sentence_model.encode(docs, show_progress_bar=False)

    # Train BERTopic and extract hierarchical topics
    topic_model = BERTopic().fit(docs, embeddings)
    hierarchical_topics = topic_model.hierarchical_topics(docs)

    # Reduce dimensionality of embeddings, this step is optional
    # reduced_embeddings = UMAP(n_neighbors=10, n_components=2, min_dist=0.0, metric='cosine').fit_transform(embeddings)

    # Run the visualization with the original embeddings
    topic_model.visualize_hierarchical_documents(docs, hierarchical_topics, embeddings=embeddings)

    # Or, if you have reduced the original embeddings already:
    topic_model.visualize_hierarchical_documents(docs, hierarchical_topics, reduced_embeddings=reduced_embeddings)
    ```

    Or if you want to save the resulting figure:

    ```python
    fig = topic_model.visualize_hierarchical_documents(docs, hierarchical_topics, reduced_embeddings=reduced_embeddings)
    fig.write_html("path/to/file.html")
    ```

    Note:
        This visualization was inspired by the scatter plot representation of Doc2Map:
        https://github.com/louisgeisler/Doc2Map

    <iframe src="../../getting_started/visualization/hierarchical_documents.html"
    style="width:1000px; height: 770px; border: 0px;""></iframe>
    N   r   d   F)sizereplacetopicdocdocument)method)UMAP
      g        cosine)n_neighborsn_componentsmin_distmetricz{UMAP is required if the embeddings are not yet reduced in dimensionality. Please install it using `pip install umap-learn`.xyloglogarithmic)startstopnumlinlinearz0level_scale needs to be one of 'log' or 'linear'	Parent_IDTlevel__      (   )
trace_name	plot_textParent_Namezmarkers+textothertextz#CFD8DC   g      ?)colorr   opacity)r%   r&   modename	hoverinfo	hovertext
showlegendmarker r6   r5   )r   r<   )r%   r&   r9   r@   r?   r>   r=   rB   updatevisible)r   labelargsrG   prefixzLevel: t)currentvaluepadstepsg333333?line)r;   r   )typex0y0x1y1rM   z#9E9E9ED1)r%   r&   r9   	showarrowyshiftD2)r&   r%   r9   rT   xshiftsimple_whitecentertop   Black)r   r;   )r9   r%   xanchoryanchorfont)sliderstemplater   r   r   )rE   )Ntopics_setnpwherearraylenintextendrandomchoicepd	DataFrame_extract_embeddingsr   to_listumapr   fit
embedding_ImportErrorModuleNotFoundErrorDistanceroundlogspacemathr'   astypetolistreversearray_splitrange
ValueError	enumerater   uniquelocr/   sort_valuesiterrowsTopicsanyitemskeysmapmaxmin	get_topic
isinstancestrjoinnextziptopic_aspects_custom_labels_	_outliersappendvaluessplitgo	Scattergldictisinsortedr%   meanr&   r9   Figure	add_tracedatarE   abs	add_shapesumadd_annotationupdate_layoutupdate_xaxesupdate_yaxes)<topic_modelr   r   r   r   r	   r
   r   r   r   r   r   r   r   r   topic_per_docindicesr   sr   dfindexembeddings_to_reducer   
umap_modelembeddings_2d	distanceslog_indicesimax_distancesmax_distancemapping	selectionrowr1   mappingskeyvaluetrace_namestopic_namesr5   wordr>   r6   
all_tracesleveltracesunique_topicsnr_traces_per_settrace_indices	nr_tracesr)   endfigtracerL   stepr`   x_rangey_ranges<                                                               u/home/sietch6/trending-topics-pipeline/venv/lib/python3.12/site-packages/bertopic/plotting/_hierarchical_documents.py visualize_hierarchical_documentsr   	   st   T  ''M ~!G]# FHHRXXm,56q9Q#s1v3s1v+?ryy''e'DEF hhwG	w 7 @A	BB*12e2BuI5<=E='=BwK ~"4"<#.#B#B266>>CS\f#B#g #- !#-g#6 $6$>#.#B#B266>>CS\f#B#g  !	!"1sS[\``auvJ&11M
 
	 2 >*73	.:* AqD!BsGAqD!BsG $,,446Ie{m;HH((1b/#i.1"4b9! VC[VX 	 	/:;!1;;		!824..sK^G_A`bk2l
'.Igbk"

B$ KLL(7 H|-/XX__->?E5%<??'++,?,H,HL,XZ[,[\	'1188=	))+6	%%' 	2CQ 2!$Q!1!12	2
 #**QD**(m#,W]]_#= (<CGLLN*se|#*5>GCL"'HQK	( (m $&88<<#8VEAI; #%uqyk&:#;#B#B3#GVEAI; -H2 KK*44;;C@DDFG +&0077<@@BB$$U+mS1$)7!sxxS+"<"<]"KE"RSTUWVWX0 "J !//;!,!;!;EKDYDY<Y!ZJ$)7!sxxR]RgRghmRn8owtQcr8oprqr8s/t!tJ",Sb/!+CR&E" "":. '%))*=*G*G3u:*UWd*dellmnop  
8H8H8Mbq8Q!R$s)!RSI("o&s^"K z*5+: Js=)* 4"   MMffb6%!)!56"<sBCffR& 45;S@A' $Uhbffb6%!))=&>"&Du%LMnr$iaE	 f 59:I"IuWXykFZ<[<b<b<d#e5CJ#efM"BPQ	{?S<T<[<[<]#^5CJ#^_M" 	E{ "VEAI;+?(@E(Ibhhmm\bNc'dfg'g hI "rF519+*>'?5'H!'K LI'7;IMM#i.!"34(*If%=F[[=M=M=OIMM#i.1"4c"9:=F[[=M=M=OIMM#i.1"4c"9:@KCPUJ@WXc@dIMM#i.1"4f"<=LL#++#++3CY^^7J)--PT"((U4\B+#C8		6 	&!i4"n 4>>V>>*1-./M%&7&;< +ye$Q'%eS\*+ ))+C ! 	!EMM% 	!! s388}% ,%a((&+CHHUO#,
 E#M2 we*ugCHH567

 71:
23 	BE=ADLOI&uwqz'9:	BT (I!6S"IUSTG 	
S"$$((*,--

S"$$((*,--G
 	
S"$$((*,--

S"$$((*,--G MMw<!1:w<!1:	+   MM1:w<!1:w<!	+   s7|a'7de\^_s7|a'7de\^_ gb0
    U#U#Jm 3=* 01 	% N 	8 <
 @ +2 9p "S< $f#^@ ?sN   =s6s;-t  "t%t#7
t(	t-(t2 t8
8t=
2u
>u t)NNNNFTr   r-   Fz(<b>Hierarchical Documents and Topics</b>i  i  )numpyrd   pandasrl   plotly.graph_objectsgraph_objectsr   rx   typingr   r   r   rm   rh   ndarrayfloatboolr   r        r   <module>r      s     !    $!%)'+" $&+;n
s)n n I	n
 

n 

n %*$n n n n n s#n n n n  YY!nr   