Skip to content
heliaEDGE
Reference
HELIA

distiller

This module contains the implementation of a distiller trainer that can be used to train a student model

Classes

Name Description
Distiller A trainer for distillation

Machine-readable model

  • DistillerclassTrain a student using target labels and a teacher's softened predictions.
class

Distiller

Python

Train a student using target labels and a teacher's softened predictions.

helia_edge/trainers/distiller.py:16

Distiller(student: keras.models.Model, teacher: keras.models.Model)

Train a student using target labels and a teacher’s softened predictions.

call() returns the student’s output. The distillation objective evaluates the teacher with training=False and applies temperature-scaled softmax along axis 1, so predictions must have their class dimension on that axis. Configure the optimizer, two loss functions and mixing weight with compile(). The objective does not apply sample_weight.

Parameters of Distiller
NameTypeDefaultDescription
studentkeras.models.ModelRequiredKeras model whose predictions are returned and scored.
teacherkeras.models.ModelRequiredKeras model providing the reference predictions.
method

compile

Python

Configure the distiller.

helia_edge/trainers/distiller.py:39

compile(
optimizer: keras.optimizers.Optimizer,
metrics: list[keras.metrics.Metric],
student_loss_fn: keras.losses.Loss,
distillation_loss_fn: keras.losses.Loss,
alpha: float = 0.1,
temperature: float = 3,
)

Configure the distiller.

Parameters of compile
NameTypeDefaultDescription
optimizerkeras.optimizers.OptimizerRequiredKeras optimizer for the student weights
metricslist[keras.metrics.Metric]RequiredKeras metrics for evaluation
student_loss_fnkeras.losses.LossRequiredLoss function of difference between student predictions and ground-truth
distillation_loss_fnkeras.losses.LossRequiredLoss function of difference between soft student predictions and soft teacher predictions
alphafloat0.1weight to student_loss_fn and 1-alpha to distillation_loss_fn. Defaults to 0.1.
temperaturefloat3Temperature for softening probability distributions. Defaults to 3.