Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 6 additions & 18 deletions optimism/SparseCholesky.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@

import numpy as onp

from sksparse.cholmod import analyze, cholesky
from sksparse import cholmod
from sksparse.cholmod import CholmodNotPositiveDefiniteError as NotPosDefError

from optimism.JaxConfig import *
Expand All @@ -21,15 +21,15 @@ def factorize(self):
# we can improve this later if we are inclined
assert isspmatrix_csc(self.A), \
"Preconditioner matrix is not in a valid sparse format"
self.Precond = analyze(self.A, mode='supernodal',
ordering_method='nesdis')
self.Precond = cholmod.CholeskyFactor(self.A, sym_kind="sym", supernodal_mode="supernodal",
order="metis")

attempt = 0
maxAttempts = 10
while attempt < maxAttempts:
try:
print('Factorizing preconditioner')
self.Precond.cholesky_inplace(self.A)
self.Precond.factorize(self.A)
except NotPosDefError:
attempt += 1
print('Cholesky failed, assembling preconditioner', attempt)
Expand All @@ -41,7 +41,7 @@ def factorize(self):
if attempt == maxAttempts:
print("Cholesky failed too many times, using identity preconditioner")
self.A = identity(self.A.shape[0], format='csc')
self.Precond.cholesky_inplace(self.A)
self.Precond.factorize(self.A)


def update(self, new_stiffness_func):
Expand All @@ -50,9 +50,7 @@ def update(self, new_stiffness_func):


def apply(self, b):
if type(b) == type(np.array([])):
b = onp.array(b, copy=False)
return self.Precond(b)
return np.asarray(self.Precond.solve(onp.array(b, copy=True)))


def apply_transpose(self, b):
Expand All @@ -67,16 +65,6 @@ def multiply_by_transpose(self, x):
return self.A.T.dot(x)


def check_stability(self, x, p):
A = self.stiffness_func(x, p)
try:
self.Precond.cholesky(A)
print("Jacobian is stable.")
except NotPosDefError as e:
print(e)
print("Jacobian is unstable.")


def get_diagonal_stiffness(self):
return self.A.diagonal()

Expand Down
Loading