Skip to content

Add single-layer UCJ energy - #684

Open
hkbelagali wants to merge 52 commits into
qiskit-community:mainfrom
hkbelagali:main
Open

hkbelagali wants to merge 52 commits into
qiskit-community:mainfrom
hkbelagali:main

Conversation

@hkbelagali

Copy link
Copy Markdown

No description provided.

@CLAassistant

CLAassistant commented Aug 3, 2026 •

Copy link
Copy Markdown

CLA assistant check
All committers have signed the CLA.

@hkbelagali
hkbelagali marked this pull request as ready for review August 8, 2026 02:38
@hkbelagali

Copy link
Copy Markdown
Author

@kevinsung I was looking into why the test cases are getting stuck in CI, I believe it can be narrowed down to a deadlock in jax 0.10.2. This error does not occur on my laptop when I run the test cases, but if I restrict to 4 cores like the GitHub Actions runners, then I am able to reproduce the deadlock in the UCJ algorithm implementation. This code also produces the same deadlock on jax 0.10.2, but works fine with 0.9.2.

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

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

n_calls = 16
batch = 1024
n  = 5 

rng = np.random.default_rng(0)
A = jnp.asarray(rng.normal(size=(batch, n, n)) + 1j * rng.normal(size=(batch, n, n)))

def f(t):
    return sum(jnp.real(jnp.sum(jnp.linalg.det(A * jnp.exp(1j * (t + k)))))
               for k in range(n_calls))
c = jax.jit(f).lower(0.3).compile()
t = time.time(); o = jax.block_until_ready(c(0.3))

print(f"OK exec {time.time()-t:.2f}s val={float(o):.3f}")
JAX_PLATFORMS=cpu taskset -c 0-3 python test.py

on jax==0.9.2, this prints OK exec 0.01s val=-3054.291, but it never finishes running on jax==0.10.2. I think this is because jax parallelized LAPACK operations in 0.10.0, and XLA's source code here has a comment on the safety of this. I believe the fix for this is also live on the jax main branch right now according to this PR. The deadlock disappeared when I used a nightly build of jax. Would it be possible to temporarily pin jax<0.10 until the next release comes out?

@kevinsung

Copy link
Copy Markdown
Collaborator

Sure, you can go ahead and edit pyproject.toml to restrict to working JAX versions.

@hkbelagali

Copy link
Copy Markdown
Author

Sounds good, thanks!

@kevinsung kevinsung left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@hkbelagali Thanks for the contribution! My first request is that you use the newly introduced rotate_one_body_tensor and rotate_two_body_tensor functions from https://github.com/qiskit-community/ffsim/blob/main/python/ffsim/linalg/util.py. I think these can replace the _propagate_through_orbital_rotations and _propagate_spin_sector_tensor functions you introduced here. Note, however, that the orbital rotation convention is transposed from your convention (please check this), so you either need to pass u.T.conj() everywhere, or rework your logic to align with the ffsim convention.

@kevinsung kevinsung left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

To keep things simple for now, let's get rid of the high-level dispatcher functions like ucj_energy, ucj_energy_and_grad, and optimize_ucj_energy and just force the user to use the appropriate function for their operator and Hamiltonian type.

Comment thread python/ffsim/variational/ucj_energy.py Outdated

@kevinsung kevinsung left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I wonder if the variable names with single letters like q, h, and g can be made more descriptive. If you can't think of better names, it's fine though.

Comment thread python/ffsim/variational/ucj_energy.py
Comment thread tests/python/variational/ucj_energy_test.py
Comment thread pyproject.toml
Comment thread tests/python/variational/ucj_energy_test.py Outdated
Comment thread tests/python/variational/ucj_energy_test.py Outdated
Comment thread tests/python/variational/ucj_energy_test.py Outdated
Comment thread python/ffsim/variational/ucj_energy.py Outdated
Comment thread python/ffsim/variational/ucj_energy.py Outdated
Comment thread python/ffsim/variational/ucj_energy.py Outdated
@hkbelagali

Copy link
Copy Markdown
Author

I wonder if the variable names with single letters like q, h, and g can be made more descriptive. If you can't think of better names, it's fine though.

I was trying to match symbols from the paper's equations/lemmas in the code, but I'll change the names to be more descriptive.

@kevinsung

Copy link
Copy Markdown
Collaborator

the lint CI failure should be fixed after merging main

@hkbelagali
hkbelagali requested a review from kevinsung September 5, 2026 15:11
Comment thread tests/python/variational/ucj_energy_test.py
Comment thread docs/how-to-guides/simulate-ucj.ipynb
Comment thread python/ffsim/variational/ucj_energy.py Outdated
Comment thread docs/how-to-guides/simulate-ucj.ipynb Outdated
Comment thread docs/how-to-guides/simulate-ucj.ipynb Outdated
Comment thread python/ffsim/variational/ucj_energy.py
Comment thread python/ffsim/variational/ucj_energy.py
Comment thread python/ffsim/variational/ucj_energy.py Outdated

@kevinsung kevinsung left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for your patience @hkbelagali ! We are getting close. One more small change, as well as a numerical instability found by GPT which I would like you to review.

Comment thread python/ffsim/variational/ucj_energy.py
Comment thread python/ffsim/variational/ucj_energy.py Outdated
Comment on lines +943 to +950
phase_factors = jnp.exp(1j * diagonal_phases)
phased_coeffs = phase_factors[:, :, None] * occ_coeffs[None, :, :]
overlap = jnp.einsum("pi,bpj->bij", occ_coeffs_conj, phased_coeffs)
det = jnp.linalg.det(overlap)
overlap_rhs = jnp.broadcast_to(
occ_coeffs_conj.T, (diagonal_phases.shape[0], n_occ, norb)
)
transition_density = phased_coeffs @ jnp.linalg.solve(overlap, overlap_rhs)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The calculation here is unstable when overlap is close to being singular. I used GPT to find this issue and implement a fix in kevinsung@50dbd7e. The fixed code avoids using solve and seems to be faster. Would you mind reviewing this change to see if it makes sense, and if it is a good change, incorporate it into your PR?

Here is GPT's full writeup, including a repro for the instability causing large energy error:

Problem

For a diagonal phase matrix D and occupied-orbital coefficient matrix Q, the
implementation forms the occupied-space overlap

M = Q† D Q

and then computes

det = jnp.linalg.det(overlap)
transition_density = phased_coeffs @ jnp.linalg.solve(overlap, overlap_rhs)

The energy expressions subsequently multiply the transition density by
det(overlap). Algebraically, these products can remain finite when M is singular:

det(M) M⁻¹ = adj(M).

The implementation nevertheless calculates det(M) and M⁻¹ separately. A
singular M makes solve undefined, while a nearly singular M produces very large
intermediate values that are expected to cancel against a very small determinant.
That cancellation is numerically unstable. The same problem is more pronounced for
the two-body contraction, which contains products of two inverse-derived transition
densities.

D being unitary does not guarantee that Q† D Q is invertible: the phase-rotated
occupied subspace can be partly or completely orthogonal to the original occupied
subspace. These are valid UCJ parameter values, and an optimizer can also pass close
to them.

Observed impact

A spinless example with four orbitals and two electrons reproduces the issue. Use a
normalized 4-by-4 Hadamard matrix for the orbital rotation, and choose a real
symmetric Jastrow matrix such that one two-body transition has phase factors
approximately (1, 1, 1, -1). The corresponding occupied-space overlap has rank one.

The following is a self-contained reproduction against the reviewed, solve-based
implementation:

import numpy as np

import ffsim

norb = 4
nelec = 2

# Q is formed from the first two columns of this Hadamard rotation. For one of
# the two-body transitions below, Q† D Q has rank one.
orbital_rotation = 0.5 * np.array(
    [
        [1, 1, 1, 1],
        [1, -1, 1, -1],
        [1, 1, -1, -1],
        [1, -1, -1, 1],
    ],
    dtype=complex,
)
diag_coulomb_mat = (np.pi / 6) * np.array(
    [
        [0, -1, 1, -2],
        [-1, 0, 1, -2],
        [1, 1, 0, 2],
        [-2, -2, 2, 0],
    ],
    dtype=float,
)
ucj_op = ffsim.UCJOpSpinless(
    diag_coulomb_mats=diag_coulomb_mat[None],
    orbital_rotations=orbital_rotation[None],
)

# Move a tiny distance from the exactly rank-deficient point. The physical
# energy is smooth, but det(M) and solve(M, ...) become badly scaled.
params = ucj_op.to_parameters()
params[0] -= 1e-10
ucj_op = ffsim.UCJOpSpinless.from_parameters(
    params,
    norb=norb,
    n_reps=1,
)

hamiltonian = ffsim.random.random_molecular_hamiltonian_spinless(
    norb,
    seed=0,
)
reference = ffsim.slater_determinant(norb, range(nelec))
state = ffsim.apply_unitary(
    reference,
    ucj_op,
    norb=norb,
    nelec=nelec,
)
linop = ffsim.linear_operator(hamiltonian, norb=norb, nelec=nelec)
statevector_energy = np.vdot(state, linop @ state).real
backpropagated_energy = ffsim.ucj_energy_spinless(
    ucj_op,
    hamiltonian,
    nelec,
)

print(f"state vector:    {statevector_energy:.16f}")
print(f"backpropagation: {backpropagated_energy:.16f}")
print(f"absolute error:  {abs(backpropagated_energy - statevector_energy):.3e}")

Before the bordered-determinant fix, this prints:

state vector:    -1.1225081045966403
backpropagation: -1.1173267624850860
absolute error:  5.181e-03

With the fix described below, the same reproduction gives
-1.1225081045966392 and -1.1225081045966385, respectively, an absolute error
below 1e-15.

At the singular point, the backpropagated energy can happen to agree with the
state-vector result because of floating-point rounding. Tiny valid perturbations show
the instability clearly. In one tested random Hamiltonian:

Parameter perturbation State-vector energy Backpropagated energy Error
-1e-6 -1.1225086523792236 -1.1223683782355460 1.40e-4
-1e-10 -1.1225081045966403 -1.1173267624850860 5.18e-3

The autodifferentiated gradient is also unreliable near this point. For example, one
Jastrow parameter whose state-vector finite-difference derivative is zero returned a
gradient of approximately -0.0703. In exact-singular cases, backend-dependent
NaN or infinite values are also possible.

This affects all three public implementations because spin-balanced,
spin-unbalanced, and spinless evaluation all call _transition_batch.

Recommended fix

Do not represent the determinant-weighted contractions as a determinant multiplied by
an independently computed inverse. Compute the finite polynomial quantities directly.
The downstream code needs three distinct objects:

  1. The overlap det(M).
  2. The determinant-weighted one-body transition density,
    det(M) * rho.
  3. The determinant-weighted antisymmetrized two-body transition density,
    det(M) * (rho[q, p] * rho[s, r] - rho[s, p] * rho[q, r]).

The one-body object can be formed with the adjugate:

det(M) rho = (D Q) adj(M) Q†.

The two-body object can similarly be formed from second-order cofactors. A convenient
way to obtain both quantities without explicitly constructing cofactor tensors is to
use bordered determinants. Define

X = D Q
Y = Q†
M = Y X.

For selected orbital rows a and columns b, the Schur-complement identity gives

-det([[M, Y[:, b]], [X[a, :], 0]])
    = det(M) rho[a, b].

For two selected rows (a, c) and columns (b, d), it gives

det([[M, Y[:, (b, d)]], [X[(a, c), :], 0]])
    = det(M) * (rho[a, b] * rho[c, d] - rho[a, d] * rho[c, b]).

Unlike the Schur-complement derivation, the determinant on the left is still defined
when M is singular. It is exactly the cofactor polynomial needed by the matrix
element. The bordered matrix is only one or two rows and columns larger than M, so
its determinant has the same cubic asymptotic cost as the current determinant and
solve.

Concrete implementation sketch

Replace _transition_batch with helpers that return only quantities which remain
defined at singular overlaps:

def _transition_factors(
    diagonal_phases: jax.Array,
    occ_coeffs: jax.Array,
) -> tuple[jax.Array, jax.Array, jax.Array]:
    """Return M = Q† D Q, X = D Q, and Y = Q† for a phase batch."""
    q_dag = jnp.conj(occ_coeffs).T
    phase_factors = jnp.exp(1j * diagonal_phases)
    phased_coeffs = phase_factors[:, :, None] * occ_coeffs[None, :, :]
    overlap = jnp.einsum("ip,bpj->bij", q_dag, phased_coeffs)
    return overlap, phased_coeffs, q_dag


def _overlap_batch(
    diagonal_phases: jax.Array,
    occ_coeffs: jax.Array,
) -> jax.Array:
    """Return <Q|D|Q> for each diagonal phase in a batch."""
    overlap, _, _ = _transition_factors(diagonal_phases, occ_coeffs)
    return jnp.linalg.det(overlap)


def _det_weighted_transition_batch(
    diagonal_phases: jax.Array,
    occ_coeffs: jax.Array,
    row_orbitals: jax.Array,
    col_orbitals: jax.Array,
) -> jax.Array:
    """Return determinant-weighted one- or two-body transitions.

    ``row_orbitals`` and ``col_orbitals`` both have shape ``(batch, k)``.
    ``k=1`` returns ``det(M) * rho[row, col]``. ``k=2`` returns the
    determinant-weighted antisymmetrized product of the two selected rows and
    columns.
    """
    overlap, phased_coeffs, q_dag = _transition_factors(
        diagonal_phases, occ_coeffs
    )
    batch_size = diagonal_phases.shape[0]
    body_order = row_orbitals.shape[1]

    # X[(a, c), :] for each member of the batch: (batch, k, n_occ).
    selected_rows = phased_coeffs[
        jnp.arange(batch_size)[:, None], row_orbitals
    ]

    # Y[:, (b, d)] for each member of the batch: (batch, n_occ, k).
    # Advanced indexing first produces (n_occ, batch, k).
    selected_cols = jnp.transpose(q_dag[:, col_orbitals], (1, 0, 2))

    zeros = jnp.zeros(
        (batch_size, body_order, body_order), dtype=overlap.dtype
    )
    upper = jnp.concatenate([overlap, selected_cols], axis=2)
    lower = jnp.concatenate([selected_rows, zeros], axis=2)
    bordered = jnp.concatenate([upper, lower], axis=1)

    # det([[M, Y], [X, 0]]) = (-1)^k det(M) det(X M^-1 Y).
    return (-1) ** body_order * jnp.linalg.det(bordered)

The helper above was checked against the existing det * rho and
det * antisymmetrized(rho * rho) expressions on random, well-conditioned complex
inputs. The maximum discrepancy was approximately 1.1e-16, so the index ordering
and the (-1) ** body_order sign agree with the current implementation away from the
singular case.

This implementation also naturally handles zero- and one-electron sectors. For
example, a two-body transition with fewer than two occupied orbitals produces a
rank-deficient bordered determinant and therefore zero, without special-casing an
inverse of a 0-by-0 or 1-by-1 overlap.

The current one-body call site can then change from

det_sector, rho_sector = _transition_batch(phi[:, sector_slice], q_sector)
det_other, _ = _transition_batch(phi[:, other_slice], q_other)

return jnp.sum(
    h_sector[p, q]
    * const
    * det_sector
    * det_other
    * rho_sector[rows, q, p]
)

to

weighted_transition = _det_weighted_transition_batch(
    phi[:, sector_slice],
    q_sector,
    row_orbitals=q[:, None],
    col_orbitals=p[:, None],
)
det_other = _overlap_batch(phi[:, other_slice], q_other)

return jnp.sum(h_sector[p, q] * const * weighted_transition * det_other)

The same-spin two-body call site should request a second-order bordered determinant:

weighted_wick = _det_weighted_transition_batch(
    phi[:, sector_slice],
    q_sector,
    row_orbitals=jnp.stack([q, s], axis=1),
    col_orbitals=jnp.stack([p, r], axis=1),
)
det_other = _overlap_batch(phi[:, other_slice], q_other)

return 0.5 * coeff * const * weighted_wick * det_other

For opposite-spin terms, each spin sector contributes a determinant-weighted
one-body transition, so no separate determinant factors are needed:

weighted_left = _det_weighted_transition_batch(
    phi[:, _spin_slice(left_spin, norb)],
    q_left,
    row_orbitals=q[:, None],
    col_orbitals=p[:, None],
)
weighted_right = _det_weighted_transition_batch(
    phi[:, _spin_slice(right_spin, norb)],
    q_right,
    row_orbitals=s[:, None],
    col_orbitals=r[:, None],
)

return 0.5 * g_flat[indices] * const * weighted_left * weighted_right

The spinless one-body and two-body paths use the same substitutions. For the
spinless two-body path, use rows stack([q, s]) and columns stack([p, r]), matching
the existing Wick contraction.

The bordered-determinant path can be used unconditionally: it removes solve, keeps
the same asymptotic scaling, and avoids a condition-number branch whose threshold
would introduce another numerical decision into the differentiated objective. If
benchmarking shows that a well-conditioned solve path is materially faster, retain
it only as an optimization and select the bordered path for ill-conditioned batches.
A diagonal regularizer or pseudoinverse is not a correct fallback because it changes
the finite cofactor limit at exactly singular overlaps.

The code should verify JAX gradients at exactly rank-deficient bordered matrices. JAX
currently produces cofactor-based derivatives for jnp.linalg.det, but a local
custom-JVP determinant based on minors would be the fallback if a supported backend
returns non-finite determinant derivatives. Returning an unweighted rho should be
avoided because it inherently requires division by a possibly zero overlap.

Tests to add

  • A structured spinless UCJ case whose occupied-space overlap is rank deficient,
    compared against state-vector energy evaluation.
  • Small perturbations on both sides of the singular point, again compared against the
    state-vector result.
  • Gradient comparisons near the singular point, using finite differences of the
    state-vector energy rather than finite differences of the same backpropagation
    implementation.
  • Chunked and unchunked variants, to ensure the fallback is independent of batching.
  • At least one spinful case, since the spinful paths combine overlap factors from two
    sectors differently.

Implemented fix and measured performance

The working tree now implements the bordered-determinant helpers and updates every
spinless and spinful call site described above. It also adds
test_ucj_energy_singular_overlap_spinless, which contains the structured
rank-deficient example, checks chunked and unchunked energies against state-vector
evaluation, and compares five autodifferentiated gradient entries with finite
differences of the state-vector energy.

The regression improves the previously observed 5.18e-3 energy error to less than
1e-15. The five sampled gradient discrepancies are at most approximately 2.4e-9;
the solve-based implementation had errors as large as 7e-2 in the same case.

Performance was measured on the JAX CPU backend with JAX 0.11.1, NumPy 2.5.3, and an
x86-64 Intel Xeon Cascadelake processor. Both implementations were loaded in the same
Python process: the original implementation came from HEAD, while the working-tree
module contained the bordered-determinant implementation. After compiling both, the
benchmark alternated old and new calls and reports median synchronous wall time. The
four-orbital cases used chunk_size=16 and 20 samples per callable.

Four-orbital case Operation Solve-based Bordered determinant Change
Spinless, 2 electrons Energy 2.568 ms 2.404 ms -6.4%
Spinless, 2 electrons Energy + gradient 1.522 ms 1.305 ms -14.3%
Balanced, (2, 2) electrons Energy 9.292 ms 7.408 ms -20.3%
Balanced, (2, 2) electrons Energy + gradient 11.462 ms 6.654 ms -41.9%
Unbalanced, (2, 2) electrons Energy 10.891 ms 9.563 ms -12.2%
Unbalanced, (2, 2) electrons Energy + gradient 9.585 ms 5.947 ms -38.0%

A second warmed benchmark used six orbitals, chunk_size=64, and 10 samples:

Six-orbital case Operation Solve-based Bordered determinant Change
Spinless, 3 electrons Energy 3.532 ms 2.876 ms -18.6%
Spinless, 3 electrons Energy + gradient 4.530 ms 3.017 ms -33.4%
Balanced, (3, 3) electrons Energy 30.633 ms 19.032 ms -37.9%
Balanced, (3, 3) electrons Energy + gradient 67.691 ms 32.020 ms -52.7%

Separate-process first-call measurements, which include JAX tracing and compilation,
varied between a 2% improvement and a 15% regression depending on the ansatz. That
range is noisy and compilation dominates it; it does not indicate a material
steady-state regression. The paired warmed results consistently favor the new path.
This is plausible because each Hamiltonian term now computes only the selected scalar
transition via one bordered determinant rather than solving for the entire norb by
norb transition-density matrix.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks! I think this works well, I benchmarked it on larger systems and it seems to be faster there as well. I've added it to my branch. Should we have test cases for this as well?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants