Source code for pytc.test.test_kernel_process_ovvv_block

"""Numerical regression test for ``kernel_process_ovvv_block``.

The production kernel has been rewritten three times in response to
XLA failures on pCV5Z-class (≈30 GB f64 ``ovvv_blk``) inputs:

* v1 (textbook) — paired ``'kdac'``/``'kcad'`` einsums.  XLA fused the
  axis-1↔3 transpose into downstream HLOs and the autotuner aborted
  with ``NOT_FOUND: No valid config found!`` on
  ``f64[1179, 3070116]``.

* v2 (commit ``eeb39b2``) — precompute ``ovvv_swap`` and
  ``ovvv_t1sym = 2*ovvv_blk - ovvv_swap`` inside ``@jax.jit``.  XLA
  re-fused the transpose into a different downstream GEMM
  (``tmp_a``/``tmp_b`` against ``tau``) and blew the BFC allocator on a
  27 GB scratch tensor (``f64[21, 146196, 1179]``).

* v3 (current) — barrier-wrap ``ovvv_swap`` to prevent re-fusion, drop
  the ``ovvv_t1sym`` intermediate (save 30 GB), and pre-transpose
  ``tau_swap`` so the ``tmp_a``/``tmp_b`` contractions don't require a
  reshape-after-transpose of ``ovvv_blk``.

This test pins numerical equivalence against the textbook form so
every future rewrite (including further XLA-workaround tweaks) cannot
silently drift from the CCSD amplitude update's defined output.
"""

import unittest

import numpy as np
import jax
import jax.numpy as jnp

jax.config.update("jax_enable_x64", True)

from pytc.solver.jax_xtc_ccsd import kernel_process_ovvv_block


[docs] def _reference_kernel_process_ovvv_block(ovvv_blk, t1, t2, tau): """Textbook implementation with paired 'kdac'/'kcad' einsums. This is the implementation the production kernel used before the Fix-B transpose refactor. Kept here verbatim as the ground truth for the regression test. """ t1_upd = 2 * jnp.einsum('kdac,ikcd->ia', ovvv_blk, t2) t1_upd -= jnp.einsum('kcad,ikcd->ia', ovvv_blk, t2) t1_upd += 2 * jnp.einsum('kdac,kd,ic->ia', ovvv_blk, t1, t1) t1_upd -= jnp.einsum('kcad,kd,ic->ia', ovvv_blk, t1, t1) Lvv_blk = 2 * jnp.einsum('kdac,kd->ac', ovvv_blk, t1) Lvv_blk -= jnp.einsum('kcad,kd->ac', ovvv_blk, t1) Wvoov_blk = jnp.einsum('kcad,id->akic', ovvv_blk, t1) Wvovo_blk = jnp.einsum('kdac,id->akci', ovvv_blk, t1) tmp_a_blk = jnp.einsum('kdac,ijcd->kaij', ovvv_blk, tau) tmp_b_blk = jnp.einsum('kcbd,ijcd->kbij', ovvv_blk, tau) return t1_upd, Lvv_blk, Wvoov_blk, Wvovo_blk, tmp_a_blk, tmp_b_blk
_OUTPUT_NAMES = ("t1_upd", "Lvv_blk", "Wvoov_blk", "Wvovo_blk", "tmp_a_blk", "tmp_b_blk")
[docs] class TestKernelProcessOvvvBlock(unittest.TestCase): def _run_both(self, nocc, nvir, blk, seed=0): rng = np.random.default_rng(seed) # ovvv_blk: (nocc, nvir, blk, nvir) — sliced along axis 2 (`a`). ovvv_blk = rng.standard_normal((nocc, nvir, blk, nvir)) t1 = rng.standard_normal((nocc, nvir)) t2 = rng.standard_normal((nocc, nocc, nvir, nvir)) tau = rng.standard_normal((nocc, nocc, nvir, nvir)) ref = _reference_kernel_process_ovvv_block( jnp.asarray(ovvv_blk), jnp.asarray(t1), jnp.asarray(t2), jnp.asarray(tau), ) new = kernel_process_ovvv_block( jnp.asarray(ovvv_blk), jnp.asarray(t1), jnp.asarray(t2), jnp.asarray(tau), ) return ref, new def _assert_equal(self, ref, new, rtol=1e-12, atol=1e-12): for name, r, n in zip(_OUTPUT_NAMES, ref, new): r_np = np.asarray(r) n_np = np.asarray(n) self.assertEqual( r_np.shape, n_np.shape, f"{name}: shape mismatch ref={r_np.shape} new={n_np.shape}", ) max_abs = float(np.max(np.abs(r_np - n_np))) denom = max(float(np.max(np.abs(r_np))), 1.0) max_rel = max_abs / denom self.assertTrue( np.allclose(r_np, n_np, rtol=rtol, atol=atol), f"{name}: max|Δ|={max_abs:.3e} max rel={max_rel:.3e}", )
[docs] def test_square_block(self): """blk == nvir exercises the perfectly-square case.""" ref, new = self._run_both(nocc=4, nvir=6, blk=6, seed=1) self._assert_equal(ref, new)
[docs] def test_partial_block(self): """blk < nvir is the common case mid-pipeline.""" ref, new = self._run_both(nocc=3, nvir=7, blk=4, seed=2) self._assert_equal(ref, new)
[docs] def test_block_of_one(self): """blk == 1 is the tail tile at the end of a slab.""" ref, new = self._run_both(nocc=3, nvir=5, blk=1, seed=3) self._assert_equal(ref, new)
[docs] def test_nocc_equals_nvir(self): """A degenerate shape where nocc == nvir catches label-vs-axis bugs that go undetected when every dimension has a unique size.""" ref, new = self._run_both(nocc=4, nvir=4, blk=4, seed=4) self._assert_equal(ref, new)
[docs] def test_large_randomized(self): """Slightly bigger shape to make sure the refactor still agrees when the contractions produce non-trivial amounts of cancellation.""" ref, new = self._run_both(nocc=5, nvir=9, blk=5, seed=5) # Looser tolerance here because the reference kernel evaluates # 2*A - B as two separate einsums while the new kernel folds # the linear combination into a single contraction, so # floating-point reassociation differs at the last bit or two. self._assert_equal(ref, new, rtol=1e-11, atol=1e-11)
if __name__ == "__main__": unittest.main()