如何无需双重循环计算多文本间两两公共词元的数量?
如何高效计算多文本间公共词元数量的上三角矩阵?
我手里有大量已经完成分词的小型文本,想要计算所有文本对之间的公共词元数量,把结果存为上三角矩阵M(也可以是对称矩阵),其中M[i,j]代表文本i和文本j的公共词元数量。
目前我只能通过双重循环实现,但这种方法的时间效率特别低,尤其是文本数量多的时候。以下是我当前的实现代码:
from scipy.sparse import lil_matrix n = len(subjects) M = lil_matrix((n, n)) i = 0 for subj_1 in subjects: j = 0 for subj_2 in subjects[i+1:]: inter_len = len(list(set(subj_1).intersection(subj_2))) if inter_len>0: M[i,j+i+1] = inter_len j += 1 i += 1
注:subjects是存储各文本词元列表的列表,每个元素是一个文本的分词结果(词元组成的列表)。
有没有更高效的实现方式?
高效解决方案:利用稀疏矩阵乘法
这个问题我之前处理过!双重循环的问题在于时间复杂度是O(n²*k)(n是文本数量,k是单文本平均词元数),当n稍微大一点(比如上千个文本),速度会慢到难以接受。其实我们可以借助稀疏矩阵的矩阵乘法来大幅提升效率,核心思路是把文本的词元存在转化为one-hot的词-文档矩阵,再通过矩阵乘法直接批量计算所有文本对的公共词元数。
核心原理
两个文本的公共词元数量,本质上是它们的one-hot词向量的点积:如果词元在两个文本中都出现,点积就加1,最终总和就是公共词元数。而所有文本对的点积结果,正好是「词-文档矩阵的转置」乘以「词-文档矩阵」得到的对称矩阵,其中(i,j)位置的值就是文本i和j的公共词元数。
代码实现(推荐用sklearn快速构建矩阵)
from scipy.sparse import csr_matrix, triu from sklearn.feature_extraction.text import CountVectorizer # 先把每个词元列表转成空格分隔的字符串(CountVectorizer要求输入格式) texts = [' '.join(subj) for subj in subjects] # 构建one-hot词-文档矩阵(binary=True表示只记录词是否出现,不统计次数,对应原代码的set去重) vectorizer = CountVectorizer(binary=True) doc_term_matrix = vectorizer.fit_transform(texts) # 计算所有文本对的公共词元数:得到的是对称矩阵,(i,j)就是文本i和j的公共词元数 common_words_matrix = doc_term_matrix.T.dot(doc_term_matrix) # 提取上三角部分(k=1表示跳过对角线,和原代码的上三角范围一致) upper_tri_matrix = triu(common_words_matrix, k=1) # 如果你需要和原代码一样的lil_matrix格式,可以转换 upper_tri_lil = upper_tri_matrix.tolil()
手动构建稀疏矩阵(无需sklearn)
如果不想依赖sklearn,也可以手动构建词-文档矩阵:
from scipy.sparse import csr_matrix, triu from collections import defaultdict # 给每个词分配唯一ID word_to_id = defaultdict(int) current_id = 0 rows = [] cols = [] for subj_idx, subj in enumerate(subjects): # 先去重(和原代码的set操作对应) unique_words = set(subj) for word in unique_words: if word not in word_to_id: word_to_id[word] = current_id current_id += 1 rows.append(subj_idx) cols.append(word_to_id[word]) # 构建CSR格式的稀疏矩阵(CSR是最适合乘法的稀疏格式) doc_term_matrix = csr_matrix(([1]*len(rows), (rows, cols)), shape=(len(subjects), len(word_to_id))) # 后续步骤和上面一致 common_words_matrix = doc_term_matrix.T.dot(doc_term_matrix) upper_tri_matrix = triu(common_words_matrix, k=1)
为什么这个方法更快?
- 稀疏矩阵乘法是用底层优化过的算法实现的(比如C/C++加速),时间复杂度远低于双重循环;
- 批量处理所有文本对,避免了循环中重复的集合交集计算;
- CSR格式的稀疏矩阵内存占用小,适合处理大量文本的场景。
内容的提问来源于stack exchange,提问作者ThetaGamma
相关产品推荐
相关产品推荐

