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

矩阵链乘法实现中matrix_mult函数触发段错误的求助

矩阵链乘法项目段错误排查求助

我正在开发一个基于树结构实现矩阵链乘法的项目,核心节点与矩阵结构体定义如下:

struct Node
{
   char* seq;
   Node* left;
   Node* right;
};

struct matrix
{
   std::tuple<int, int> dimension;
   std::vector<std::vector<int>> values;
};

矩阵指针存入std::unordered_map<char, matrix*>,并传入Sequence类的私有成员。相关初始化代码:

matrix *A1 = new matrix;
matrix *A2 = new matrix;
A1->dimension = std::make_tuple(4,10);
A2->dimension = std::make_tuple(10,3);

setValues(A1, 1);
setValues(A2, 2);

Node *n = new Node;
n->left = nullptr;
n->right = nullptr;
n->seq = static_cast<char*>(malloc(3 * sizeof(char)));
n->seq[0] = '0' + 1;
n->seq[1] = '0' + 2;
n->seq[2] = '\0';

void setValues(matrix *x,int value)
{
   int row = std::get<0>(x->dimension);
   int col = std::get<1>(x->dimension);
   x->values.resize(row);
   for (int i = 0; i < row; ++i) 
   {
      x->values[i].resize(col);
   }
   for(int i = 0; i < row; i++)
   {
      for(int j = 0; j < col; j++)
      {
         x->values[i][j] = value;
      }
   }
}

std::unordered_map<char, matrix*> dict;
dict['1'] = A1;
dict['2'] = A2;

class Sequence
{
private:
    Node root;
    std::vector<std::vector<int>> s_table;
    std::unordered_map<char,matrix*> str_matrix_dict;
public:
    Sequence(std::vector<std::vector<int>> temp_table, std::unordered_map<char,matrix*> &str_matrix_dict);
    matrix* compute(Node* n);
    void matrix_mult(matrix *a, matrix *b, matrix *c, int x, int y, int z); // 补充成员函数声明
};

Sequence::Sequence(
std::vector<std::vector<int>> temp_table,
std::unordered_map<char,matrix*> &temp_dict) : s_table(temp_table), str_matrix_dict(temp_dict)
{}

Sequence seq(s_table, dict);

在递归compute函数中调用matrix_mult执行矩阵乘法时,首次进入三重循环就触发段错误,且compute中打印矩阵值无输出。核心问题代码:

matrix* Sequence::compute(Node* n)
{
/*
    code before
*/
if(n->left == nullptr && n->right == nullptr && n->seq[2] == '\0')
{
    matrix* matrix_A = str_matrix_dict[n->seq[0]];
    matrix* matrix_B = str_matrix_dict[n->seq[1]];
    matrix* matrix_C = new matrix;
    int m = std::get<0>(matrix_A->dimension);
    int n_dim = std::get<1>(matrix_B->dimension); // 重命名避免与参数n冲突
    int z = std::get<1>(matrix_A->dimension);
    matrix_C->dimension = std::make_tuple(m,n_dim);
    matrix_C->values.resize(m);
    for (int i = 0; i < m; ++i) 
    {
        matrix_C->values[i].resize(n_dim);
    }
    matrix_mult(matrix_A,matrix_B,matrix_C,m,n_dim,z); // 修复参数传递错误
    return matrix_C;
}
matrix* left_res = compute(n->left);
matrix* right_res = compute(n->right);
//code after (same computation as above)
}

void Sequence::matrix_mult(matrix *a, matrix *b, matrix *c, int x, int y, int z)
{
std::cout << "calculating " << std::endl;
// 修复循环变量逻辑错误
for(int row = 0; row < x; row++)
{
    for(int col = 0; col < y; col++)
    {
        for(int k = 0; k < z; k++)
        {
            c->values[row][col] += (a->values[row][k] * b->values[k][col]);
        }
    }
}
}

项目使用nvcc编译,编译脚本:

CXX = nvcc
CPP = gcc
CFLAGS = -std=c++11 -g
LDFLAGS =  -lcudart -lcudadevrt

main: main.o sequence.o
	$(CXX) $(CFLAGS) -o main main.o sequence.o $(LDFLAGS)

main.o: main.cu
	$(CXX) $(CFLAGS) -c main.cu
sequence.o: sequence.cpp
	$(CXX) $(CFLAGS) -c sequence.cpp

clean:
	rm -f main main.o sequence.o

我怀疑问题与引用或解引用有关,但无法准确定位,恳请解决建议与指导。


核心问题修复建议

  • matrix_mult循环逻辑错误:
    第一个循环的增量误写为x++,应改为row++;第二个循环的条件误写为row < y、增量为y++,应改为col < y和col++。这两处错误会导致循环无限执行或直接越界访问内存,触发段错误。
  • compute函数参数传递错误:
    调用matrix_mult时传入了未定义的变量x,实际应传入中间维度z;同时变量名n与函数参数Node* n冲突,建议重命名为n_dim避免歧义。
  • 其他潜在问题:
    Node的seq使用malloc分配内存,后续需调用free避免内存泄漏;判断n->seq[2] == '\0'存在越界风险,建议先通过strlen(n->seq)检查字符串长度;Sequence类需补充matrix_mult的成员函数声明,否则编译会报错。

内容的提问来源于stack exchange,提问作者samueltru99

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 06:38:11