essos.losses¶
Classes¶
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)¶