You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

求解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

核心问题分析

  1. 未使用词-主题概率项:代码中定义了prob_item = self.WK[word] + self.beta,但完全没有将其与文档-主题项相乘,违反了LDA Gibbs采样的核心公式——每个主题的概率是文档主题分布项和词主题分布项的乘积。
  2. 归一化逻辑错误: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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.19 15:27:53