从布尔值Pandas DataFrame提取线性无关列子集(复刻Matlab licols)
提取布尔值DataFrame的线性无关列(Matlab licols函数的Python实现)
问题背景
我有一个形状为(n,m)的Pandas DataFrame,其中n为观测数,m为特征数。所有特征均为布尔值(0,1),由分类变量独热编码生成。我需要提取该矩阵的子集,移除其中的线性相关列(即可以表示为其他列线性组合的列)。
我希望实现与以下Matlab licols 函数完全一致的解决方案:
function [Xsub,idx]=licols(X,tol) %Extract a linearly independent set of columns of a given matrix X % % [Xsub,idx]=licols(X) % in: % % X: The given input matrix % tol: A rank estimation tolerance. Default=1e-10 % out: % % Xsub: The extracted columns of X % idx: The indices (into X) of the extracted columns if ~nnz(X) %X has no non-zeros and hence no independent columns Xsub=[]; idx=[]; return end if nargin<2, tol=1e-10; end [Q, R, E] = qr(X,0); if ~isvector(R) diagr = abs(diag(R)); else diagr = abs(R(1)); end %Rank estimation r = find(diagr >= tol*diagr(1), 1, 'last'); %rank estimation idx=sort(E(1:r)); Xsub=X(:,idx);
我尝试在Python 3.9.x版本的Jupyter Lab中调用该Matlab函数,但Matlab引擎不支持此Python版本,因此需要用Python实现等效逻辑。
Python等效实现
以下是完全匹配Matlab licols 逻辑的Python代码,使用numpy和pandas实现:
import numpy as np import pandas as pd def licols(X, tol=1e-10): # 转换为numpy矩阵(兼容DataFrame或numpy数组输入) if isinstance(X, pd.DataFrame): mat = X.values cols = X.columns else: mat = X cols = None # 检查矩阵是否全零 if np.count_nonzero(mat) == 0: return (pd.DataFrame(), []) if cols is not None else (np.array([]), []) # 执行带列置换的经济版QR分解(对应Matlab的qr(X,0)) Q, R, E = np.linalg.qr(mat, mode='economic', pivoting=True) # 获取R的对角线绝对值 diagr = np.abs(R[0]) if R.ndim == 1 else np.abs(np.diag(R)) # 估计矩阵的秩:找到最后一个满足阈值条件的元素索引 mask = diagr >= tol * diagr[0] r = np.max(np.where(mask)[0]) if np.any(mask) else 0 # 获取原始列索引并排序 idx = E[:r+1] idx_sorted = np.sort(idx) # 提取线性无关列并返回对应格式 if cols is not None: Xsub = X.iloc[:, idx_sorted] idx_sorted = cols[idx_sorted].tolist() else: Xsub = mat[:, idx_sorted] return Xsub, idx_sorted
使用示例
# 创建包含线性相关列的布尔值DataFrame data = { 'col0': [1,0,1,0], 'col1': [0,1,0,1], 'col2': [1,1,1,1], # col2 = col0 + col1,线性相关 'col3': [1,0,0,1] } df = pd.DataFrame(data) # 提取线性无关列 df_sub, idx = licols(df) print("提取的列索引:", idx) print("提取的DataFrame:") print(df_sub)
运行结果:
提取的列索引: ['col0', 'col1', 'col3'] 提取的DataFrame: col0 col1 col3 0 1 0 1 1 0 1 0 2 1 0 0 3 0 1 1
内容的提问来源于stack exchange,提问作者NC520
相关产品推荐
相关产品推荐

