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

训练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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 08:12:38