"""Optimizer implementations for VMC."""
import logging
import jax
import jax.numpy as jnp
import jax.scipy.sparse.linalg as spla
import optax
import folx
from jax.tree_util import tree_map
import jax.flatten_util
logger = logging.getLogger(__name__)
[docs]
class NewtonOptimizer:
"""Newton Optimizer (formerly Matrix-Free Optimizer).
Supports:
- Stochastic Reconfiguration (SR) / Natural Gradient for Energy Minimization
(curvature="fisher")
- Gauss-Newton for Variance Minimization
(curvature="gauss_newton")
Solvers:
- "cg": Conjugate Gradient (iterative, matrix-free)
- "exact" or "cholesky": Exact matrix inversion
"""
def __init__(self, value_and_grad_func, learning_rate, damping=1e-3, maxiter=100, curvature_type="fisher", max_vmap_batch_size=0, solver="exact", solve_kwargs=None, jacobian_sample_size=0, clip_multiplier=5.0, jac_row_clip_multiplier=0.0, max_delta_norm=None, mesh=None):
self.value_and_grad_func = value_and_grad_func
self.learning_rate = learning_rate
self.damping = damping
self.maxiter = maxiter
self.curvature_type = curvature_type
self.max_vmap_batch_size = max_vmap_batch_size
self.solver = solver
self.solve_kwargs = solve_kwargs if solve_kwargs is not None else {}
self.jacobian_sample_size = jacobian_sample_size
self.clip_multiplier = clip_multiplier
# Opt-in hardening against huge-but-finite Jacobian rows from
# near-nodal walkers: jac_row_clip_multiplier rescales outlier rows;
# max_delta_norm is a trust-region radius RELATIVE to
# max(1, ||params||). Both OFF by default (0/None) to preserve the
# numerics of existing calculations; enable explicitly for runs that
# chain many sub-optimizations.
self.jac_row_clip_multiplier = jac_row_clip_multiplier
self.max_delta_norm = max_delta_norm
# Thread one Mesh instance through so _get_vmap uses the same mesh
# as walker init/MCMC instead of re-deriving one when mesh=None.
self.mesh = mesh
def _get_vmap(self):
"""Return the appropriate vmap implementation.
Automatically detects multi-GPU environments via get_vmap_fn.
"""
from .sharding import get_vmap_fn
return get_vmap_fn(max_vmap_batch_size=self.max_vmap_batch_size, mesh=self.mesh)
def _get_unbatched_vmap(self):
"""Return a per-batch vmap without nested folx batching."""
from .sharding import get_vmap_fn
return get_vmap_fn(max_vmap_batch_size=0, mesh=self.mesh)
def _get_effective_batch_size(self, n_walkers: int) -> int:
"""Return a batch size compatible with the current execution mode."""
if self.max_vmap_batch_size <= 0:
return n_walkers
batch_size = min(self.max_vmap_batch_size, n_walkers)
from .sharding import is_multi_gpu, n_devices
if is_multi_gpu():
ndev = n_devices()
if batch_size % ndev != 0:
batch_size = max(ndev, (batch_size // ndev) * ndev)
batch_size = min(batch_size, n_walkers)
return batch_size
@staticmethod
def _slice_walkers(walkers, start: int, stop: int):
"""Slice a walker pytree along the walker dimension."""
return jax.tree_util.tree_map(lambda x: x[start:stop], walkers)
@staticmethod
def _pad_walkers_to_batch_multiple(walkers, batch_size: int):
"""Pad walkers along axis 0 so they reshape cleanly into batches."""
n_walkers = walkers.shape[0]
n_batches = (n_walkers + batch_size - 1) // batch_size
padded_n = n_batches * batch_size
pad_count = padded_n - n_walkers
if pad_count == 0:
mask = jnp.ones((padded_n,), dtype=bool)
return walkers, mask, n_batches
padded_walkers = jax.tree_util.tree_map(
lambda x: jnp.concatenate(
[x, jnp.repeat(x[:1], pad_count, axis=0)],
axis=0,
),
walkers,
)
mask = jnp.concatenate(
[jnp.ones((n_walkers,), dtype=bool), jnp.zeros((pad_count,), dtype=bool)],
axis=0,
)
return padded_walkers, mask, n_batches
@staticmethod
def _flatten_jacobian(jacobian, n_walkers: int):
"""Flatten a per-walker Jacobian pytree to an ``(N, P)`` matrix."""
jac_flat, _ = jax.tree_util.tree_flatten(jacobian)
return jnp.concatenate(
[jnp.reshape(leaf, (n_walkers, -1)) for leaf in jac_flat],
axis=1,
)
@staticmethod
def _clip_jacobian_rows(jac_mat, clip_multiplier, mask=None):
"""Clip each walker's Jacobian row to at most
``clip_multiplier * median(row_norm)`` (over the unmasked rows),
rescaling the whole row down (direction preserved) rather than
clamping individual entries -- same spirit as the existing
energy MAD-clipping, applied to row *scale* so a single
near-nodal walker's huge-but-finite row can't dominate
``sum_jte``/``sum_jtj``. Median (not mean) since it is itself
robust to the exact walkers this is meant to guard against.
Statistics are computed within the same batch/set being clipped
(no extra Jacobian pass) -- cheap, at the cost of being a local
rather than a whole-walker-population estimate.
"""
if clip_multiplier is None or clip_multiplier <= 0:
return jac_mat
# Defense-in-depth: callers are expected to have already excluded
# non-finite rows (finite-masking happens before this is called in
# NewtonOptimizer.step), but guard here too -- an Inf row gives
# row_norm=inf, and threshold/inf=0, so scale*row = 0*inf = NaN
# without this, silently manufacturing a NaN from the clip itself
# rather than the walker that was already broken.
finite_row = jnp.all(jnp.isfinite(jac_mat), axis=1)
jac_mat = jnp.where(finite_row[:, None], jac_mat, 0.0)
row_norm = jnp.linalg.norm(jac_mat, axis=1)
if mask is not None:
mask = mask & finite_row
else:
mask = finite_row
# Masked (padding/non-finite) rows are already zero; excluding
# them from the median keeps the threshold meaningful.
valid_norm = jnp.where(mask, row_norm, jnp.nan)
median_norm = jnp.nanmedian(valid_norm)
threshold = clip_multiplier * median_norm
# If every row is masked/non-finite, nanmedian is NaN and the clip
# would itself manufacture NaNs; fall back to "no clipping" and let
# the finite-masking above (rows already zeroed) carry the batch.
threshold = jnp.where(jnp.isfinite(threshold), threshold, jnp.inf)
scale = jnp.minimum(1.0, threshold / jnp.maximum(row_norm, 1e-300))
return jac_mat * scale[:, None]
[docs]
def init(self, params, rng, batch):
return jnp.array(0, dtype=jnp.int32) # step count
[docs]
def step(self, params, state, rng, batch, global_step_int=None):
walkers, ansatz = batch
if self.solver == "exact" or self.solver == "cholesky":
# Exact inversion: (M + lambda I) delta = -g
if self.curvature_type == "fisher":
# SR: S = Cov(grad_log_psi)
# Need loss + grads from value_and_grad_func, plus Jacobian separately.
(loss, aux_data), grads = self.value_and_grad_func(params, batch)
# J_i = d(log_psi(w_i))/dp
def single_log_psi_grad(w, p):
return jax.grad(lambda pp: ansatz(w, pp)[0][1])(p)
# Compute Jacobian for all walkers: Shape (N, P)
vmap_fn = self._get_vmap()
jac = vmap_fn(single_log_psi_grad, in_axes=(0, None))(walkers, params)
jac_flat, params_treedef = jax.tree_util.tree_flatten(jac)
jac_mat = jnp.concatenate([jnp.reshape(leaf, (walkers.shape[0], -1)) for leaf in jac_flat], axis=1)
jac_centered = jac_mat - jnp.mean(jac_mat, axis=0, keepdims=True)
# S = 1/N * J.T @ J
n_walkers = walkers.shape[0]
curvature_mat = (jac_centered.T @ jac_centered) / n_walkers
# Finite-masking/row-clipping (below) is gauss_newton-specific;
# no dropped-walker tracking on this path.
n_dropped = jnp.array(0, dtype=jnp.int32)
elif self.curvature_type == "gauss_newton":
# GN: G = 2/M * J_centered.T @ J_centered
# J_i = d(E_L(w_i))/dp
#
# Optimization: compute Jacobian and local energies in one pass,
# then derive the variance loss and gradient analytically:
# variance = sum((E - mean(E))^2) / (M - 1)
# grad_variance = 2/(M-1) * J^T @ (E - mean(E))
#
# When jacobian_sample_size > 0, a random subset of walkers is
# used for the Jacobian (gradient + curvature), reducing cost
# from O(N) to O(M) local-energy differentiations per step.
def single_local_energy_and_grad(w, p):
"""Compute both E_L(w) and grad_p E_L(w) in one pass."""
return jax.value_and_grad(lambda pp: ansatz.local_energy(w, pp)[0])(p)
def single_local_energy(w, p):
return ansatz.local_energy(w, p)[0]
n_walkers_total = walkers.shape[0]
if self.jacobian_sample_size > 0 and self.jacobian_sample_size < n_walkers_total:
sample_size = self.jacobian_sample_size
# In multi-device mode, shard_map requires the sharded axis
# length to be divisible by device count.
from .sharding import is_multi_gpu, n_devices
if is_multi_gpu():
ndev = n_devices()
if sample_size % ndev != 0:
sample_size = max(ndev, (sample_size // ndev) * ndev)
sample_size = min(sample_size, n_walkers_total)
indices = jax.random.choice(rng, n_walkers_total,
shape=(sample_size,), replace=False)
indices = jnp.sort(indices).astype(jnp.int32) # sort for deterministic gather
sub_walkers = jax.tree_util.tree_map(lambda x: x[indices], walkers)
else:
sample_size = n_walkers_total
sub_walkers = walkers
n_walkers = sample_size
if self.max_vmap_batch_size > 0:
batch_vmap_fn = self._get_unbatched_vmap()
batch_size = self._get_effective_batch_size(n_walkers)
padded_walkers, padded_mask, n_batches = self._pad_walkers_to_batch_multiple(
sub_walkers,
batch_size,
)
batched_walkers = jax.tree_util.tree_map(
lambda x: x.reshape((n_batches, batch_size) + x.shape[1:]),
padded_walkers,
)
batched_mask = padded_mask.reshape((n_batches, batch_size))
params_vec, _ = jax.flatten_util.ravel_pytree(params)
param_dtype = params_vec.dtype
def masked_energy_batch(batch_walkers, batch_mask):
energies_batch = batch_vmap_fn(
single_local_energy,
in_axes=(0, None),
)(batch_walkers, params)
# A walker can be genuinely non-finite (not just huge) --
# e.g. a near-nodal walker whose updated params push some
# Jastrow term into a mathematically undefined regime.
# Clipping downstream can only rescale huge-but-finite
# values; a NaN/inf must be excluded here instead, same
# treatment as a padding row, or it silently poisons
# every accumulated sum it touches regardless of clipping.
batch_mask = batch_mask & jnp.isfinite(energies_batch)
return jnp.where(batch_mask, energies_batch, 0.0)
def masked_energy_jacobian_batch(batch_walkers, batch_mask, clip_lo, clip_hi):
orig_mask = batch_mask
energies_batch, jac_batch = batch_vmap_fn(
single_local_energy_and_grad,
in_axes=(0, None),
)(batch_walkers, params)
jac_mat_batch = self._flatten_jacobian(jac_batch, batch_size)
# See masked_energy_batch: exclude genuinely non-finite
# walkers (energy OR any jacobian component) before any
# clipping runs, so nanmedian in _clip_jacobian_rows can't
# confuse "real broken walker" with "intentional padding
# sentinel", and NaN can't survive a finite rescale.
batch_mask = (
orig_mask
& jnp.isfinite(energies_batch)
& jnp.all(jnp.isfinite(jac_mat_batch), axis=1)
)
energies_batch = jnp.where(batch_mask, energies_batch, 0.0)
if clip_lo is not None and clip_hi is not None:
clipped = jnp.clip(energies_batch, clip_lo, clip_hi)
energies_batch = jnp.where(batch_mask, clipped, 0.0)
jac_mat_batch = jnp.where(batch_mask[:, None], jac_mat_batch, 0.0)
jac_mat_batch = self._clip_jacobian_rows(
jac_mat_batch, self.jac_row_clip_multiplier, mask=batch_mask
)
# Walkers dropped for non-finiteness specifically, not
# counting padding rows (which orig_mask already excludes).
n_dropped_batch = jnp.sum((~batch_mask) & orig_mask).astype(jnp.int32)
return energies_batch, jac_mat_batch, n_dropped_batch
clip_lo = None
clip_hi = None
if self.clip_multiplier > 0:
def mean_scan_body(carry, xs):
batch_walkers, batch_mask = xs
energies_batch = masked_energy_batch(batch_walkers, batch_mask)
return carry + jnp.sum(energies_batch), None
sum_e_raw, _ = jax.lax.scan(
mean_scan_body,
jnp.array(0.0, dtype=param_dtype),
(batched_walkers, batched_mask),
)
e_mean_raw = sum_e_raw / n_walkers
def mad_scan_body(carry, xs):
batch_walkers, batch_mask = xs
energies_batch = masked_energy_batch(batch_walkers, batch_mask)
mad_terms = jnp.where(
batch_mask,
jnp.abs(energies_batch - e_mean_raw),
0.0,
)
return carry + jnp.sum(mad_terms), None
mad_sum, _ = jax.lax.scan(
mad_scan_body,
jnp.array(0.0, dtype=param_dtype),
(batched_walkers, batched_mask),
)
e_std_raw = mad_sum / n_walkers
clip_lo = e_mean_raw - self.clip_multiplier * e_std_raw
clip_hi = e_mean_raw + self.clip_multiplier * e_std_raw
init_carry = (
jnp.array(0.0, dtype=param_dtype),
jnp.array(0.0, dtype=param_dtype),
jnp.zeros_like(params_vec),
jnp.zeros_like(params_vec),
jnp.zeros((params_vec.shape[0], params_vec.shape[0]), dtype=param_dtype),
jnp.array(0, dtype=jnp.int32),
)
def stats_scan_body(carry, xs):
sum_e, sum_e2, sum_j, sum_jte, sum_jtj, n_dropped = carry
batch_walkers, batch_mask = xs
energies_batch, jac_mat_batch, n_dropped_batch = masked_energy_jacobian_batch(
batch_walkers,
batch_mask,
clip_lo,
clip_hi,
)
sum_e = sum_e + jnp.sum(energies_batch)
sum_e2 = sum_e2 + jnp.sum(energies_batch**2)
sum_j = sum_j + jnp.sum(jac_mat_batch, axis=0)
sum_jte = sum_jte + jac_mat_batch.T @ energies_batch
sum_jtj = sum_jtj + jac_mat_batch.T @ jac_mat_batch
n_dropped = n_dropped + n_dropped_batch
return (sum_e, sum_e2, sum_j, sum_jte, sum_jtj, n_dropped), None
(sum_e, sum_e2, sum_j, sum_jte, sum_jtj, n_dropped), _ = jax.lax.scan(
stats_scan_body,
init_carry,
(batched_walkers, batched_mask),
)
# Dropped (non-finite) walkers contribute zeros to every
# sum above; dividing by the full n_walkers would treat
# them as real zero-energy/zero-Jacobian samples and bias
# the step. Use the valid count (floored to avoid /0).
n_valid = jnp.maximum(n_walkers - n_dropped, 2)
e_mean = sum_e / n_valid
e_std = jnp.sqrt(jnp.maximum(sum_e2 / n_valid - e_mean**2, 0.0))
mean_j = sum_j / n_valid
loss = (sum_e2 - n_valid * e_mean**2) / (n_valid - 1)
aux_data = (e_mean, e_std)
grads_vec = (2.0 / (n_valid - 1)) * (
sum_jte - n_valid * mean_j * e_mean
)
curvature_mat = (2.0 / n_valid) * (
sum_jtj - n_valid * jnp.outer(mean_j, mean_j)
)
else:
vmap_fn = self._get_vmap()
energies, jac = vmap_fn(
single_local_energy_and_grad,
in_axes=(0, None)
)(sub_walkers, params)
jac_mat = self._flatten_jacobian(jac, n_walkers)
# Exclude genuinely non-finite walkers (energy OR any
# jacobian component) BEFORE computing any clip statistics --
# a NaN/inf energy would otherwise poison e_mean_raw/
# e_std_raw themselves, and clipping can only rescale
# huge-but-finite values, not NaN/inf. Same treatment as a
# padding row in the batched path (zero + excluded from
# stats, n_walkers denominator unchanged).
finite_mask = jnp.isfinite(energies) & jnp.all(jnp.isfinite(jac_mat), axis=1)
n_dropped = jnp.sum(~finite_mask)
if n_dropped.dtype != jnp.int32:
n_dropped = n_dropped.astype(jnp.int32)
energies = jnp.where(finite_mask, energies, 0.0)
jac_mat = jnp.where(finite_mask[:, None], jac_mat, 0.0)
# Clip energies to suppress outliers. The gradient is
# 2/(M-1) * J^T @ (E - mean(E)), so clipping energies
# naturally limits the influence of extreme walkers.
if self.clip_multiplier > 0:
e_mean_raw = jnp.sum(energies) / jnp.maximum(jnp.sum(finite_mask), 1)
e_std_raw = jnp.sum(jnp.where(finite_mask, jnp.abs(energies - e_mean_raw), 0.0)) / jnp.maximum(jnp.sum(finite_mask), 1)
clipped = jnp.clip(
energies,
e_mean_raw - self.clip_multiplier * e_std_raw,
e_mean_raw + self.clip_multiplier * e_std_raw,
)
energies = jnp.where(finite_mask, clipped, 0.0)
jac_mat = self._clip_jacobian_rows(jac_mat, self.jac_row_clip_multiplier, mask=finite_mask)
# Compute variance loss and auxiliary data analytically
# from energies. Dropped walkers were zeroed above, so
# denominators use the valid count and centered
# quantities are re-masked to keep dropped rows at zero.
n_valid = jnp.maximum(jnp.sum(finite_mask), 2)
e_mean = jnp.sum(energies) / n_valid
energy_diff = jnp.where(finite_mask, energies - e_mean, 0.0)
e_std = jnp.sqrt(jnp.sum(energy_diff**2) / n_valid)
loss = jnp.sum(energy_diff**2) / (n_valid - 1)
aux_data = (e_mean, e_std)
# Compute variance gradient analytically:
# grad_variance = 2/(M-1) * J^T @ (E - mean(E))
grads_vec = (2.0 / (n_valid - 1)) * (jac_mat.T @ energy_diff)
mean_jac = jnp.sum(jac_mat, axis=0, keepdims=True) / n_valid
jac_centered = jnp.where(finite_mask[:, None], jac_mat - mean_jac, 0.0)
curvature_mat = (2.0 / n_valid) * (jac_centered.T @ jac_centered)
else:
raise ValueError(f"Unknown curvature type: {self.curvature_type}")
curvature_mat = curvature_mat + self.damping * jnp.eye(curvature_mat.shape[0])
# Flatten gradients to match matrix (for fisher, grads are pytree; for gauss_newton, already flat)
if self.curvature_type == "fisher":
grads_vec, unravel_fn = jax.flatten_util.ravel_pytree(grads)
else:
# gauss_newton: grads_vec is already flat, need unravel_fn
_, unravel_fn = jax.flatten_util.ravel_pytree(params)
# (M + lambda I) delta = -g
solve_kwargs = self.solve_kwargs.copy()
if "assume_a" not in solve_kwargs:
solve_kwargs["assume_a"] = "pos"
delta_vec = jax.scipy.linalg.solve(curvature_mat, -grads_vec, **solve_kwargs)
# Trust-region cap (opt-in): rescale the whole step (direction
# preserved) when a near-flat curvature direction yields a huge
# but finite step; the bound is relative to max(1, ||params||)
# so it tracks the parameter scale.
if self.max_delta_norm is not None and self.max_delta_norm > 0:
params_vec_for_scale, _ = jax.flatten_util.ravel_pytree(params)
trust_radius = self.max_delta_norm * jnp.maximum(1.0, jnp.linalg.norm(params_vec_for_scale))
delta_norm = jnp.linalg.norm(delta_vec)
delta_scale = jnp.minimum(1.0, trust_radius / jnp.maximum(delta_norm, 1e-300))
delta_vec = delta_vec * delta_scale
delta = unravel_fn(delta_vec)
lr = self.learning_rate(state) if callable(self.learning_rate) else self.learning_rate
new_params = jax.tree_util.tree_map(lambda p, d: p + lr * d, params, delta)
return new_params, state + 1, {"loss": loss, "aux": aux_data, "lr": lr, "n_dropped_walkers": n_dropped}
(loss, aux_data), grads = self.value_and_grad_func(params, batch)
if self.curvature_type == "fisher":
# SR: S = Cov(grad_log_psi)
def single_log_psi(w, p):
return ansatz(w, p)[0][1]
def mvp(v):
# Forward: w = J v
# Compute JVP per walker: d(log_psi)/dp * v
def compute_jvp(w):
_, tangent = jax.jvp(lambda p: single_log_psi(w, p), (params,), (v,))
return tangent
vmap_fn = self._get_vmap()
w = vmap_fn(compute_jvp)(walkers)
w_centered = w - jnp.mean(w)
# Backward: J.T w_centered
# Compute VJP per walker: (d(log_psi)/dp)^T * w_i
def compute_vjp(w_el, w_val):
_, vjp_fun = jax.vjp(lambda p: single_log_psi(w_el, p), params)
return vjp_fun(w_val)[0]
per_walker_grads = vmap_fn(compute_vjp)(walkers, w_centered)
u = jax.tree_util.tree_map(lambda x: jnp.sum(x, axis=0), per_walker_grads)
# S = 1/N * J.T @ (J @ v centered)
n_walkers = walkers.shape[0]
return jax.tree_util.tree_map(lambda x: x / n_walkers, u)
elif self.curvature_type == "gauss_newton":
# GN: G = 2/N * J.T @ J
def single_local_energy(w, p):
return ansatz.local_energy(w, p)[0]
def mvp(v):
# Forward: w = J v
def compute_jvp(w):
_, tangent = jax.jvp(lambda p: single_local_energy(w, p), (params,), (v,))
return tangent
vmap_fn = self._get_vmap()
w = vmap_fn(compute_jvp)(walkers)
w_centered = w - jnp.mean(w)
# Backward: J.T w
def compute_vjp(w_el, w_val):
_, vjp_fun = jax.vjp(lambda p: single_local_energy(w_el, p), params)
return vjp_fun(w_val)[0]
per_walker_grads = vmap_fn(compute_vjp)(walkers, w_centered)
u = jax.tree_util.tree_map(lambda x: jnp.sum(x, axis=0), per_walker_grads)
n_walkers = walkers.shape[0]
return jax.tree_util.tree_map(lambda x: 2.0 * x / n_walkers, u)
else:
raise ValueError(f"Unknown curvature type: {self.curvature_type}")
def damped_mvp(v):
mvp_val = mvp(v)
return jax.tree_util.tree_map(lambda x, y: x + self.damping * y, mvp_val, v)
rhs = jax.tree_util.tree_map(lambda x: -x, grads)
delta, info = spla.cg(
damped_mvp,
rhs,
maxiter=self.maxiter
)
lr = self.learning_rate(state) if callable(self.learning_rate) else self.learning_rate
new_params = jax.tree_util.tree_map(lambda p, d: p + lr * d, params, delta)
return new_params, state + 1, {"loss": loss, "aux": aux_data, "lr": lr}
[docs]
def create_optimizer(optimizer_type, learning_rate, opt_kwargs=None):
"""Create an optimizer based on specified type and parameters.
Args:
optimizer_type: "adam", "sgd", "rmsprop", "lion", or "newton"
learning_rate: Initial learning rate (float) or an optax schedule (callable).
opt_kwargs: Additional optimizer parameters.
For optax optimizers: can include 'decay_rate' and 'transition_steps'
to customize the default 1/(1+t) schedule.
For Newton: can include 'min_learning_rate' (default 0.01).
"""
if opt_kwargs is None:
opt_kwargs = {}
merged_kwargs = {**opt_kwargs}
if not callable(learning_rate):
def schedule_lr(step):
decay_rate = merged_kwargs.get("decay_rate", 1.0)
transition_steps = merged_kwargs.get("transition_steps", 100)
return learning_rate / (1.0 + (step / transition_steps) * decay_rate)
else:
schedule_lr = learning_rate
if optimizer_type.lower() == "adam":
return optax.chain(
optax.scale_by_adam(
b1=merged_kwargs.get("b1", 0.9),
b2=merged_kwargs.get("b2", 0.999),
eps=merged_kwargs.get("eps", 1e-8)
),
optax.scale_by_learning_rate(schedule_lr),
)
elif optimizer_type.lower() == "sgd":
return optax.chain(
optax.sgd(learning_rate=1.0), # Rescaled by schedule below
optax.scale_by_learning_rate(schedule_lr),
)
elif optimizer_type.lower() == "rmsprop":
return optax.chain(
optax.scale_by_rms(decay=merged_kwargs.get("decay", 0.9), eps=merged_kwargs.get("eps", 1e-8)),
optax.scale_by_learning_rate(schedule_lr),
)
elif optimizer_type.lower() == "lion":
# Lion currently doesn't easily chain with custom schedules in this simple way if using optax.lion
# but we can wrap it if needed. For now using constant or optax native.
return optax.lion(learning_rate=schedule_lr, b1=merged_kwargs.get("b1", 0.9), b2=merged_kwargs.get("b2", 0.99))
elif optimizer_type.lower() == "newton":
if "value_and_grad_func" not in merged_kwargs:
raise ValueError("Newton optimizer requires value_and_grad_func in opt_kwargs")
min_lr = merged_kwargs.get("min_learning_rate", 0.01)
def newton_schedule(step):
current_lr = schedule_lr(step) if callable(schedule_lr) else schedule_lr
return jnp.maximum(current_lr, min_lr)
return NewtonOptimizer(
value_and_grad_func=merged_kwargs["value_and_grad_func"],
learning_rate=newton_schedule,
damping=merged_kwargs.get("damping", 1e-3),
maxiter=merged_kwargs.get("maxiter", 100),
curvature_type=merged_kwargs.get("curvature", "fisher"),
max_vmap_batch_size=merged_kwargs.get("max_vmap_batch_size", 0),
solver=merged_kwargs.get("solver", "exact"),
solve_kwargs=merged_kwargs.get("solve_kwargs", None),
jacobian_sample_size=merged_kwargs.get("jacobian_sample_size", 0),
clip_multiplier=merged_kwargs.get("clip_multiplier", 5.0),
jac_row_clip_multiplier=merged_kwargs.get("jac_row_clip_multiplier", 0.0),
max_delta_norm=merged_kwargs.get("max_delta_norm", None),
mesh=merged_kwargs.get("mesh", None),
)
else:
raise ValueError(f"Unsupported optimizer type: {optimizer_type}")
[docs]
def create_gradient_mask(ansatz, params, frozen_params):
"""Create a trainable-leaf mask for the combined parameter PyTree.
Assumes params = [jastrow_params, linear_coeffs]. The mask is applied
only to the jastrow_params part based on frozen_params identifiers.
The linear_coeffs part of the mask is always True (not frozen).
Args:
ansatz: The wavefunction ansatz object.
params: The combined parameters PyTree [jastrow_params, linear_coeffs].
frozen_params: A list of identifiers (int index or str name/type)
for Jastrow factors whose parameters should be frozen.
Returns:
A boolean PyTree with the same structure as params. ``False`` leaves
are frozen and ``True`` leaves remain trainable. Returns ``None``
when no parameters are frozen.
"""
if not frozen_params:
return None
if not isinstance(params, (list, tuple)) or len(params) != 2:
raise ValueError("`params` must be a list or tuple: [jastrow_params, linear_coeffs]")
jastrow_params = params[0]
linear_coeffs = params[1]
logger.info(f"Creating gradient mask for frozen Jastrow parameters: {frozen_params}")
jastrows = ansatz.jastrow.jastrows
if not isinstance(jastrow_params, (list, tuple)) or len(jastrow_params) != len(jastrows):
raise TypeError(
f"Jastrow params structure (length {len(jastrow_params)}) does not "
f"match jastrows (length {len(jastrows)})"
)
frozen_indices = set()
for identifier in frozen_params:
if isinstance(identifier, bool) or not isinstance(identifier, (int, str)):
raise TypeError("frozen_params identifiers must be integer indices or strings")
if isinstance(identifier, int):
matches = [identifier] if 0 <= identifier < len(jastrows) else []
else:
matches = [
i
for i, jastrow in enumerate(jastrows)
if identifier
in (jastrow.__class__.__name__, getattr(jastrow, "name", None))
]
if not matches:
raise ValueError(f"Unknown frozen Jastrow parameter: {identifier!r}")
frozen_indices.update(matches)
jastrow_mask_leaves = []
for i, (param_pytree, jastrow) in enumerate(zip(jastrow_params, jastrows)):
should_freeze = i in frozen_indices
if should_freeze:
logger.info(
f" Freezing Jastrow {i}: type={jastrow.__class__.__name__}, "
f"name={getattr(jastrow, 'name', None)}"
)
jastrow_mask_leaves.append(
tree_map(lambda _: not should_freeze, param_pytree)
)
jastrow_mask = (
tuple(jastrow_mask_leaves)
if isinstance(jastrow_params, tuple)
else jastrow_mask_leaves
)
mask = (jastrow_mask, tree_map(lambda _: True, linear_coeffs))
return mask if isinstance(params, tuple) else list(mask)
[docs]
def apply_gradient_mask(grads, mask):
"""Apply gradient mask to gradients to freeze parameters.
Args:
grads: The gradient PyTree
mask: The mask PyTree created by create_gradient_mask
Returns:
A PyTree with the same structure as grads, where gradients for frozen
parameters are set to zero.
"""
if mask is None:
return grads
return tree_map(
lambda gradient, trainable: jnp.where(
trainable, gradient, jnp.zeros_like(gradient)
),
grads,
mask,
)