NumPy矩阵乘法维度不匹配错误排查及相关代码问题咨询
NumPy矩阵乘法维度不匹配错误排查及相关代码问题咨询
各位大佬好,我最近在实现一个基于图拉普拉斯的优化算法时,频繁遇到NumPy矩阵乘法维度不匹配的报错,折腾了好一阵没搞定,想请大家帮忙看看问题出在哪。先把我的代码片段贴出来:
import numpy as np from numpy.linalg import inv from scipy.linalg import pinv # Define necessary functions def create_laplacian_from_adjacency(adj_matrix): degree_matrix = np.diag(adj_matrix.sum(axis=1)) laplacian_matrix = degree_matrix - adj_matrix return laplacian_matrix def dirichlet_energy(L, X): """Dirichlet energy for smoothness quantification.""" return np.trace(X.T @ L @ X) def update_C(X, X_tilde, L, C, gamma, alpha, lam, J): """Updating C using gradient descent with a majorized function approximation.""" p, k = C.shape C_old = np.copy(C) gradient_f = (-2 * gamma * L @ C_old @ inv(C_old.T @ L @ C_old + J) + alpha * (C_old @ X_tilde - X) @ X_tilde.T + 2 * L @ C_old @ X_tilde @ X_tilde.T + lam * C_old @ np.ones((k, k))) # Majorized function optimization step (simplified approach) t = 0.01 # Learning rate, needs tuning based on problem specifics # 后面的更新逻辑还没写完,但目前前面的梯度计算已经报维度不匹配了
我现在遇到的具体问题:
- 调用
update_C时,矩阵乘法的部分经常报错,比如L @ C_old、C_old.T @ L @ C_old这些操作,我不确定输入的各矩阵维度是否符合算法要求 - 想确认下
dirichlet_energy函数的实现逻辑有没有问题?毕竟这是后续优化的核心损失项之一 - 梯度计算里的各项维度是否匹配?比如
inv(C_old.T @ L @ C_old + J)这里,J应该设成什么维度的矩阵?我现在是随便给的单位矩阵,但不确定对不对
麻烦各位帮忙分析下,谢谢啦!
备注:内容来源于stack exchange,提问作者Aradhya_009
相关产品推荐
相关产品推荐

