"""Production DF/THC algebra and X-store layouts.
The JAX section implements the panelled LS-THC algebra used by the solver.
``extract_vv_df_factor`` builds its metric-applied virtual DF input. The
X-store section converts ISDF X factors between the
rank-innermost store layout ``(nmo, nmo, rank)`` and the rank-major layout
``(rank, nmo, nmo)`` (panel-contiguous), in place or into a new store.
"""
from __future__ import annotations
import os
from functools import partial
from types import SimpleNamespace
import h5py
import jax
import jax.numpy as jnp
import numpy as np
from numpy.typing import NDArray
# ---------------- DF-factor extraction ----------------
Float64Array = NDArray[np.float64]
# ---------------- JAX production path ----------------
# JAX is float32 by default. Enabling x64 must happen before any array is created;
# doing it at import time is deliberate, and require_float64() verifies it actually
# took effect rather than assuming the flag was honoured.
jax.config.update("jax_enable_x64", True)
class Float64NotEnabled(RuntimeError):
"""Raised when x64 is off -- a float32 run would look like ~1e-7 parity."""
[docs]
def require_float64():
"""Fail closed if JAX is not in x64 mode.
Without this the computation runs in float32: parity comes out around
1e-7 instead of 1e-12. A wrong dtype is not a small error here; it is a
different calculation.
"""
probe = jnp.zeros(1, dtype=jnp.float64)
if probe.dtype != jnp.float64:
raise Float64NotEnabled(
"jax_enable_x64 is not in effect (probe dtype "
f"{probe.dtype}). The sandwich is an FP64 calculation; running it in "
"float32 yields ~1e-7 agreement that can be mistaken for parity. "
"Set JAX_ENABLE_X64=1 before importing jax, or call "
"jax.config.update('jax_enable_x64', True) earlier.")
return True
[docs]
def _as_fp64_jax(name: str, value: object, ndim: int):
"""Validate rank while landing an operand on device as float64."""
array = jnp.asarray(value, dtype=jnp.float64)
if array.ndim != ndim:
raise ValueError(f"{name} must be a {ndim}D array; got {array.ndim}D "
f"{array.shape}")
return array
[docs]
def _as_fp64_host(name: str, value: object, ndim: int):
"""FP64 as a HOST array -- for operands the panel loops only ever slice.
A wholesale device cast of the full 3-index B block costs its entire
size on device even though every consumer reads it one aux panel at a
time. Kept on the host, each panel slice is uploaded by the jitted
kernel that receives it, so peak device cost is one panel.
"""
array = np.asarray(value, dtype=np.float64)
if array.ndim != ndim:
raise ValueError(f"{name} must be a {ndim}D array; got {array.ndim}D "
f"{array.shape}")
return array
# ---------------------------------------------------------------- kernels ---
# Each jitted body is one panel's arithmetic. The explicit subscripts encode
# the (a,c,b,d) source ordering that the RCCSD pair swaps depend on.
@partial(jax.jit, donate_argnums=(0,))
def _exact_accumulate(out, b_panel, t2):
"""out + one aux panel's exact sandwich, donating the accumulator.
Donating the accumulator avoids keeping three full-size buffers
live per panel (out, panel result, sum).
"""
right = jnp.einsum("ijcd,bdq->ijcbq", t2, b_panel)
return out + jnp.einsum("acq,ijcbq->ijab", b_panel, right)
@partial(jax.jit, donate_argnums=(0,))
def _acc_into(acc, block):
"""acc + block with the accumulator's buffer donated (see above)."""
return acc + block
[docs]
def exact_df_panelled(b, t2, aux_panel: int):
"""Exact current-DF sandwich, panelled over the auxiliary axis.
The ``(nocc^2, nvir, nvir, q)`` intermediate inside
:func:`_exact_accumulate` costs ``nocc^2 * nvir^2 * q * 8`` bytes --
~4.9 GiB per aux column at large systems, so the caller's
``aux_panel=32`` would demand ~157 GiB. The
panel step is therefore clamped so the intermediate stays under
``PYTC_EXACT_PANEL_CAP_GB`` (default 8 GiB, read at call time); at q=1
the GEMM shapes are unchanged (the batch is the occupied-pair axis, not
q), so the clamp costs no efficiency. Small cases are unaffected.
"""
per_q = (t2.shape[0] * t2.shape[1] * t2.shape[2] * t2.shape[3]
* np.dtype(np.float64).itemsize)
cap = int(float(os.environ.get("PYTC_EXACT_PANEL_CAP_GB", "8")) * 1024 ** 3)
q_step = max(1, min(int(aux_panel), cap // max(per_q, 1)))
out = jnp.zeros_like(t2)
for q0 in range(0, b.shape[2], q_step):
q1 = min(q0 + q_step, b.shape[2])
out = _exact_accumulate(out, b[:, :, q0:q1], t2)
return out
@jax.jit
def _endpoint_panel(b_q, y_mq):
# (m, b, d): rank-leading so the cross kernels get batch-leading GEMMs.
return jnp.einsum("mq,bdq->mbd", y_mq, b_q)
@partial(jax.jit, static_argnames=("occupied_pair_batch_size",))
def _t2_left_contract_pairs_jit(t2_pairs, factor_panel, *,
occupied_pair_batch_size):
"""tau_t[m, n, d] = sum_c t2_pairs[n, c, d] F[c, m], pair-blocked.
Whole-t2 middle-axis contractions make XLA materialize a physical
transpose of the full t2 at large shapes. Pair-blocking keeps every
temporary at one pair block, and rank-leading output makes every
downstream GEMM batch-leading. ``t2_pairs`` is pair-flattened,
pair-padded t2.
"""
n_padded_pairs, nvir, _ = t2_pairs.shape
n_rank = factor_panel.shape[1]
n_pair_blocks = n_padded_pairs // occupied_pair_batch_size
def pair_body(pair_block, acc):
pair0 = pair_block * occupied_pair_batch_size
tau = jax.lax.dynamic_slice(
t2_pairs, (pair0, 0, 0),
(occupied_pair_batch_size, nvir, nvir))
block = jnp.einsum("cm,ncd->mnd", factor_panel, tau)
return jax.lax.dynamic_update_slice(acc, block, (0, pair0, 0))
return jax.lax.fori_loop(
0, n_pair_blocks, pair_body,
jnp.zeros((n_rank, n_padded_pairs, nvir), dtype=t2_pairs.dtype))
@partial(jax.jit, static_argnames=("occupied_pair_batch_size",))
def _t2_right_contract_pairs_jit(t2_pairs, factor_panel, *,
occupied_pair_batch_size):
"""tau_t[m, n, c] = sum_d t2_pairs[n, c, d] F[d, m]; see the left twin."""
n_padded_pairs, nvir, _ = t2_pairs.shape
n_rank = factor_panel.shape[1]
n_pair_blocks = n_padded_pairs // occupied_pair_batch_size
def pair_body(pair_block, acc):
pair0 = pair_block * occupied_pair_batch_size
tau = jax.lax.dynamic_slice(
t2_pairs, (pair0, 0, 0),
(occupied_pair_batch_size, nvir, nvir))
block = jnp.einsum("dm,ncd->mnc", factor_panel, tau)
return jax.lax.dynamic_update_slice(acc, block, (0, pair0, 0))
return jax.lax.fori_loop(
0, n_pair_blocks, pair_body,
jnp.zeros((n_rank, n_padded_pairs, nvir), dtype=t2_pairs.dtype))
@jax.jit
def _cross_panel(p_panel, endpoint_t, tau_left_t, tau_right_t):
# All rank axes (m) leading: every dot is batch-leading both sides.
# tau_left_t[m,i j,d] / endpoint_t[m,b,d] -> right_t[m,ij,b]
right_t = jnp.einsum("mnd,mbd->mnb", tau_left_t, endpoint_t)
fit_left = jnp.einsum("am,mnb->nab", p_panel, right_t)
# endpoint_t[m,a,c] / tau_right_t[m,ij,c] -> left_t[m,ij,a]
left_t = jnp.einsum("mac,mnc->mna", endpoint_t, tau_right_t)
df_left = jnp.einsum("bm,mna->nab", p_panel, left_t)
return fit_left, df_left
[docs]
def partial_thc_crosses_panelled(b, p_virtual, y, t2, rank_panel: int,
aux_panel: int):
"""Both partial-THC crosses with a precontracted DF endpoint.
The endpoint D[m,b,d] is accumulated over aux panels BEFORE the
occupied-pair contractions, so no (ij,m,b,Qp) intermediate is ever
retained -- that property is the point of the panelled formulation and
is preserved here. t2 enters only through pair-blocked slices (see
:func:`_t2_left_contract_pairs_jit` for why).
"""
nocc_i, nocc_j, nvir, _ = t2.shape
n_pairs = nocc_i * nocc_j
occupied_pair_batch_size = 8
n_pair_blocks = (n_pairs + occupied_pair_batch_size - 1) // occupied_pair_batch_size
padded_pairs = n_pair_blocks * occupied_pair_batch_size
t2_pairs = jnp.pad(
jnp.asarray(t2).reshape(n_pairs, nvir, nvir),
((0, padded_pairs - n_pairs), (0, 0), (0, 0)))
fit_left_df_right = jnp.zeros((padded_pairs, nvir, nvir), dtype=jnp.float64)
df_left_fit_right = jnp.zeros((padded_pairs, nvir, nvir), dtype=jnp.float64)
for m0 in range(0, p_virtual.shape[1], rank_panel):
m1 = min(m0 + rank_panel, p_virtual.shape[1])
p_panel = p_virtual[:, m0:m1]
endpoint_t = jnp.zeros((m1 - m0, b.shape[0], b.shape[1]),
dtype=jnp.float64)
for q0 in range(0, b.shape[2], aux_panel):
q1 = min(q0 + aux_panel, b.shape[2])
endpoint_t = endpoint_t + _endpoint_panel(
b[:, :, q0:q1], y[m0:m1, q0:q1])
tau_left_t = _t2_left_contract_pairs_jit(
t2_pairs, p_panel,
occupied_pair_batch_size=occupied_pair_batch_size)
tau_right_t = _t2_right_contract_pairs_jit(
t2_pairs, p_panel,
occupied_pair_batch_size=occupied_pair_batch_size)
fit_left, df_left = _cross_panel(
p_panel, endpoint_t, tau_left_t, tau_right_t)
fit_left_df_right = _acc_into(fit_left_df_right, fit_left)
df_left_fit_right = _acc_into(df_left_fit_right, df_left)
return (fit_left_df_right[:n_pairs].reshape(t2.shape),
df_left_fit_right[:n_pairs].reshape(t2.shape))
@jax.jit
def _full_thc_block(p_m, p_n, tau_m_t, rank_metric):
# tau_m_t[m, ij, d]; rank axis leading throughout (see the crosses).
tau_mn = jnp.einsum("mjd,dn->mjn", tau_m_t, p_n)
weighted = tau_mn * rank_metric[:, None, :]
right_t = jnp.einsum("mjn,bn->mjb", weighted, p_n)
return jnp.einsum("am,mjb->jab", p_m, right_t)
[docs]
def full_thc_panelled(p_virtual, y, t2, rank_panel: int, aux_panel: int):
"""B_tilde[a,c,Q] B_tilde[b,d,Q] t2[ij,c,d] without ever forming B_tilde."""
nocc_i, nocc_j, nvir, _ = t2.shape
n_pairs = nocc_i * nocc_j
occupied_pair_batch_size = 8
n_pair_blocks = (n_pairs + occupied_pair_batch_size - 1) // occupied_pair_batch_size
padded_pairs = n_pair_blocks * occupied_pair_batch_size
t2_pairs = jnp.pad(
jnp.asarray(t2).reshape(n_pairs, nvir, nvir),
((0, padded_pairs - n_pairs), (0, 0), (0, 0)))
out = jnp.zeros((padded_pairs, nvir, nvir), dtype=jnp.float64)
n_rank = p_virtual.shape[1]
for m0 in range(0, n_rank, rank_panel):
m1 = min(m0 + rank_panel, n_rank)
p_m = p_virtual[:, m0:m1]
tau_m_t = _t2_left_contract_pairs_jit(
t2_pairs, p_m,
occupied_pair_batch_size=occupied_pair_batch_size)
for n0 in range(0, n_rank, rank_panel):
n1 = min(n0 + rank_panel, n_rank)
p_n = p_virtual[:, n0:n1]
rank_metric = jnp.zeros((m1 - m0, n1 - n0), dtype=jnp.float64)
for q0 in range(0, y.shape[1], aux_panel):
q1 = min(q0 + aux_panel, y.shape[1])
rank_metric = rank_metric + y[m0:m1, q0:q1] @ y[n0:n1, q0:q1].T
out = _acc_into(out, _full_thc_block(p_m, p_n, tau_m_t, rank_metric))
return out[:n_pairs].reshape(t2.shape)
# ------------------------------------------------------------ entry point ---
[docs]
class ScalableDirectSandwichesJax:
"""Exact, cross, full-THC, and robust production contractions."""
__slots__ = ("exact", "fit_left_df_right", "df_left_fit_right", "full_thc",
"robust")
def __init__(self, exact, fit_left_df_right, df_left_fit_right, full_thc):
self.exact = exact
self.fit_left_df_right = fit_left_df_right
self.df_left_fit_right = df_left_fit_right
self.full_thc = full_thc
# Explicitly fit-left/DF-right + DF-left/fit-right - full-THC, retaining
# the two crosses separately so signs and RCCSD pair swaps stay auditable.
self.robust = fit_left_df_right + df_left_fit_right - full_thc
[docs]
def df_sandwiches_jax(b, fit, t2, *, rank_panel: int,
aux_panel: int):
"""Evaluate the panelled production DF/THC contractions.
``fit`` exposes the fitted ``p_virtual`` and ``y`` factors.
`b` is deliberately kept HOST-resident (see :func:`_as_fp64_host`): the
panel loops below only read aux-axis slices, so uploading the whole
block would cost its full size in device memory for no benefit.
"""
require_float64()
b_h = _as_fp64_host("b", b, 3)
t2_j = _as_fp64_jax("t2", t2, 4)
p_j = _as_fp64_jax("p_virtual", fit.p_virtual, 2)
y_j = _as_fp64_jax("y", fit.y, 2)
nvir = t2_j.shape[2]
if t2_j.shape[2:] != (nvir, nvir):
raise ValueError(f"t2 virtual axes must be square; got {t2_j.shape}")
if b_h.shape[:2] != (nvir, nvir):
raise ValueError(f"B virtual dimensions do not match t2: {b_h.shape} vs "
f"{t2_j.shape}")
if p_j.shape[0] != nvir or y_j.shape != (p_j.shape[1], b_h.shape[2]):
raise ValueError("implicit fit dimensions do not match B and t2")
# Panel sizes are explicit bounded production inputs.
for name, size, upper in (("rank_panel", rank_panel, p_j.shape[1]),
("aux_panel", aux_panel, b_h.shape[2])):
if not 1 <= int(size) <= upper:
raise ValueError(f"{name} must be in [1, {upper}]; got {size}")
exact = exact_df_panelled(b_h, t2_j, int(aux_panel))
fit_left, df_left = partial_thc_crosses_panelled(
b_h, p_j, y_j, t2_j, int(rank_panel), int(aux_panel))
full = full_thc_panelled(p_j, y_j, t2_j, int(rank_panel), int(aux_panel))
return ScalableDirectSandwichesJax(exact, fit_left, df_left, full)
[docs]
def fit_lsthc_jax(p_virtual, b, *, rcond: float, virtual_panel: int):
"""JAX FP64 normal-equation LS-THC fit with host-resident DF factors."""
require_float64()
if not np.isfinite(rcond) or not 0.0 < float(rcond) <= 1.0:
raise ValueError(f"rcond must be in (0, 1]; got {rcond!r}")
p = _as_fp64_jax("p_virtual", p_virtual, 2)
b_h = _as_fp64_host("b", b, 3)
if b_h.shape[:2] != (p.shape[0], p.shape[0]):
raise ValueError("B and P virtual dimensions differ")
if not 1 <= int(virtual_panel) <= p.shape[0]:
raise ValueError("virtual_panel is out of bounds")
overlap = p.T @ p
gram = overlap * overlap
cross = jnp.zeros((p.shape[1], b_h.shape[2]), dtype=jnp.float64)
for a0 in range(0, p.shape[0], int(virtual_panel)):
a1 = min(a0 + int(virtual_panel), p.shape[0])
b_panel = _as_fp64_jax("b_panel", b_h[a0:a1], 3)
cross = cross + jnp.einsum(
"am,cm,acq->mq", p[a0:a1], p, b_panel)
eigenvalues, eigenvectors = jnp.linalg.eigh(0.5 * (gram + gram.T))
order = jnp.argsort(eigenvalues)[::-1]
eigenvalues, eigenvectors = jnp.maximum(eigenvalues[order], 0.0), eigenvectors[:, order]
singular_values = jnp.sqrt(eigenvalues)
# This is a requested dense-C threshold, not a dimension-dependent policy
# decision for this dispatch.
floor = float(np.sqrt(np.finfo(np.float64).eps))
threshold = max(float(rcond), floor) * singular_values[0]
keep = singular_values > threshold
inv = jnp.where(keep, 1.0 / jnp.where(eigenvalues > 0, eigenvalues, 1.0), 0.0)
y = eigenvectors @ (inv[:, None] * (eigenvectors.T @ cross))
return SimpleNamespace(p_virtual=p, y=y, gram=gram, cross=cross,
rcond=float(rcond), resolved_rcond=max(float(rcond), floor))
# ---------------- X-store layouts ----------------
[docs]
def convert_x_to_rank_major(src, dst, *, dataset="X", row_block=8):
"""Convert one X dataset to rank-major layout, chunked over an nmo axis.
``src``/``dst`` are paths or open ``h5py.File`` objects. The loop
materializes one ``(nmo, row_block, rank)`` slab at a time, so peak
host memory is ``nmo * row_block * rank * 8`` bytes (~1.6 GiB at the
default block). Returns the destination path or
file unchanged. The destination dataset is written contiguous (no
HDF5 chunking) so panel reads stay single sequential extents.
"""
close_src = not isinstance(src, h5py.File)
close_dst = not isinstance(dst, h5py.File)
src_f = h5py.File(src, "r") if close_src else src
dst_f = h5py.File(dst, "w") if close_dst else dst
try:
_convert_x_dataset(src_f, dst_f, dataset, row_block)
return dst
finally:
if close_src:
src_f.close()
if close_dst:
dst_f.close()
[docs]
def _convert_x_dataset(src_f, dst_f, dataset, row_block, out_name=None):
if int(row_block) < 1:
raise ValueError(f"row_block must be >= 1; got {row_block}")
x_in = src_f[dataset]
if x_in.ndim != 3 or x_in.shape[0] != x_in.shape[1]:
raise ValueError(
f"source dataset {dataset!r} must be (nmo, nmo, rank); "
f"got {x_in.shape}")
nmo, _, rank = x_in.shape
x_out = dst_f.create_dataset(out_name or dataset,
shape=(rank, nmo, nmo), dtype=np.float64)
for key, val in x_in.attrs.items():
x_out.attrs[key] = val
x_out.attrs["x_layout"] = "rank_major"
for j0 in range(0, nmo, int(row_block)):
j1 = min(j0 + int(row_block), nmo)
slab = np.asarray(x_in[j0:j1, :, :], dtype=np.float64)
x_out[:, j0:j1, :] = np.ascontiguousarray(
slab.transpose(2, 0, 1))
[docs]
def convert_store_to_rank_major(src, dst, *, x_dataset="X", x_rm_dataset="X_rm",
row_block=8):
"""Whole-store conversion: every dataset copied; X kept AND a rank-major
twin appended.
The ISDF store carries more than X (K1/K3/D kernels, metadata); a
store the driver can actually load needs all of it. Top-level
datasets and file attributes are copied verbatim -- including
``x_dataset`` itself, since legacy consumers read it as
``(nmo, nmo, rank)`` -- and a rank-major twin ``x_rm_dataset`` is
added alongside (with the ``x_layout`` attribute and a provenance
stamp of the source X).
Writes go to ``<dst>.tmp`` and are atomically renamed into place, so
an interrupted conversion never leaves a partial file at the
destination.
"""
tmp = f"{dst}.tmp"
with h5py.File(src, "r") as src_f, h5py.File(tmp, "w") as dst_f:
for key, val in src_f.attrs.items():
dst_f.attrs[key] = val
for key in src_f:
src_f.copy(key, dst_f)
_convert_x_dataset(src_f, dst_f, x_dataset, row_block,
out_name=x_rm_dataset)
dst_f[x_rm_dataset].attrs["x_source_stamp"] = (
_x_source_stamp(src_f[x_dataset]))
os.replace(tmp, dst)
return dst
[docs]
def _x_source_stamp(x):
"""Provenance stamp for an X dataset: sha256 over shape plus EVERY row
slab, so any changed entry changes the stamp (no sampling shortcuts)."""
import hashlib
h = hashlib.sha256(str(tuple(x.shape)).encode())
nmo = x.shape[0]
step = max(1, nmo // 64)
for i0 in range(0, nmo, step):
h.update(np.asarray(x[i0:i0 + step], dtype=np.float64).tobytes())
return h.hexdigest()
[docs]
def add_rank_major(store, *, x_dataset="X", out_dataset="X_rm", row_block=8):
"""Append a rank-major copy of X to an existing store, in place.
The store keeps ``x_dataset`` (innermost) untouched for legacy
consumers (eris build, fingerprint manifest); the factorized
contraction reads ``out_dataset`` when present. An existing
``out_dataset`` is kept only when it is shape/attr-valid AND its
provenance stamp matches the current ``x_dataset``; a stale twin
(X rewritten after the twin was built) is rebuilt.
"""
with h5py.File(store, "r+") as fh:
x_in = fh[x_dataset]
if x_in.ndim != 3 or x_in.shape[0] != x_in.shape[1]:
raise ValueError(
f"{x_dataset!r} must be (nmo, nmo, rank); got {x_in.shape}")
nmo, _, rank = x_in.shape
if out_dataset in fh:
x_rm = fh[out_dataset]
valid = (tuple(x_rm.shape) == (rank, nmo, nmo)
and x_rm.attrs.get("x_layout") == "rank_major")
if not valid:
raise ValueError(
f"{out_dataset!r} exists but is not a valid rank-major X: "
f"shape {x_rm.shape}, attrs {dict(x_rm.attrs)}")
if x_rm.attrs.get("x_source_stamp") == _x_source_stamp(x_in):
return store
del fh[out_dataset]
# Write to a temp dataset and rename into place: an interrupted
# conversion must never leave a correctly-shaped but partially
# written X_rm behind.
tmp_name = out_dataset + ".tmp"
if tmp_name in fh:
del fh[tmp_name]
_convert_x_dataset(fh, fh, x_dataset, row_block, out_name=tmp_name)
fh[tmp_name].attrs["x_source_stamp"] = _x_source_stamp(x_in)
fh.move(tmp_name, out_dataset)
return store