1- """Module for utility functions for fitting BTMs"""
1+ """Module for utility functions for fitting BTMs. """
22
33import random
44from typing import Dict , Tuple , TypeVar
55
66import numba
77import numpy as np
88from numba import njit
9+
910from tweetopic ._prob import norm_prob , sample_categorical
1011
1112
1213@njit
1314def doc_unique_biterms (
14- doc_unique_words : np .ndarray , doc_unique_word_counts : np .ndarray
15+ doc_unique_words : np .ndarray ,
16+ doc_unique_word_counts : np .ndarray ,
1517) -> Dict [Tuple [int , int ], int ]:
1618 (n_max_unique_words ,) = doc_unique_words .shape
1719 biterm_counts = dict ()
@@ -42,7 +44,7 @@ def doc_unique_biterms(
4244
4345@njit
4446def nb_add_counter (dest : Dict [T , int ], source : Dict [T , int ]):
45- """Adds one counter dict to another in place with Numba"""
47+ """Adds one counter dict to another in place with Numba. """
4648 for key in source :
4749 if key in dest :
4850 dest [key ] += source [key ]
@@ -52,25 +54,28 @@ def nb_add_counter(dest: Dict[T, int], source: Dict[T, int]):
5254
5355@njit
5456def corpus_unique_biterms (
55- doc_unique_words : np .ndarray , doc_unique_word_counts : np .ndarray
57+ doc_unique_words : np .ndarray ,
58+ doc_unique_word_counts : np .ndarray ,
5659) -> Dict [Tuple [int , int ], int ]:
5760 n_documents , _ = doc_unique_words .shape
5861 biterm_counts = doc_unique_biterms (
59- doc_unique_words [0 ], doc_unique_word_counts [0 ]
62+ doc_unique_words [0 ],
63+ doc_unique_word_counts [0 ],
6064 )
6165 for i_doc in range (1 , n_documents ):
6266 doc_unique_words_i = doc_unique_words [i_doc ]
6367 doc_unique_word_counts_i = doc_unique_word_counts [i_doc ]
6468 doc_biterms = doc_unique_biterms (
65- doc_unique_words_i , doc_unique_word_counts_i
69+ doc_unique_words_i ,
70+ doc_unique_word_counts_i ,
6671 )
6772 nb_add_counter (biterm_counts , doc_biterms )
6873 return biterm_counts
6974
7075
7176@njit
7277def compute_biterm_set (
73- biterm_counts : Dict [Tuple [int , int ], int ]
78+ biterm_counts : Dict [Tuple [int , int ], int ],
7479) -> np .ndarray :
7580 return np .array (list (biterm_counts .keys ()))
7681
@@ -115,7 +120,12 @@ def add_biterm(
115120 topic_biterm_count : np .ndarray ,
116121) -> None :
117122 add_remove_biterm (
118- True , i_biterm , i_topic , biterms , topic_word_count , topic_biterm_count
123+ True ,
124+ i_biterm ,
125+ i_topic ,
126+ biterms ,
127+ topic_word_count ,
128+ topic_biterm_count ,
119129 )
120130
121131
@@ -128,7 +138,12 @@ def remove_biterm(
128138 topic_biterm_count : np .ndarray ,
129139) -> None :
130140 add_remove_biterm (
131- False , i_biterm , i_topic , biterms , topic_word_count , topic_biterm_count
141+ False ,
142+ i_biterm ,
143+ i_topic ,
144+ biterms ,
145+ topic_word_count ,
146+ topic_biterm_count ,
132147 )
133148
134149
@@ -146,7 +161,11 @@ def init_components(
146161 i_topic = random .randint (0 , n_components - 1 )
147162 biterm_topic_assignments [i_biterm ] = i_topic
148163 add_biterm (
149- i_biterm , i_topic , biterms , topic_word_count , topic_biterm_count
164+ i_biterm ,
165+ i_topic ,
166+ biterms ,
167+ topic_word_count ,
168+ topic_biterm_count ,
150169 )
151170 return biterm_topic_assignments , topic_word_count , topic_biterm_count
152171
@@ -360,7 +379,10 @@ def predict_docs(
360379 )
361380 biterms = doc_unique_biterms (words , word_counts )
362381 prob_topic_given_document (
363- pred , biterms , topic_distribution , topic_word_distribution
382+ pred ,
383+ biterms ,
384+ topic_distribution ,
385+ topic_word_distribution ,
364386 )
365387 predictions [i_doc , :] = pred
366388 return predictions
0 commit comments