essos.losses

Classes

base_loss

custom_loss

composite_loss

Module Contents

class essos.losses.base_loss
losses
_dependencies
_dependencies_buffer = None
_starting_dofs = None
_dofs_to_pytree = None
clear_cache()
property dependencies
property dependencies_buffer
__add__(other)
__iter__()
abstractmethod __mul__(other)
__rmul__(other)
class essos.losses.custom_loss(fun, *args_names, **kwargs)

Bases: base_loss

fun
args_names = ()
kwargs
_dofs_to_args = None
clear_cache()
_ensure_unravelers()
property starting_dofs
property dofs_to_pytree
__call__(dofs: jax.numpy.ndarray) → float
call_pytree(dofs_pytree) → float
grad(dofs: jax.numpy.ndarray) → jax.numpy.ndarray
value_and_grad(dofs: jax.numpy.ndarray)
grad_pytree(dofs_pytree) → dict
__mul__(other)
class essos.losses.composite_loss(losses: list)

Bases: base_loss

losses
property dependencies
property starting_dofs
property dofs_to_pytree
__call__(dofs: jax.numpy.ndarray) → float
grad(dofs: jax.numpy.ndarray) → jax.numpy.ndarray