OptimHistory.update_position_history()

OptimHistory.update_position_history()#

static OptimHistory.update_position_history(i, position_history, position)[source]#

Writes a position into a position history at one epoch index.

Parameters:
Return type:

Position (dict[str, Any])

Returns:

Position – Updated position history.

Examples

>>> import jax.numpy as jnp
>>> from liesel.optim import OptimHistory
>>> from liesel.optim.types import Position
>>> position = Position({"theta": jnp.array([1.0, 2.0])})
>>> history = OptimHistory.init_position_history(position, epochs=2)
>>> updated = OptimHistory.update_position_history(1, history, position)
>>> updated["theta"].tolist()
[[0.0, 0.0], [1.0, 2.0]]