pytc.vmc.optimizer.NewtonOptimizer

class pytc.vmc.optimizer.NewtonOptimizer(value_and_grad_func, learning_rate, damping=0.001, maxiter=100, curvature_type='fisher', max_vmap_batch_size=0, solver='exact', solve_kwargs=None, jacobian_sample_size=0, clip_multiplier=5.0)[source]

Bases: object

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

Methods

__delattr__

Implement delattr(self, name).

__dir__

Default dir() implementation.

__eq__

Return self==value.

__format__

Default object formatter.

__ge__

Return self>=value.

__getattribute__

Return getattr(self, name).

__getstate__

Helper for pickle.

__gt__

Return self>value.

__hash__

Return hash(self).

__init__

__init_subclass__

This method is called when a class is subclassed.

__le__

Return self<=value.

__lt__

Return self<value.

__ne__

Return self!=value.

__new__

__reduce__

Helper for pickle.

__reduce_ex__

Helper for pickle.

__repr__

Return repr(self).

__setattr__

Implement setattr(self, name, value).

__sizeof__

Size of object in memory, in bytes.

__str__

Return str(self).

__subclasshook__

Abstract classes can override this to customize issubclass().

_flatten_jacobian

Flatten a per-walker Jacobian pytree to an (N, P) matrix.

_get_effective_batch_size

Return a batch size compatible with the current execution mode.

_get_unbatched_vmap

Return a per-batch vmap without nested folx batching.

_get_vmap

Return the appropriate vmap implementation.

_pad_walkers_to_batch_multiple

Pad walkers along axis 0 so they reshape cleanly into batches.

_slice_walkers

Slice a walker pytree along the walker dimension.

init

step

Attributes

__annotations__

__dict__

__doc__

__module__

__weakref__

list of weak references to the object

init(params, rng, batch)[source]
step(params, state, rng, batch, global_step_int=None)[source]