{
  "$schema": "https://ambiqai.github.io/helia-ui/schema/reference-model-1.json",
  "generatedFrom": {
    "sourceCommit": "f341fb11f7f5d77a4974ba8273c4cd55d67c21f0",
    "tool": "pyref",
    "version": "1.7.3"
  },
  "language": "python",
  "modules": [
    {
      "description": "# Distiller Trainer API\n\nThis module contains the implementation of a distiller trainer that can be used to train a student model\n\n**Classes**\n\n| Name | Description |\n| --- | --- |\n| `Distiller` | A trainer for distillation |",
      "name": "distiller",
      "path": "helia_edge.trainers.distiller",
      "submodules": [],
      "summary": "Distiller Trainer API",
      "symbols": [
        {
          "description": "Train a student using target labels and a teacher's softened predictions.\n\ncall() returns the student's output. The distillation objective evaluates\nthe teacher with training=False and applies temperature-scaled softmax\nalong axis 1, so predictions must have their class dimension on that axis.\nConfigure the optimizer, two loss functions and mixing weight with compile().\nThe objective does not apply sample_weight.",
          "examples": [],
          "id": "helia_edge.trainers.distiller.Distiller",
          "kind": "class",
          "language": "python",
          "members": [
            {
              "description": "",
              "examples": [],
              "id": "helia_edge.trainers.distiller.Distiller.teacher",
              "kind": "attribute",
              "language": "python",
              "members": [],
              "name": "teacher",
              "params": [],
              "raises": [],
              "returns": [],
              "signature": "teacher: keras.models.Model = teacher",
              "source": {
                "line": 36,
                "path": "helia_edge/trainers/distiller.py",
                "url": "https://github.com/AmbiqAI/helia-edge/blob/f341fb11f7f5d77a4974ba8273c4cd55d67c21f0/helia_edge/trainers/distiller.py#L36"
              },
              "summary": ""
            },
            {
              "description": "",
              "examples": [],
              "id": "helia_edge.trainers.distiller.Distiller.student",
              "kind": "attribute",
              "language": "python",
              "members": [],
              "name": "student",
              "params": [],
              "raises": [],
              "returns": [],
              "signature": "student: keras.models.Model = student",
              "source": {
                "line": 37,
                "path": "helia_edge/trainers/distiller.py",
                "url": "https://github.com/AmbiqAI/helia-edge/blob/f341fb11f7f5d77a4974ba8273c4cd55d67c21f0/helia_edge/trainers/distiller.py#L37"
              },
              "summary": ""
            },
            {
              "description": "Configure the distiller.",
              "examples": [],
              "id": "helia_edge.trainers.distiller.Distiller.compile",
              "kind": "method",
              "language": "python",
              "members": [],
              "name": "compile",
              "params": [
                {
                  "description": "Keras optimizer for the student weights",
                  "name": "optimizer",
                  "type": "keras.optimizers.Optimizer"
                },
                {
                  "description": "Keras metrics for evaluation",
                  "name": "metrics",
                  "type": "list[keras.metrics.Metric]"
                },
                {
                  "description": "Loss function of difference between student\npredictions and ground-truth",
                  "name": "student_loss_fn",
                  "type": "keras.losses.Loss"
                },
                {
                  "description": "Loss function of difference between soft\nstudent predictions and soft teacher predictions",
                  "name": "distillation_loss_fn",
                  "type": "keras.losses.Loss"
                },
                {
                  "default": "0.1",
                  "description": "weight to student_loss_fn and 1-alpha to distillation_loss_fn. Defaults to 0.1.",
                  "name": "alpha",
                  "type": "float"
                },
                {
                  "default": "3",
                  "description": "Temperature for softening probability distributions. Defaults to 3.",
                  "name": "temperature",
                  "type": "float"
                }
              ],
              "raises": [],
              "returns": [],
              "signature": "compile(\n    optimizer: keras.optimizers.Optimizer,\n    metrics: list[keras.metrics.Metric],\n    student_loss_fn: keras.losses.Loss,\n    distillation_loss_fn: keras.losses.Loss,\n    alpha: float = 0.1,\n    temperature: float = 3,\n)",
              "source": {
                "line": 39,
                "path": "helia_edge/trainers/distiller.py",
                "url": "https://github.com/AmbiqAI/helia-edge/blob/f341fb11f7f5d77a4974ba8273c4cd55d67c21f0/helia_edge/trainers/distiller.py#L39"
              },
              "summary": "Configure the distiller."
            },
            {
              "description": "",
              "examples": [],
              "id": "helia_edge.trainers.distiller.Distiller.compute_loss",
              "kind": "method",
              "language": "python",
              "members": [],
              "name": "compute_loss",
              "params": [],
              "raises": [],
              "returns": [],
              "signature": "compute_loss(x=None, y=None, y_pred=None, sample_weight=None, allow_empty=False)",
              "source": {
                "line": 66,
                "path": "helia_edge/trainers/distiller.py",
                "url": "https://github.com/AmbiqAI/helia-edge/blob/f341fb11f7f5d77a4974ba8273c4cd55d67c21f0/helia_edge/trainers/distiller.py#L66"
              },
              "summary": ""
            },
            {
              "description": "",
              "examples": [],
              "id": "helia_edge.trainers.distiller.Distiller.call",
              "kind": "method",
              "language": "python",
              "members": [],
              "name": "call",
              "params": [],
              "raises": [],
              "returns": [],
              "signature": "call(x)",
              "source": {
                "line": 78,
                "path": "helia_edge/trainers/distiller.py",
                "url": "https://github.com/AmbiqAI/helia-edge/blob/f341fb11f7f5d77a4974ba8273c4cd55d67c21f0/helia_edge/trainers/distiller.py#L78"
              },
              "summary": ""
            }
          ],
          "name": "Distiller",
          "params": [
            {
              "description": "Keras model whose predictions are returned and scored.",
              "name": "student",
              "type": "keras.models.Model"
            },
            {
              "description": "Keras model providing the reference predictions.",
              "name": "teacher",
              "type": "keras.models.Model"
            }
          ],
          "raises": [],
          "returns": [],
          "signature": "Distiller(student: keras.models.Model, teacher: keras.models.Model)",
          "source": {
            "line": 16,
            "path": "helia_edge/trainers/distiller.py",
            "url": "https://github.com/AmbiqAI/helia-edge/blob/f341fb11f7f5d77a4974ba8273c4cd55d67c21f0/helia_edge/trainers/distiller.py#L16"
          },
          "summary": "Train a student using target labels and a teacher's softened predictions."
        }
      ]
    }
  ],
  "name": "helia_edge"
}
