基于Pandas DataFrame构建带行列索引映射的Scipy稀疏矩阵实现方案咨询
基于Pandas DataFrame构建带行列索引映射的Scipy稀疏矩阵实现方案咨询
嗨,看起来你已经迈出了不错的第一步!你的核心需求其实是构建带元数据(行列标签映射)的稀疏矩阵,同时保证矩阵运算的效率,下面给你几个实用的方向和优化方案:
一、保留行列标签与索引的映射关系
你当前的index函数已经实现了标签到索引的映射,但更好的方式是把映射关系单独存储下来,方便后续反向查找(比如根据矩阵的行号找到对应的BranchNumber)。这里推荐用Pandas的Categorical类型,它底层做了优化,比手动写字典映射更高效:
import pandas as pd from scipy.sparse import csr_matrix # 将分支和型号列转为分类类型,自动生成编码并保留原始标签 df['BranchNumber_cat'] = pd.Categorical(df['BranchNumber']) df['ModelArticleNumber_cat'] = pd.Categorical(df['ModelArticleNumber']) # 获取行列的原始标签集合(对应矩阵的行/列含义) row_labels = df['BranchNumber_cat'].categories col_labels = df['ModelArticleNumber_cat'].categories # 获取每个数据点对应的行/列索引编码 row_indices = df['BranchNumber_cat'].codes col_indices = df['ModelArticleNumber_cat'].codes
这样后续你想查矩阵第i行对应的分支号,直接用row_labels[i]即可;想查某个分支对应的行索引,用row_labels.get_loc(branch_num)就行,列的操作同理。
二、高效构建目标稀疏矩阵
基于上面的编码结果,构建CSR矩阵时最好显式指定shape参数,确保矩阵是你需要的463×5235尺寸(避免因数据中存在未出现的分支/型号导致维度不符):
sparse_mat = csr_matrix( (df['ActualSellingPrice'], (row_indices, col_indices)), shape=(len(row_labels), len(col_labels)) )
CSR格式本身非常适合后续的转置、矩阵乘法等操作,Scipy已经对这些运算做了深度优化,完全能支持你需要的sparse_mat.T.dot(sparse_mat)这类操作。
三、矩阵运算结果与标签关联
当你完成矩阵运算后,如果需要把结果和原始标签对应起来,可以将稀疏矩阵转为Pandas DataFrame(如果内存允许的话):
# 计算转置乘原矩阵(型号间的关联矩阵) mat_product = sparse_mat.T.dot(sparse_mat) # 将结果转为带标签的DataFrame product_df = pd.DataFrame( mat_product.todense(), index=col_labels, columns=col_labels )
如果矩阵太大(比如后续运算结果维度更高),不想转为稠密矩阵,也可以通过row_labels和col_labels直接关联索引与标签,按需查询稀疏矩阵中的元素。
四、封装成复用工具(可选)
如果这类操作你会频繁使用,可以把逻辑封装成一个类,让使用更便捷:
class LabeledSparseMatrix: def __init__(self, df, row_col, col_col, value_col): self.row_cat = pd.Categorical(df[row_col]) self.col_cat = pd.Categorical(df[col_col]) self.row_labels = self.row_cat.categories self.col_labels = self.col_cat.categories self.matrix = csr_matrix( (df[value_col], (self.row_cat.codes, self.col_cat.codes)), shape=(len(self.row_labels), len(self.col_labels)) ) # 根据索引获取标签 def get_row_label(self, idx): return self.row_labels[idx] def get_col_label(self, idx): return self.col_labels[idx] # 根据标签获取索引 def get_row_index(self, label): return self.row_labels.get_loc(label) def get_col_index(self, label): return self.col_labels.get_loc(label)
使用示例:
# 初始化带标签的稀疏矩阵 labeled_mat = LabeledSparseMatrix(df, 'BranchNumber', 'ModelArticleNumber', 'ActualSellingPrice') # 查询分支号123对应的行索引 labeled_mat.get_row_index(123) # 执行矩阵运算 product_matrix = labeled_mat.matrix.T.dot(labeled_mat.matrix)
备注:内容来源于stack exchange,提问作者Scott Deerwester
相关产品推荐
相关产品推荐

