Source code for pytc.vmc.test.test_optimizer_mask

"""Tests for freezing selected Jastrow parameters during optimization."""

from types import SimpleNamespace
import unittest

import jax
jax.config.update("jax_enable_x64", True)
import jax.numpy as jnp
import numpy as np
import optax

from pytc.vmc.optimization import make_opt_update_step, optimize
from pytc.vmc.optimizer import apply_gradient_mask, create_gradient_mask


[docs] class FrozenJastrow: pass
[docs] class TrainableJastrow: name = "trainable"
[docs] def _ansatz(): jastrow = SimpleNamespace(jastrows=[FrozenJastrow(), TrainableJastrow()]) return SimpleNamespace(jastrow=jastrow)
[docs] def _params(): return [ [ {"coefficient": jnp.array([1.0, 2.0])}, {"coefficient": jnp.array([3.0])}, ], jnp.array([4.0]), ]
[docs] class TestGradientMask(unittest.TestCase):
[docs] def test_freezes_only_selected_jastrow(self): params = _params() grads = jax.tree_util.tree_map(jnp.ones_like, params) mask = create_gradient_mask(_ansatz(), params, ["FrozenJastrow"]) masked_grads = apply_gradient_mask(grads, mask) np.testing.assert_array_equal(masked_grads[0][0]["coefficient"], 0.0) np.testing.assert_array_equal(masked_grads[0][1]["coefficient"], 1.0) np.testing.assert_array_equal(masked_grads[1], 1.0)
[docs] def test_optax_step_keeps_frozen_parameters_unchanged(self): params = _params() mask = create_gradient_mask(_ansatz(), params, [0]) def loss_fn(current_params, _): leaves = jax.tree_util.tree_leaves(current_params) return sum(jnp.sum(leaf**2) for leaf in leaves), () optimizer = optax.sgd(learning_rate=0.1) opt_step = make_opt_update_step(loss_fn, optimizer, gradient_mask=mask) opt_state = optimizer.init(params) new_params, _, _, _ = opt_step( None, params, None, opt_state, jax.random.PRNGKey(0) ) np.testing.assert_array_equal( new_params[0][0]["coefficient"], params[0][0]["coefficient"] ) np.testing.assert_allclose( new_params[0][1]["coefficient"], jnp.array([2.4]) ) np.testing.assert_allclose(new_params[1], jnp.array([3.2]))
[docs] def test_lion_weight_decay_cannot_move_frozen_parameters(self): params = _params() mask = create_gradient_mask(_ansatz(), params, [0]) def loss_fn(current_params, _): leaves = jax.tree_util.tree_leaves(current_params) return sum(jnp.sum(leaf**2) for leaf in leaves), () optimizer = optax.lion(learning_rate=0.1, weight_decay=0.01) opt_step = make_opt_update_step(loss_fn, optimizer, gradient_mask=mask) opt_state = optimizer.init(params) new_params, _, _, _ = opt_step( None, params, None, opt_state, jax.random.PRNGKey(0) ) np.testing.assert_array_equal( new_params[0][0]["coefficient"], params[0][0]["coefficient"] ) self.assertFalse( np.array_equal( new_params[0][1]["coefficient"], params[0][1]["coefficient"], ) )
[docs] def test_empty_frozen_params_does_not_create_mask(self): self.assertIsNone(create_gradient_mask(_ansatz(), _params(), []))
[docs] def test_mask_preserves_tuple_parameter_structure(self): list_params = _params() params = (tuple(list_params[0]), list_params[1]) grads = jax.tree_util.tree_map(jnp.ones_like, params) mask = create_gradient_mask(_ansatz(), params, [0]) self.assertIsInstance(mask, tuple) self.assertIsInstance(mask[0], tuple) apply_gradient_mask(grads, mask)
[docs] def test_unknown_frozen_parameter_is_rejected(self): with self.assertRaisesRegex(ValueError, "Unknown frozen Jastrow parameter"): create_gradient_mask(_ansatz(), _params(), ["typo"])
[docs] def test_newton_rejects_frozen_params_before_setup(self): with self.assertRaisesRegex(NotImplementedError, "Newton optimizer"): optimize(_ansatz(), optimizer_type="newton", frozen_params=[0])
[docs] class TestPublicOptimizeFrozenParams(unittest.TestCase): """Public-path regression for the original production defect: if the gradient_mask wiring inside ``optimize()`` is removed, the Optax path silently trains frozen leaves — and until this test existed, every other test in this file still passed. This test exercises the public ``optimize(..., frozen_params=...)`` call end to end on a real ansatz and must fail the moment that wiring is dropped."""
[docs] @classmethod def setUpClass(cls): from pyscf import gto, scf from pytc.ansatz.sj import SlaterJastrow from pytc.ansatz.det import SlaterDet from pytc.jastrow import NuclearCusp, CompositeJastrow from pytc.jastrow.bha import BoysHandyAnalytical mol = gto.M(atom="H 0 0 0; H 0 0 1.8; H 0 0 3.6; H 0 0 5.4", basis="sto-3g", unit="Bohr", verbose=0) mf = scf.RHF(mol).run() det = SlaterDet.create(mol, mf.mo_coeff) ncusp = NuclearCusp.create(mol, name="ncusp") bha = BoysHandyAnalytical.create(mol) jastrow = CompositeJastrow.create([ncusp, bha]) cls.sj = SlaterJastrow.create(mol, jastrow, [det])
[docs] def test_public_optimize_keeps_frozen_factor_fixed(self): initial = [self.sj.jastrow.init_params(), jnp.ones(1)] frozen_before = jax.tree_util.tree_map( lambda x: np.asarray(x).copy(), initial[0][0]) result = optimize( self.sj, n_walkers=16, n_steps=8, burn_in_steps=8, n_opt_steps=3, learning_rate=0.05, optimizer_type="adam", frozen_params=[0], params=initial, adaptive_step_size=False, ) final = result["params"][-1] # The frozen factor's parameters must be bitwise unchanged after # real optimization steps on the public path. frozen_after = jax.tree_util.tree_map(np.asarray, final[0][0]) for key in frozen_before: np.testing.assert_array_equal( frozen_after[key], frozen_before[key], err_msg=f"frozen leaf {key!r} moved on the public optimize() path") # The trainable factor and the linear coefficients must have moved # (otherwise the test would also pass with optimization broken). bha_moved = any( not np.array_equal(np.asarray(final[0][1][k]), np.asarray(initial[0][1][k])) for k in final[0][1]) self.assertTrue(bha_moved, "trainable Jastrow factor did not update") self.assertFalse( bool(np.array_equal(np.asarray(final[1]), np.asarray(initial[1]))), "linear coefficients did not update")
if __name__ == "__main__": unittest.main()