求解LDA Gibbs采样代码中的逻辑错误问题
LDA Gibbs采样代码逻辑错误排查
在实现LDA Gibbs采样时,运行结果与标准Gibbs采样结果不符,核心问题出在_gibbs_sampling函数的概率计算逻辑上,以下是具体问题和修正方案:
原代码片段
class LDAGibbs: def __init__(self, num_topics, doc_file_path, vocas, output_dir_name, alpha=0.1, beta=0.01): """ Constructor method Do not modify this function. :param num_topics: the number of topics :param doc_file_path: BOW document file path :param vocas: vocabulary list :param output_dir_name: output directory name :param alpha: alpha value in LDA :param beta: beta value in LDA :return: void """ self.docs = self.read_bow(doc_file_path) self.words = vocas self.K = num_topics self.D = len(self.docs) self.W = len(vocas) self.output_dir_name = output_dir_name # Hyper-parameters # We use symmetric hyper-parameters self.alpha = alpha self.beta = beta # self.WK: Words by Topics matrix # self.DK: Documents by Topics matrix self.WK = np.zeros([self.W, self.K]) self.DK = np.zeros([self.D, self.K]) # Random initialization of topics np.random.seed(108) # Topic index of each word in each document self.doc_topics = list() for di in range(self.D): doc = self.docs[di] topics = np.random.randint(self.K, size=len(doc)) self.doc_topics.append(topics) for wi in range(len(doc)): topic = topics[wi] word = doc[wi] self.WK[word, topic] += 1 self.DK[di, topic] += 1 def run(self, max_iter=2000, do_print_log=False): """ Run Collapsed Gibbs sampling for LDA Do not modify this function. :param max_iter: Maximum number of Gibbs sampling iterations :param do_print_log: Print loglikelihood and run time :return: void """ if do_print_log: prev = time.clock() for iteration in range(max_iter): print(iteration, time.clock() - prev, self.loglikelihood()) prev = time.clock() self._gibbs_sampling() if iteration % 100 == 99: self.export_result(output_file_name="iter_{}".format(iteration)) else: for iteration in range(max_iter): self._gibbs_sampling() if iteration % 100 == 99: self.export_result(output_file_name="iter_{}".format(iteration)) def _gibbs_sampling(self): for di in range(self.D): doc = self.docs[di] for wi in range(len(doc)): word = doc[wi] old_topic = self.doc_topics[di][wi] self.DK[di, old_topic] -= 1 self.WK[word, old_topic] -= 1 prob_vec = self.DK[di] + self.alpha prob_item = self.WK[word] + self.beta prob_vec /= np.sum(self.WK + self.beta, axis=0) new_topic = self._sampling_from_dist(prob_vec) self.doc_topics[di][wi] = new_topic self.DK[di, new_topic] += 1 self.WK[word, new_topic] += 1
核心问题分析
- 未使用词-主题概率项:代码中定义了
prob_item = self.WK[word] + self.beta,但完全没有将其与文档-主题项相乘,违反了LDA Gibbs采样的核心公式——每个主题的概率是文档主题分布项和词主题分布项的乘积。 - 归一化逻辑错误:
prob_vec /= np.sum(self.WK + self.beta, axis=0)的操作不符合公式要求,正确做法是先计算两个项的乘积,再对整个概率向量做归一化(除以所有主题概率的总和)。
修正后的_gibbs_sampling函数
def _gibbs_sampling(self): for di in range(self.D): doc = self.docs[di] for wi in range(len(doc)): word = doc[wi] old_topic = self.doc_topics[di][wi] # 移除当前词的旧主题计数 self.DK[di, old_topic] -= 1 self.WK[word, old_topic] -= 1 # 计算每个主题的概率:(文档d中主题k的计数 + alpha) * (词w在主题k中的计数 + beta) doc_topic_term = self.DK[di] + self.alpha word_topic_term = self.WK[word] + self.beta prob_vec = doc_topic_term * word_topic_term # 归一化概率向量 prob_vec /= np.sum(prob_vec) # 采样新主题 new_topic = self._sampling_from_dist(prob_vec) # 更新计数和主题分配 self.doc_topics[di][wi] = new_topic self.DK[di, new_topic] += 1 self.WK[word, new_topic] += 1
补充说明
修正后的代码严格遵循LDA Collapsed Gibbs采样公式:对每个词,计算其分配到各个主题的概率时,结合当前文档的主题分布(含alpha平滑)和当前词的主题分布(含beta平滑),再做归一化后采样。若_sampling_from_dist是基于概率向量的多项式采样实现(如np.random.choice),修正后的逻辑会得到符合预期的结果。
内容的提问来源于stack exchange,提问作者Inzamamul Alam
相关产品推荐
相关产品推荐

