Source code for pytc.df.isdf

"""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