class
Distiller
PythonTrain a student using target labels and a teacher's softened predictions.
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
| Name | Type | Default | Description |
|---|---|---|---|
student | keras.models.Model | Required | Keras model whose predictions are returned and scored. |
teacher | keras.models.Model | Required | Keras model providing the reference predictions. |
attribute
teacher
Pythonteacher: keras.models.Model = teacherattribute
student
Pythonstudent: keras.models.Model = studentmethod
compile
PythonConfigure the distiller.
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
| Name | Type | Default | Description |
|---|---|---|---|
optimizer | keras.optimizers.Optimizer | Required | Keras optimizer for the student weights |
metrics | list[keras.metrics.Metric] | Required | Keras metrics for evaluation |
student_loss_fn | keras.losses.Loss | Required | Loss function of difference between student predictions and ground-truth |
distillation_loss_fn | keras.losses.Loss | Required | Loss function of difference between soft student predictions and soft teacher predictions |
alpha | float | 0.1 | weight to student_loss_fn and 1-alpha to distillation_loss_fn. Defaults to 0.1. |
temperature | float | 3 | Temperature for softening probability distributions. Defaults to 3. |
method
compute_loss
Pythoncompute_loss(x=None, y=None, y_pred=None, sample_weight=None, allow_empty=False)method
call
Pythoncall(x)