Skip to content
heliaEDGE
User guide
HELIA

Training and callbacks

Use normal Keras compile() and fit() for supervised models. heliaEDGE adds components you can select individually: architectures, augmentation, task metrics, a progress callback and specialized trainers.

Choose a loss that matches the model’s outputs and your target encoding. Use the same deterministic preprocessing for validation and deployment, and enable random augmentation only where intended.

For the sequence model in the first-model walkthrough, targets contain one integer class ID per time step. Compile it with a loss that consumes logits:

import keras
model.compile(
optimizer="adam",
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
metrics=[keras.metrics.SparseCategoricalAccuracy()],
)

This assumes the walkthrough’s (batch, 240, 2) output and integer targets shaped (batch, 240). Use a different objective for window classification, regression or one-hot targets.

For record-based data, helia_edge.data reads records with Grain and hands the batches to a TensorFlow dataset or a Torch data loader. The generator helpers create TensorFlow datasets only.

Use TQDMProgressBar with a compiled Keras model and your training data:

from helia_edge.callbacks import TQDMProgressBar
history = model.fit(
train_x,
train_y,
validation_data=(validation_x, validation_y),
epochs=10,
callbacks=[TQDMProgressBar()],
verbose=0,
)

model and the arrays come from your task. Keras’s own callbacks provide checkpointing, early stopping and learning-rate scheduling. Add EDGE’s progress callback alongside the Keras callbacks your workflow needs.

Add checkpointing and early stopping

Replace the callbacks list above with:

callbacks = [
TQDMProgressBar(),
keras.callbacks.ModelCheckpoint("best.keras", monitor="val_loss", save_best_only=True),
keras.callbacks.EarlyStopping(monitor="val_loss", patience=5, restore_best_weights=True),
]

These callbacks require validation data to produce val_loss. Preserve preprocessing and architecture settings alongside the saved model.

Training approach Components Where to start
Masked reconstruction MaskedAutoencoder and patch layers Masked-autoencoder guide
Contrastive representation learning ContrastiveTrainer, SimCLRTrainer, SimCLR loss Trainer API
Knowledge distillation Distiller Distiller API
Discrete latent representations GSAutoencoder, VQAutoencoder and quantizer layers Trainer API

Backend support is component-specific. Masked-autoencoder paths have separate TensorFlow/Torch checks; contrastive training remains TensorFlow-specific in this source version. Consult Backend support before adopting a trainer.

Keep the data split, preprocessing configuration and architecture parameters with the trained model so the result can be reproduced.