已合并
matrix_inverse算子的golden代码 #830
matrix_inverse算子的golden代码 #830
已合并
dx创建于 7月7日
共 2 个文件变更+192-1
@@ -1,5 +1,5 @@
1/**1/**
2- * Copyright (c) 2025 Huawei Technologies Co., Ltd.2+ * Copyright (c) 2025-2026 Huawei Technologies Co., Ltd.
3 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of3 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
S
Ssunhao_hw7月23日

已有文件若进行修改,Copyright (c)改为2025-2026,而不是直接改成2026

likedislike
dx
dx
7月23日 评论:
4 * CANN Open Software License Agreement Version 2.0 (the "License").4 * CANN Open Software License Agreement Version 2.0 (the "License").
5 * Please refer to the License for details. You may not use this file except in compliance with the License.5 * Please refer to the License for details. You may not use this file except in compliance with the License.
@@ -15,5 +15,6 @@
15#include "golden/fill_data.hpp"15#include "golden/fill_data.hpp"
16#include "golden/matmul.hpp"16#include "golden/matmul.hpp"
17#include "golden/conv2d.hpp"17#include "golden/conv2d.hpp"
18+#include "golden/matrix_inverse.hpp"
18 19 
19#endif // EXAMPLES_COMMON_GOLDEN_HPP20#endif // EXAMPLES_COMMON_GOLDEN_HPP
@@ -0,0 +1,190 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+#ifndef EXAMPLES_COMMON_GOLDEN_MATRIX_INVERSE_HPP
12+#define EXAMPLES_COMMON_GOLDEN_MATRIX_INVERSE_HPP
13+ 
14+#include <cmath>
15+#include <limits>
16+#include <utility>
17+#include <vector>
atomgit-bot
atomgit-botatomgit-bot7月7日

🟠 High Priority

变更行:第 25、77、157 行使用了 int32_t,但文件中只包含了 <cmath>、<limits>、<utility>、<vector>,没有包含 <cstdint>(或 <stdint.h>)。虽然某些标准库实现会通过 <vector> 等头文件间接引入 int32_t,但 C++ 标准并不保证这一点。在不保证传递包含的平台上(例如某些嵌入式工具链或严格模式的编译器),这将导致编译失败。

失败模式:编译错误 —— int32_t 未声明。 严重程度:该文件通过 golden.hpp 被 30+ 个示例文件间接包含,影响面较广。

建议:在现有头文件包含之后添加 #include <cstdint>,以保证 int32_t 在所有平台上均可用。

改动建议
17
+ #include <cmath>
18
+ #include <cstdint>
19
+ #include <limits>
20
+ #include <utility>
17
21
  #include <vector>
应用建议
likedislike
不准确?
18+#include <cstdint>
19+ 
20+namespace Catlass::golden {
21+ 
22+// Row-major LU decomposition with partial pivoting (analogous to LAPACK sgetrf).
23+// A is N x N in row-major; ipiv receives pivot indices (0-based).
24+// Returns 0 on success, or the row index (+1) where a zero pivot was encountered.
25+template <class Element>
26+int SgetrfRowMajor(int N, std::vector<Element>& A, std::vector<int32_t>& ipiv)
27+{
28+ ipiv.resize(N);
29+ 
30+ // Pivot magnitude below this threshold is treated as zero (numerically singular).
31+ constexpr Element kPivotThreshold = Element(64) * std::numeric_limits<Element>::epsilon();
32+ 
33+ for (int k = 0; k < N; ++k) {
34+ // Find pivot
35+ Element maxVal = std::fabs(A[k * N + k]);
36+ int maxRow = k;
37+ for (int i = k + 1; i < N; ++i) {
38+ Element absVal = std::fabs(A[i * N + k]);
39+ if (absVal > maxVal) {
40+ maxVal = absVal;
41+ maxRow = i;
42+ }
43+ }
44+ ipiv[k] = maxRow;
45+ 
46+ if (maxVal <= kPivotThreshold) {
47+ // Numerically singular matrix
48+ return k + 1; // 1-based info
49+ }
50+ 
51+ // Swap rows k and maxRow
52+ if (maxRow != k) {
53+ for (int j = 0; j < N; ++j) {
54+ std::swap(A[k * N + j], A[maxRow * N + j]);
55+ }
56+ }
57+ 
58+ // Compute multipliers and update trailing submatrix
59+ Element invPivot = Element(1) / A[k * N + k];
60+ for (int i = k + 1; i < N; ++i) {
61+ A[i * N + k] *= invPivot; // store L factor
62+ Element factor = A[i * N + k];
63+ for (int j = k + 1; j < N; ++j) {
64+ A[i * N + j] -= factor * A[k * N + j];
65+ }
66+ }
67+ }
68+ 
69+ return 0;
70+}
71+ 
72+// Solve A * X = B using LU factors and pivot info from SgetrfRowMajor.
73+// A contains LU factors (unit lower L, upper U) in row-major.
74+// B is N x nrhs in row-major; result overwrites B.
75+// Returns 0 on success, or -1 if `trans` is a null pointer (invalid argument).
76+template <class Element>
77+int SgetrsRowMajor(
78+ const char* trans, int N, int nrhs, const std::vector<Element>& A, const std::vector<int32_t>& ipiv,
79+ std::vector<Element>& B)
80+{
81+ // Runtime error handling: a null `trans` is an invalid argument and must be
82+ // rejected explicitly (not via assert, which is stripped in release builds).
83+ if (trans == nullptr) {
84+ return -1;
85+ }
86+ bool notrans = (trans[0] == 'N' || trans[0] == 'n');
87+ 
88+ if (notrans) {
89+ // Solve A * X = B
90+ // Step 1: Apply pivots to B (P * B)
91+ for (int k = 0; k < N; ++k) {
92+ int pivRow = ipiv[k];
93+ if (pivRow != k) {
94+ for (int j = 0; j < nrhs; ++j) {
95+ std::swap(B[k * nrhs + j], B[pivRow * nrhs + j]);
96+ }
97+ }
98+ }
99+ 
100+ // Step 2: Forward substitution L * Y = B (L is unit lower triangular)
101+ for (int k = 0; k < N; ++k) {
102+ for (int i = k + 1; i < N; ++i) {
103+ Element factor = A[i * N + k];
104+ for (int j = 0; j < nrhs; ++j) {
105+ B[i * nrhs + j] -= factor * B[k * nrhs + j];
106+ }
107+ }
108+ }
109+ 
110+ // Step 3: Back substitution U * X = Y
111+ for (int k = N - 1; k >= 0; --k) {
112+ Element invDiag = Element(1) / A[k * N + k];
113+ for (int j = 0; j < nrhs; ++j) {
114+ B[k * nrhs + j] *= invDiag;
115+ }
116+ for (int i = 0; i < k; ++i) {
117+ Element factor = A[i * N + k];
118+ for (int j = 0; j < nrhs; ++j) {
119+ B[i * nrhs + j] -= factor * B[k * nrhs + j];
120+ }
121+ }
122+ }
123+ } else {
124+ // Solve A^T * X = B (unused for inverse, but included for completeness)
125+ // Step 1: Forward with U^T
126+ for (int k = 0; k < N; ++k) {
127+ for (int j = 0; j < nrhs; ++j) {
128+ for (int i = 0; i < k; ++i) {
129+ B[k * nrhs + j] -= A[i * N + k] * B[i * nrhs + j];
130+ }
131+ B[k * nrhs + j] /= A[k * N + k];
132+ }
133+ }
134+ 
135+ // Step 2: Backward with L^T
136+ for (int k = N - 1; k >= 0; --k) {
137+ for (int j = 0; j < nrhs; ++j) {
138+ for (int i = k + 1; i < N; ++i) {
139+ B[k * nrhs + j] -= A[i * N + k] * B[i * nrhs + j];
140+ }
141+ }
142+ }
143+ 
144+ // Step 3: Apply inverse pivots
145+ for (int k = N - 1; k >= 0; --k) {
146+ int pivRow = ipiv[k];
147+ if (pivRow != k) {
148+ for (int j = 0; j < nrhs; ++j) {
149+ std::swap(B[k * nrhs + j], B[pivRow * nrhs + j]);
150+ }
151+ }
152+ }
153+ }
154+ return 0;
155+}
156+ 
157+// Compute the inverse of an N x N matrix in row-major layout.
158+// The input matrix is overwritten with its inverse.
159+// Returns 0 on success, or a positive value if the matrix is singular.
160+template <class Element>
161+int ComputeInverseInplace(int N, std::vector<Element>& A)
162+{
163+ // Step 1: LU decomposition with partial pivoting
164+ std::vector<int32_t> ipiv;
165+ int info = SgetrfRowMajor(N, A, ipiv);
166+ if (info != 0) {
167+ return info;
168+ }
169+ 
170+ // Step 2: Form the N x N identity matrix in a work buffer
171+ std::vector<Element> work(N * N, Element(0));
172+ for (int i = 0; i < N; ++i) {
173+ work[i * N + i] = Element(1);
174+ }
175+ 
176+ // Step 3: Solve A * X = I => X = A^{-1}
177+ int solveInfo = SgetrsRowMajor("N", N, N, A, ipiv, work);
178+ if (solveInfo != 0) {
179+ return solveInfo;
180+ }
181+ 
182+ // Step 4: Copy result back to A
183+ A = std::move(work);
184+ 
185+ return 0;
186+}
187+ 
188+} // namespace Catlass::golden
189+ 
190+#endif // EXAMPLES_COMMON_GOLDEN_MATRIX_INVERSE_HPP