矩阵链乘法实现中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
相关产品推荐
相关产品推荐

