训练NMF后测试集重建遇形状不兼容问题求助
问题根源分析
你的核心问题是NMF实现的维度定义不符合标准用法,导致训练得到的基向量无法适配测试集:
- 你输入的
X_train是(4000, 784)(样本数×特征数),但原nmf函数中初始化的W是(4000, m),这是把基向量定义在了样本空间而非特征空间。 - 测试集
X_test是(800, 784),样本数和训练集不同,自然无法和样本空间的基向量做矩阵运算,导致形状不兼容报错。
标准NMF的正确维度逻辑应该是:
- 输入矩阵转置为
(特征数, 样本数),比如MNIST数据转置后为(784, 4000) - 分解得到
W(特征数, m)(m个特征空间的基向量)和H(m, 样本数)(每个样本对应的系数) - 固定训练得到的
W,求解测试集对应的H_test,即可完成测试集重建
修正后的NMF实现
以下是调整维度逻辑后的代码,适配(样本数, 特征数)格式的输入:
import numpy as np def cost(X, W, H): """ 计算NMF的欧氏距离损失 X: (n_samples, n_features), W: (n_features, m), H: (m, n_samples) """ X_t = X.T diff = X_t - np.dot(W, H) cost_value = (diff * diff).sum() / (X_t.shape[0] * X_t.shape[1]) return cost_value def nmf(X, m): """ 对输入矩阵X(n_samples, n_features)执行NMF分解 返回:重建后的X(与输入形状一致)、特征基向量W、迭代损失列表 """ n_samples, n_features = X.shape X_t = X.T # 转置为(特征数, 样本数) # 初始化正确维度的W和H W = np.random.rand(n_features, m) # (特征数, 分量数) H = np.random.rand(m, n_samples) # (分量数, 样本数) cost_values = [] pseudo_count = 0.0001 for i in range(100): # 固定W,更新H numerator_H = W.T.dot(X_t) denominator_H = W.T.dot(W.dot(H)) + pseudo_count H = H * (numerator_H / denominator_H) # 固定H,更新W numerator_W = X_t.dot(H.T) denominator_W = W.dot(H.dot(H.T)) + pseudo_count W = W * (numerator_W / denominator_W) cost_values.append(cost(X, W, H)) # 重建与输入形状一致的矩阵 X_nmf = (np.dot(W, H)).T return X_nmf, W, cost_values
测试集重建方案
修正NMF后,训练得到的nmf_basis_vectors是(784, m)的特征基向量,可通过两种方式求解测试集的系数并重建:
方法1:非负最小二乘(保证系数非负)
需要依赖scipy库,对每个测试样本单独求解非负系数:
from scipy.optimize import nnls m = 10 X_train_nmf, nmf_basis_vectors, nmf_cost_values = nmf(X_train, m) # 初始化测试集系数矩阵 H_test = np.zeros((m, X_test.shape[0])) # 为每个测试样本求解非负系数 for idx in range(X_test.shape[0]): sample = X_test[idx, :] H_test[:, idx], _ = nnls(nmf_basis_vectors, sample) # 重建测试集 reconstructed_test_nmf = (nmf_basis_vectors @ H_test).T
方法2:NMF迭代更新规则(无需额外库)
复用NMF的H更新规则,固定训练得到的W,迭代优化测试集的H:
m = 10 X_train_nmf, nmf_basis_vectors, nmf_cost_values = nmf(X_train, m) X_test_t = X_test.T # 随机初始化测试集系数 H_test = np.random.rand(m, X_test.shape[0]) pseudo_count = 0.0001 # 迭代更新H_test(固定W) for i in range(50): numerator = nmf_basis_vectors.T.dot(X_test_t) denominator = nmf_basis_vectors.T.dot(nmf_basis_vectors.dot(H_test)) + pseudo_count H_test = H_test * (numerator / denominator) # 重建测试集 reconstructed_test_nmf = (nmf_basis_vectors @ H_test).T
内容的提问来源于stack exchange,提问作者laikaka
相关产品推荐
相关产品推荐

