"""JAX implementation of Density Fitting / ISDF."""
import jax
import jax.numpy as jnp
import jax.scipy.linalg as jsp_linalg
from jax.sharding import Mesh, NamedSharding, PartitionSpec as P
import numpy as np
from functools import partial
import os
import logging
import time
import h5py
import uuid
import gc
from jax.tree_util import Partial as JaxPartial
from .pivots import (
grad_columns,
grad_diagonal,
phi_columns,
phi_diagonal,
pivoted_cholesky_streaming,
validate_pivot_controls,
)
logger = logging.getLogger(__name__)
@jax.jit
def solve_normal_equations_batch(phi_piv_p: jnp.ndarray, phi_piv_q: jnp.ndarray,
phi_p_batch: jnp.ndarray, phi_q_batch: jnp.ndarray,
rcond: float = 1e-14) -> jnp.ndarray:
"""Fast solver using LU decomposition for structured least-squares.
Solves min ||C*X - B||² where:
C[pq, m] = phi_piv_p[p, m] * phi_piv_q[q, m]
B[pq, g] = phi_p_batch[p, g] * phi_q_batch[q, g]
This exploits the separable structure to avoid O(N^2) intermediates:
(C^T B)[m, g] = (sum_p phi_piv_p[p,m]*phi_p_batch[p,g]) * (sum_q phi_piv_q[q,m]*phi_q_batch[q,g])
Args:
phi_piv_p: (n_orb, n_fused) first factor of pivots
phi_piv_q: (n_orb, n_fused) second factor of pivots
phi_p_batch: (n_orb, batch_size) first factor of target
phi_q_batch: (n_orb, batch_size) second factor of target
rcond: Relative regularization strength (default 1e-14)
Returns:
X: (n_fused, batch_size) solutions
"""
# Compute A^T A efficiently using the Kronecker-like structure
gram_p = phi_piv_p.T @ phi_piv_p # (n_fused, n_fused)
gram_q = phi_piv_q.T @ phi_piv_q # (n_fused, n_fused)
ATA = gram_p * gram_q
# Compute A^T B efficiently using separable structure
term_p = jnp.matmul(phi_piv_p.T, phi_p_batch) # (n_fused, batch_size)
term_q = jnp.matmul(phi_piv_q.T, phi_q_batch) # (n_fused, batch_size)
ATB = term_p * term_q # (n_fused, batch_size)
# Use LU solve (jnp.linalg.solve) with Tikhonov regularization
diag_mean = jnp.mean(jnp.diag(ATA))
jitter = diag_mean * rcond
ATA_reg = ATA + jitter * jnp.eye(ATA.shape[0])
X = jnp.linalg.solve(ATA_reg, ATB)
return X
solve_normal_equations_batch = jax.jit(solve_normal_equations_batch, static_argnames=['rcond'])
@jax.jit
def _build_normal_matrix(phi_piv_p: jnp.ndarray, phi_piv_q: jnp.ndarray) -> jnp.ndarray:
"""Build unregularized normal-equation matrix for structured LS."""
gram_p = phi_piv_p.T @ phi_piv_p
gram_q = phi_piv_q.T @ phi_piv_q
return gram_p * gram_q
[docs]
def prepare_normal_equations_solver(phi_piv_p: jnp.ndarray, phi_piv_q: jnp.ndarray,
rcond: float = 1e-14,
max_jitter_tries: int = 8,
jitter_growth: float = 10.0):
"""Prepare robust Cholesky factor for repeated batched solves.
Uses adaptive jitter escalation to guarantee numerically SPD matrices.
"""
ata = _build_normal_matrix(phi_piv_p, phi_piv_q)
ata = 0.5 * (ata + ata.T)
diag_mean = float(jnp.mean(jnp.diag(ata)))
eps_scale = float(jnp.finfo(ata.dtype).eps) * max(diag_mean, 1.0)
base_jitter = max(diag_mean * rcond, eps_scale)
eye = jnp.eye(ata.shape[0], dtype=ata.dtype)
last_chol = None
for attempt in range(max_jitter_tries):
jitter = base_jitter * (jitter_growth ** attempt)
chol, lower = jsp_linalg.cho_factor(ata + jitter * eye, lower=True)
if bool(jnp.all(jnp.isfinite(chol))):
if attempt > 0:
logger.warning(
"Cholesky jitter escalated: base=%.3e final=%.3e tries=%d",
base_jitter, jitter, attempt + 1
)
return chol, bool(lower)
last_chol = chol
raise np.linalg.LinAlgError(
f"Adaptive Cholesky failed after {max_jitter_tries} tries; "
f"base_jitter={base_jitter:.3e}, last_nonfinite={bool(jnp.any(jnp.isnan(last_chol)))}"
)
@partial(jax.jit, static_argnames=('lower',))
def solve_normal_equations_batch_prepared(chol: jnp.ndarray, lower: bool,
phi_piv_p: jnp.ndarray, phi_piv_q: jnp.ndarray,
phi_p_batch: jnp.ndarray, phi_q_batch: jnp.ndarray) -> jnp.ndarray:
"""Solve batched normal equations using precomputed Cholesky factor."""
term_p = jnp.matmul(phi_piv_p.T, phi_p_batch)
term_q = jnp.matmul(phi_piv_q.T, phi_q_batch)
atb = term_p * term_q
return jsp_linalg.cho_solve((chol, lower), atb)
[docs]
def _pivoted_cholesky_phi(phi_weighted, n_rank, shift):
"""Historical exact-greedy density selector (batch-size-one contract)."""
return pivoted_cholesky_streaming(
phi_diagonal(phi_weighted),
JaxPartial(phi_columns, phi_weighted),
shift,
n_rank=n_rank,
batch_size=1,
candidate_oversampling=1,
n_topup=0,
)
[docs]
def _pivoted_cholesky_grad(phi_weighted, grad_phi_weighted, n_rank, shift):
"""Historical exact-greedy gradient selector (batch-size-one contract)."""
return pivoted_cholesky_streaming(
grad_diagonal(phi_weighted, grad_phi_weighted),
JaxPartial(grad_columns, phi_weighted, grad_phi_weighted),
shift,
n_rank=n_rank,
batch_size=1,
candidate_oversampling=1,
n_topup=0,
)
[docs]
def isdf_decompose(phi, grad_phi, n_rank_phi, n_rank_grad, weights=None,
grid_batch_size=4096, rcond=1e-14,
is_incore=False, save_path=None, fixed_pivots=None,
batch_size=1, candidate_oversampling=2, n_topup=0):
"""Perform ISDF decomposition of orbitals and their gradients.
Memory-efficient implementation using pivoted Cholesky and normal-equation
Cholesky solves to avoid materializing large C matrices (n_orb² × n_fused).
Args:
phi: Orbitals on grid (n_orb, n_grid)
grad_phi: Orbital gradients on grid (n_orb, n_grid, 3)
n_rank_phi: Rank for phi decomposition
n_rank_grad: Rank for gradient decomposition
weights: Optional (n_grid,) array of integration weights.
If provided, pivot selection is weighted by these weights.
grid_batch_size: Number of grid points to process in each batch
rcond: Scales the Tikhonov jitter in prepare_normal_equations_solver
(default 1e-14).
batch_size: Exact columns retained per blocked pivot round. The
default of one reproduces the historical greedy selector.
candidate_oversampling: Size multiplier for the stale-diagonal pool
that is exactly re-pivoted before each blocked update.
n_topup: Number of final pivots selected by exact greedy singleton
updates after the blocked rounds.
Returns:
phi_piv: (N_orb, N_fused)
xi_phi: (N_fused, N_grid)
grad_phi_piv: (N_orb, N_fused, 3)
xi_grad: (N_fused, N_grid, 3)
pivots: (N_fused,)
"""
n_grid = int(phi.shape[1])
validate_pivot_controls(
n_grid, n_rank_phi, batch_size, candidate_oversampling, n_topup
)
validate_pivot_controls(
n_grid, n_rank_grad, batch_size, candidate_oversampling, n_topup
)
if save_path is not None and os.path.exists(save_path) and fixed_pivots is None:
try:
with h5py.File(save_path, 'r') as f:
if all(k in f for k in ['xi_phi', 'xi_grad', 'pivots', 'phi_isdf', 'grad_phi_isdf']):
requested_controls = (
int(batch_size),
int(candidate_oversampling),
int(n_topup),
)
cached_controls = (
int(f.attrs.get('isdf_pivot_batch_size', 1)),
int(f.attrs.get('isdf_pivot_candidate_oversampling', 2)),
int(f.attrs.get('isdf_pivot_n_topup', 0)),
)
if cached_controls != requested_controls:
raise ValueError(
"cached ISDF pivot controls %s do not match requested %s"
% (cached_controls, requested_controls)
)
logger.info(f"Loading ISDF decomposition from {save_path}")
pivots = jnp.array(f['pivots'][:])
phi_piv = jnp.array(f['phi_isdf'][:])
grad_phi_piv = jnp.array(f['grad_phi_isdf'][:])
if is_incore:
cpu_device = jax.devices("cpu")[0]
xi_phi = jax.device_put(f['xi_phi'][:], cpu_device)
xi_grad = jax.device_put(f['xi_grad'][:], cpu_device)
else:
xi_phi = None
xi_grad = None
return phi_piv, xi_phi, grad_phi_piv, xi_grad, pivots, save_path
except Exception as e:
logger.warning(f"Failed to load ISDF from {save_path}: {e}. Recomputing...")
n_orb, n_grid = phi.shape
if weights is None:
w_sqrt = jnp.ones(n_grid)
else:
w_sqrt = jnp.sqrt(jnp.abs(weights)) # Use abs to avoid NaN
start_time = time.perf_counter()
logger.info(f"Starting ISDF decomposition with n_orb={n_orb}, n_grid={n_grid}, n_rank_phi={n_rank_phi}, n_rank_grad={n_rank_grad}")
if weights is not None:
logger.info(f" Using integration weights (min={jnp.min(weights):.3e}, max={jnp.max(weights):.3e})")
# --- 1. Phi Decomposition ---
t0 = time.perf_counter()
# Apply weights to orbitals for pivot selection
# The Gram matrix is (phi^T W phi)(phi^T W phi) where W = diag(weights)
# Equivalently: (sqrt(W) phi)^T (sqrt(W) phi) squared
phi_weighted = phi * w_sqrt # (n_orb, n_grid)
orb_sq = jnp.sum(phi_weighted**2, axis=0) # Weighted orbital norms
diag_phi = orb_sq**2
shift_phi = 1e-12 * jnp.max(jnp.abs(diag_phi))
pivots_phi = pivoted_cholesky_streaming(
diag_phi,
JaxPartial(phi_columns, phi_weighted),
shift_phi,
n_rank=n_rank_phi,
batch_size=batch_size,
candidate_oversampling=candidate_oversampling,
n_topup=n_topup,
)
t1 = time.perf_counter()
logger.debug(f"Phi decomposition completed in {t1 - t0:.4f} s")
# --- 2. Gradient Decomposition ---
t0 = time.perf_counter()
grad_phi_weighted = grad_phi * w_sqrt[:, None] # (n_orb, n_grid, 3)
A_diag = jnp.sum(phi_weighted**2, axis=0)
B_diag = jnp.sum(jnp.sum(grad_phi_weighted**2, axis=2), axis=0)
diag_grad = A_diag * B_diag
shift_grad = 1e-12 * jnp.max(jnp.abs(diag_grad))
pivots_grad = pivoted_cholesky_streaming(
diag_grad,
JaxPartial(grad_columns, phi_weighted, grad_phi_weighted),
shift_grad,
n_rank=n_rank_grad,
batch_size=batch_size,
candidate_oversampling=candidate_oversampling,
n_topup=n_topup,
)
t1 = time.perf_counter()
logger.debug(f"Grad decomposition completed in {t1 - t0:.4f} s")
# --- 3. Fuse pivots ---
t0 = time.perf_counter()
# Use numpy for unique to avoid JAX dynamic shape overhead
pivots_all = np.concatenate([np.array(pivots_phi), np.array(pivots_grad)])
pivots = jnp.array(np.unique(pivots_all))
n_fused = pivots.shape[0]
t1 = time.perf_counter()
logger.info(f"Pivots fused: {pivots_phi.shape[0]} + {pivots_grad.shape[0]} -> {n_fused} in {t1 - t0:.4f} s")
# Optionally override the device-selected pivots with an externally-supplied
# fused-pivot set; no effect on the default path (fixed_pivots=None).
if fixed_pivots is not None:
fp = np.asarray(fixed_pivots)
if fp.ndim != 1 or not np.issubdtype(fp.dtype, np.integer):
raise ValueError("fixed_pivots must be a 1-D integer array of grid indices")
if np.unique(fp).shape[0] != fp.shape[0]:
raise ValueError("fixed_pivots must be unique")
if fp.size == 0 or fp.min() < 0 or fp.max() >= n_grid:
raise ValueError(f"fixed_pivots out of range [0, {n_grid})")
pivots = jnp.asarray(fp, dtype=pivots.dtype)
n_fused = int(pivots.shape[0])
logger.info(f"isdf_decompose: overriding with {n_fused} externally-supplied fixed pivots")
# --- 4. Extract pivot values ---
t0 = time.perf_counter()
phi_piv = phi[:, pivots] # (n_orb, n_fused)
grad_phi_piv = grad_phi[:, pivots, :] # (n_orb, n_fused, 3)
t1 = time.perf_counter()
logger.debug(f"Pivot values extracted in {t1 - t0:.4f} s")
# --- 5. Solve for xi_phi and xi_grad using fast normal equations solver ---
t0 = time.perf_counter()
logger.info("Using fast normal equations solver")
cpu_device = jax.devices("cpu")[0]
grid_batch_size = min(grid_batch_size, n_grid)
n_batches = (n_grid + grid_batch_size - 1) // grid_batch_size if grid_batch_size > 0 else 0
# Pre-factor normal-equation matrices once and reuse for all grid batches.
# This avoids rebuilding/re-factorizing ATA in every batch.
phi_chol, phi_lower = prepare_normal_equations_solver(phi_piv, phi_piv, rcond=rcond)
grad_chol = []
grad_lower = []
for c in range(3):
chol_c, lower_c = prepare_normal_equations_solver(grad_phi_piv[:, :, c], phi_piv, rcond=rcond)
grad_chol.append(chol_c)
grad_lower.append(lower_c)
# Multi-device: shard the grid axis of each batch across local devices;
# replicate factors. Use local_devices() so this is safe under multi-process
# JAX (arrays can only be placed on devices visible to this process).
local_devices = jax.local_devices()
n_devices = len(local_devices)
use_sharding = n_devices > 1
if use_sharding:
mesh = Mesh(np.array(local_devices), ('g',))
grid_shard = NamedSharding(mesh, P(None, 'g'))
repl = NamedSharding(mesh, P())
phi_chol = jax.device_put(phi_chol, repl)
phi_piv_d = jax.device_put(phi_piv, repl)
grad_chol = [jax.device_put(c, repl) for c in grad_chol]
grad_phi_piv_d = jax.device_put(grad_phi_piv, repl)
logger.info(f" Multi-device sharding enabled across {n_devices} devices (grid axis)")
else:
phi_piv_d = phi_piv
grad_phi_piv_d = grad_phi_piv
h5_file = None
if is_incore:
logger.info(f" Processing {n_batches} batches of size {grid_batch_size} (In-core)")
xi_phi_storage = np.zeros((n_fused, n_grid), dtype=phi.dtype)
xi_grad_storage = np.zeros((n_fused, n_grid, 3), dtype=phi.dtype)
else:
if save_path is None:
save_path = f"isdf_temp_{uuid.uuid4().hex[:8]}.h5"
logger.info(f" No save_path provided, creating temporary HDF5: {save_path}")
h5_file = h5py.File(save_path, 'a')
logger.info(f" Processing {n_batches} batches of size {grid_batch_size} (HDF5: {save_path})")
h5_file.attrs['isdf_pivot_batch_size'] = int(batch_size)
h5_file.attrs['isdf_pivot_candidate_oversampling'] = int(candidate_oversampling)
h5_file.attrs['isdf_pivot_n_topup'] = int(n_topup)
for name, shape in [('xi_phi', (n_fused, n_grid)), ('xi_grad', (n_fused, n_grid, 3))]:
if name in h5_file: del h5_file[name]
h5_file.create_dataset(name, shape=shape, dtype=phi.dtype)
for name, data in [('pivots', pivots), ('phi_isdf', phi_piv), ('grad_phi_isdf', grad_phi_piv)]:
if name in h5_file: del h5_file[name]
h5_file.create_dataset(name, data=np.array(data))
xi_phi_storage = h5_file['xi_phi']
xi_grad_storage = h5_file['xi_grad']
# Warm up JIT so the sharded-program compile cost doesn't dominate short loops
# (on a 7-batch benzene-5Z run the sharded compile otherwise ate the steady-state
# speedup from parallel GPUs).
if n_batches > 1:
t_warm = time.perf_counter()
warm_batch = jnp.zeros((n_orb, grid_batch_size), dtype=phi.dtype)
if use_sharding:
warm_batch = jax.device_put(warm_batch, grid_shard)
warm = solve_normal_equations_batch_prepared(
phi_chol, phi_lower, phi_piv_d, phi_piv_d, warm_batch, warm_batch
)
jax.block_until_ready(warm)
for c in range(3):
warm = solve_normal_equations_batch_prepared(
grad_chol[c], grad_lower[c], grad_phi_piv_d[:, :, c], phi_piv_d,
warm_batch, warm_batch
)
jax.block_until_ready(warm)
del warm, warm_batch
logger.debug(f" Solve warmup (JIT compile) took {time.perf_counter() - t_warm:.2f} s")
try:
t_batch_start = time.perf_counter()
for batch_idx in range(n_batches):
g_start = batch_idx * grid_batch_size
g_end = min(g_start + grid_batch_size, n_grid)
bs = g_end - g_start
# Pad batch width to a multiple of n_devices so the grid axis shards evenly.
pad = (-bs) % n_devices if use_sharding else 0
phi_batch = phi[:, g_start:g_end]
if pad:
phi_batch = jnp.pad(phi_batch, ((0, 0), (0, pad)))
if use_sharding:
phi_batch = jax.device_put(phi_batch, grid_shard)
xi_phi_batch = solve_normal_equations_batch_prepared(
phi_chol, phi_lower, phi_piv_d, phi_piv_d, phi_batch, phi_batch
)
if pad:
xi_phi_batch = xi_phi_batch[:, :bs]
xi_phi_storage[:, g_start:g_end] = np.array(xi_phi_batch)
for c in range(3):
grad_phi_batch_c = grad_phi[:, g_start:g_end, c]
if pad:
grad_phi_batch_c = jnp.pad(grad_phi_batch_c, ((0, 0), (0, pad)))
if use_sharding:
grad_phi_batch_c = jax.device_put(grad_phi_batch_c, grid_shard)
xi_grad_batch = solve_normal_equations_batch_prepared(
grad_chol[c], grad_lower[c], grad_phi_piv_d[:, :, c], phi_piv_d,
grad_phi_batch_c, phi_batch
)
if pad:
xi_grad_batch = xi_grad_batch[:, :bs]
xi_grad_storage[:, g_start:g_end, c] = np.array(xi_grad_batch)
if batch_idx % 4 == 0 and batch_idx > 0:
elapsed = time.perf_counter() - t_batch_start
rate = batch_idx / elapsed
eta = (n_batches - batch_idx) / rate if rate > 0 else 0
logger.debug(f"Batch {batch_idx}/{n_batches} ({rate:.1f} batch/s, ETA: {eta:.1f}s)")
if is_incore:
xi_phi = jax.device_put(xi_phi_storage[:], cpu_device)
xi_grad = jax.device_put(xi_grad_storage[:], cpu_device)
else:
xi_phi = None
xi_grad = None
if is_incore:
del xi_phi_storage, xi_grad_storage
gc.collect()
finally:
if h5_file is not None:
h5_file.close()
gc.collect()
total_time = time.perf_counter() - start_time
logger.debug(f"Total fused ranks = {n_fused}")
logger.info(f"ISDF decomposition total time: {total_time:.4f} s")
return phi_piv, xi_phi, grad_phi_piv, xi_grad, pivots, save_path