From 0ea895bc84c7ed1869063f3145ffc1ec03e39abc Mon Sep 17 00:00:00 2001 From: z00512249 Date: Tue, 22 Mar 2022 09:35:27 +0800 Subject: [PATCH] add gmres dense bprop --- .../python/mindspore/scipy/sparse/linalg.py | 125 +++++++++--------- tests/st/scipy_st/sparse/test_linalg.py | 93 ++++++++++++- 2 files changed, 154 insertions(+), 64 deletions(-) diff --git a/mindspore/python/mindspore/scipy/sparse/linalg.py b/mindspore/python/mindspore/scipy/sparse/linalg.py index 0dc8a127929..ebdf28a1b3f 100644 --- a/mindspore/python/mindspore/scipy/sparse/linalg.py +++ b/mindspore/python/mindspore/scipy/sparse/linalg.py @@ -72,7 +72,7 @@ def _high_precision_cho_solve(a, b, data_type=mstype.float64): return y.astype(data_type) -def _batch_gmres(A, x0, b, tol, atol, restart, maxiter, M): +def _batch_gmres(A, b, x0, tol, restart, maxiter, M, atol): """ batched gmres: solve the least squares problem from scratch at the end of each GMRES iteration. It does not allow for early termination, but has much less overhead on GPUs. @@ -106,7 +106,7 @@ def _batch_gmres(A, x0, b, tol, atol, restart, maxiter, M): return x, F.select(residual_norm > atol, k, _INT_ZERO) -def _incremental_gmres(A, x0, b, tol, atol, restart, maxiter, M): +def _incremental_gmres(A, b, x0, tol, restart, maxiter, M, atol): """ incremental gmres: builds a QR decomposition for the Krylov subspace incrementally during the GMRES process using Givens rotations. This improves numerical stability and gives a free estimate of @@ -139,7 +139,7 @@ def _incremental_gmres(A, x0, b, tol, atol, restart, maxiter, M): while mnp.logical_and(mnp.less(k, restart), mnp.less(ptol, err)): V, R, _ = arnoldi_iteration(k, A, M, V, R) # Givens rotation - row_k = R[k, :].copy() + row_k = R[k, :] i = _INT_ZERO while i < k: row_k = rotate_vectors(row_k, i, givens[i, 0], givens[i, 1]) @@ -185,16 +185,16 @@ class GMRES(nn.Cell): self.M = M self.solve_method = solve_method - def construct(self, b, x0, tol, atol, restart, maxiter): + def construct(self, b, x0, tol, restart, maxiter, atol): # Constant tensor which avoids loop unrolling A = _normalize_matvec(self.A) M = _normalize_matvec(self.M) x = x0 info = _to_tensor(0) if self.solve_method == 'batched': - x, info = _batch_gmres(A, x0, b, tol, atol, restart, maxiter, M) + x, info = _batch_gmres(A, b, x0, tol, restart, maxiter, M, atol) elif self.solve_method == "incremental": - x, info = _incremental_gmres(A, x0, b, tol, atol, restart, maxiter, M) + x, info = _incremental_gmres(A, b, x0, tol, restart, maxiter, M, atol) else: _raise_value_error("solve_method should be in ('incremental' or 'batched'), but got ", self.solve_method, ".") @@ -210,20 +210,38 @@ class GMRESV2(nn.Cell): super(GMRESV2, self).__init__() self.solve_method = solve_method - def construct(self, A, b, x0, tol, atol, restart, maxiter, M): + def construct(self, A, b, x0, tol, restart, maxiter, M, atol): A = _normalize_matvec(A) M = _normalize_matvec(M) x = x0 info = _to_tensor(0) if self.solve_method == 'batched': - x, info = _batch_gmres(A, x0, b, tol, atol, restart, maxiter, M) + x, info = _batch_gmres(A, b, x0, tol, restart, maxiter, M, atol) elif self.solve_method == "incremental": - x, info = _incremental_gmres(A, x0, b, tol, atol, restart, maxiter, M) + x, info = _incremental_gmres(A, b, x0, tol, restart, maxiter, M, atol) else: _raise_value_error("solve_method should be in ('incremental' or 'batched'), but got ", self.solve_method, ".") return x, info + def bprop(self, A, b, x0, tol, restart, maxiter, M, atol, out, dout): + """ + Derivatives of `gmres` are implemented via implicit differentiation with + another `gmres` solve, rather than by differentiating *through* the solver. + They will be accurate only if both solves converge. + """ + n = b.shape[0] + if not isinstance(M, (Tensor, CSRTensor)): + M = F.eye(n, n, b.dtype) + grad_b, _ = self.construct(A.T, dout[0], x0, tol, restart, maxiter, M, atol) + if isinstance(A, CSRTensor): + grad_a_dense = -1 * F.reshape(grad_b, (n, 1)) * F.reshape(out[0], (1, n)) + values = F.csr_gather(A.indptr, A.indices, grad_a_dense, A.shape) + grad_a = CSRTensor(A.indptr, A.indices, values, A.shape) + else: + grad_a = -1 * F.reshape(grad_b, (n, 1)) * F.reshape(out[0], (1, n)) + return grad_a, grad_b, zeros_like(x0), zeros_like(tol), zeros_like(atol), zeros_like(maxiter), zeros_like(M) + def gmres(A, b, x0=None, *, tol=1e-5, restart=20, maxiter=None, M=None, callback=None, restrt=None, atol=0.0, callback_type=None, solve_method='batched'): @@ -322,12 +340,43 @@ def gmres(A, b, x0=None, *, tol=1e-5, restart=20, maxiter=None, if restart > size: restart = size if not is_within_graph(A): - x, info = GMRES(A, M, solve_method)(b, x0, tol, atol, restart, maxiter) + x, info = GMRES(A, M, solve_method)(b, x0, tol, restart, maxiter, atol) else: - x, info = GMRESV2(solve_method)(A, b, x0, tol, atol, restart, maxiter, M) + x, info = GMRESV2(solve_method)(A, b, x0, tol, restart, maxiter, M, atol) return x, info +def _cg(A, b, x0, tol, atol, maxiter, M): + """ + Figure 2.5 from Barrett R, et al. 'Templates for the sulution of linear systems: + building blocks for iterative methods', 1994, pg. 12-14 + """ + # Constant tensor which avoids loop unrolling + _INT_ZERO = _to_tensor(0) + atol_ = mnp.maximum(atol, tol * _norm(b)) + + r = b - A(x0) + z = p = M(r) + rho = mnp.dot(r, z) + k = _INT_ZERO + x = x0 + while k < maxiter and _norm(r) > atol_: + q = A(p) + alpha = rho / mnp.dot(p, q) + x = x + alpha * p + r = r - alpha * q + + z = M(r) + rho_ = mnp.dot(r, z) + beta = rho_ / rho + p = z + beta * p + rho = rho_ + + k += 1 + + return x, F.select(_norm(r) > atol_, k, _INT_ZERO) + + class CG(nn.Cell): """Figure 2.5 from Barrett R, et al. 'Templates for the sulution of linear systems: building blocks for iterative methods', 1994, pg. 12-14 @@ -339,34 +388,9 @@ class CG(nn.Cell): self.M = M def construct(self, b, x0, tol, atol, maxiter): - # Constant tensor which avoids loop unrolling - _INT_ZERO = _to_tensor(0) - A = _normalize_matvec(self.A) M = _normalize_matvec(self.M) - - atol_ = mnp.maximum(atol, tol * _norm(b)) - - r = b - A(x0) - z = p = M(r) - rho = mnp.dot(r, z) - k = _INT_ZERO - x = x0 - while k < maxiter and _norm(r) > atol_: - q = A(p) - alpha = rho / mnp.dot(p, q) - x = x + alpha * p - r = r - alpha * q - - z = M(r) - rho_ = mnp.dot(r, z) - beta = rho_ / rho - p = z + beta * p - rho = rho_ - - k += 1 - - return x, F.select(_norm(r) > atol_, k, _INT_ZERO) + return _cg(A, b, x0, tol, atol, maxiter, M) class CGv2(nn.Cell): @@ -378,34 +402,9 @@ class CGv2(nn.Cell): super(CGv2, self).__init__() def construct(self, A, b, x0, tol, atol, maxiter, M): - # Constant tensor which avoids loop unrolling - _INT_ZERO = _to_tensor(0) - A = _normalize_matvec(A) M = _normalize_matvec(M) - - atol_ = mnp.maximum(atol, tol * _norm(b)) - - r = b - A(x0) - z = p = M(r) - rho = mnp.dot(r, z) - k = _INT_ZERO - x = x0 - while k < maxiter and _norm(r) > atol_: - q = A(p) - alpha = rho / mnp.dot(p, q) - x = x + alpha * p - r = r - alpha * q - - z = M(r) - rho_ = mnp.dot(r, z) - beta = rho_ / rho - p = z + beta * p - rho = rho_ - - k += 1 - - return x, F.select(_norm(r) > atol_, k, _INT_ZERO) + return _cg(A, b, x0, tol, atol, maxiter, M) def bprop(self, A, b, x0, tol, atol, maxiter, M, out, dout): """ diff --git a/tests/st/scipy_st/sparse/test_linalg.py b/tests/st/scipy_st/sparse/test_linalg.py index 2975ebbe5d2..2c199840e0b 100644 --- a/tests/st/scipy_st/sparse/test_linalg.py +++ b/tests/st/scipy_st/sparse/test_linalg.py @@ -432,12 +432,103 @@ def test_gmres_against_graph_scipy(n, tensor_type, dtype, error, preconditioner, ms_output, _ = msp.sparse.linalg.gmres(a, b, x0, tol=tol, restart=restart, maxiter=maxiter, M=m, atol=atol) assert onp.allclose(scipy_output, ms_output.asnumpy(), rtol=error, atol=error) - # With in graph's construct ms_net_output, _ = TestNet(solve_method)(a, b, x0, tol, restart, maxiter, m, atol) assert onp.allclose(scipy_output, ms_net_output.asnumpy(), rtol=error, atol=error) +@pytest.mark.level0 +@pytest.mark.platform_x86_cpu +@pytest.mark.platform_x86_gpu_training +@pytest.mark.env_onecard +@pytest.mark.parametrize('tensor_type, dtype, error', [('Tensor', onp.float64, 1e-5), ('Tensor', onp.float32, 1e-4), + ('CSRTensor', onp.float32, 1e-4)]) +@pytest.mark.parametrize('preconditioner', ['identity', 'exact', 'random']) +@pytest.mark.parametrize('solve_method', ['incremental', 'batched']) +@pytest.mark.parametrize('a, b, grad_a, grad_b', [ + ([[1.96822833, 0.82204467, 1.03749232, 0.88915326, 0.44986806, 1.11167143], + [0.82204467, 2.25216591, 1.40235719, 0.70838919, 0.81377919, 1.06000368], + [1.03749232, 1.40235719, 2.90618746, 0.7126087, 0.81029544, 1.28673025], + [0.88915326, 0.70838919, 0.7126087, 2.17515263, 0.40443765, 1.02082996], + [0.44986806, 0.81377919, 0.81029544, 0.40443765, 1.60570668, 0.62292701], + [1.11167143, 1.06000368, 1.28673025, 1.02082996, 0.62292701, 2.30795277]], + [0.79363745, 0.58000418, 0.1622986, 0.70075235, 0.96455108, 0.50000836], + [[-0.07867674, -0.01521201, 0.06394698, -0.03854052, -0.13523701, 0.01326866], + [-0.03508505, -0.00678363, 0.02851647, -0.01718673, -0.06030749, 0.00591702], + [-0.00586019, -0.00113306, 0.00476305, -0.00287067, -0.01007304, 0.00098831], + [-0.07704304, -0.01489613, 0.06261914, -0.03774023, -0.13242886, 0.01299314], + [-0.14497008, -0.02802971, 0.11782896, -0.07101491, -0.24918826, 0.02444888], + [-0.01868565, -0.00361284, 0.01518735, -0.00915334, -0.03211867, 0.00315129]], + [0.22853142, 0.10191113, 0.01702201, 0.22378603, 0.42109291, 0.054276]), + ([[1.85910724, 0.73233206, 0.65960803, 1.03821349, 0.55277616], + [0.73233206, 1.69548841, 0.59992146, 1.01518264, 0.50824059], + [0.65960803, 0.59992146, 1.98169091, 1.45565213, 0.47901749], + [1.03821349, 1.01518264, 1.45565213, 3.3133049, 0.75598147], + [0.55277616, 0.50824059, 0.47901749, 0.75598147, 1.46831254]], + [0.59674531, 0.226012, 0.10694568, 0.22030621, 0.34982629], + [[-0.07498642, 0.00167461, 0.01353184, 0.01008293, -0.03770084], + [-0.09940184, 0.00221986, 0.01793778, 0.01336592, -0.04997616], + [-0.09572781, 0.00213781, 0.01727477, 0.01287189, -0.04812897], + [0.03135044, -0.00070012, -0.00565741, -0.00421549, 0.01576203], + [-0.14053766, 0.00313851, 0.02536103, 0.01889718, -0.07065797]], + [0.23398106, 0.31016481, 0.29870068, -0.09782316, 0.43852141]), +]) +def test_gmres_grad(tensor_type, dtype, error, preconditioner, solve_method, a, b, grad_a, grad_b): + """ + Feature: ALL TO ALL + Description: test cases for gmres grad [N x N] X [N X 1] + Expectation: the result match jax grad + """ + if not _is_valid_platform(tensor_type): + return + + # Input CSRTensor of gmres grad in mindspore graph or pynative mode is not supported, just ignored it. + # Root cause: CSRTensor has no distribute function of T. + if tensor_type == "CSRTensor": + return + + # Gmres grad in construct + class GmresGradNet(nn.Cell): + def __init__(self, solve_method): + super(GmresGradNet, self).__init__() + self.sum = ops.ReduceSum() + self.gmres = msp.sparse.linalg.gmres + self.solve_method = solve_method + + def construct(self, a, b, x0, tol, m, atol): + # For restart && maxiter args, we maintain default values to ensure gmres can coverage. + x, _ = self.gmres(a, b, x0, tol=tol, M=m, atol=atol, + solve_method=self.solve_method) + return self.sum(x) + + gmres_grad_net = ops.GradOperation(get_all=True)(GmresGradNet(solve_method)) + + tol = float(onp.finfo(dtype=dtype).eps) + atol = tol + + a = onp.array(a, dtype=dtype) + b = onp.array(b, dtype=dtype) + x0 = onp.zeros_like(b).astype(dtype) + m = _fetch_preconditioner(preconditioner, a) + expect_grad_a = grad_a + expect_grad_b = grad_b + a = to_tensor((a, tensor_type)) + b = Tensor(b) + x0 = Tensor(x0) + m = to_tensor((m, tensor_type)) if m is not None else m + # PyNative Mode + context.set_context(mode=context.PYNATIVE_MODE) + grad_a, grad_b = gmres_grad_net(a, b, x0, tol, m, atol)[:2] + assert onp.allclose(expect_grad_a, to_ndarray(grad_a), rtol=error, atol=error) + assert onp.allclose(expect_grad_b, to_ndarray(grad_b), rtol=error, atol=error) + # Graph Mode + context.set_context(mode=context.GRAPH_MODE) + + grad_a, grad_b = gmres_grad_net(a, b, x0, tol, m, atol)[:2] + assert onp.allclose(expect_grad_a, to_ndarray(grad_a), rtol=error, atol=error) + assert onp.allclose(expect_grad_b, to_ndarray(grad_b), rtol=error, atol=error) + + @pytest.mark.level0 @pytest.mark.platform_x86_gpu_training @pytest.mark.platform_x86_cpu