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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.17 07:39:33