Skip to content
heliaEDGE
Reference
HELIA

steps

Training steps that dispatch on the active Keras backend.

Machine-readable model

  • NotSupportedclassThe active Keras backend is not supported by this trainer.
  • require_backendfunctionReturn the active backend, or raise NotSupported naming feature and supported.
  • no_gradfunctionDisable gradient tracking on Torch; a no-op on other backends.
  • gradient_stepfunctionDifferentiate lossfn with the active backend and apply model.optimizer once.
function

Return the active backend, or raise NotSupported naming feature and supported.

helia_edge/trainers/steps.py:16

require_backend(feature: str, supported: Sequence[str]) -> str

Return the active backend, or raise NotSupported naming feature and supported.

function

no_grad

Python

Disable gradient tracking on Torch; a no-op on other backends.

helia_edge/trainers/steps.py:24

no_grad() -> Iterator[None]

Disable gradient tracking on Torch; a no-op on other backends.

function

Differentiate lossfn with the active backend and apply model.optimizer once.

helia_edge/trainers/steps.py:37

gradient_step(
model: keras.Model,
loss_fn: Callable[[], tuple[Any, ...]],
variables: Sequence[Any] | None = None,
) -> tuple[Any, ...]

Differentiate loss_fn with the active backend and apply model.optimizer once.

Parameters of gradient_step
NameTypeDefaultDescription
modelkeras.ModelRequiredCompiled model whose optimizer applies the update.
loss_fnCallable[[], tuple[Any, ...]]RequiredComputes ``(loss, *outputs)`` from the model's current weights.
variablesSequence[Any] | NoneNoneVariables to update; when None, the model's trainable weights after ``loss_fn`` runs, so variables created by a first (building) call are included. Variables without a gradient are skipped.
Returns of gradient_step
ValueTypeDescription
tupletuple[Any, ...]What ``loss_fn`` returned.
Errors raised by gradient_step
TypeDescription
NotSupportedOn backends other than TensorFlow and Torch.