{ "cells": [ { "cell_type": "markdown", "id": "7fb27b941602401d91542211134fc71a", "metadata": {}, "source": [ "# Training pipeline with callbacks\n", "\n", "- `VQCClassifier` trained with `EarlyStopping` and `ModelCheckpoint`.\n", "- Both run on both backends, unlike a Lightning callback.\n", "- Inspect what each callback recorded, resume the run from its last checkpoint, then reload the fitted preprocessing for new rows.\n", "- Pennylane backend throughout." ] }, { "cell_type": "markdown", "id": "1b64d729", "metadata": {}, "source": [ "### 1. Data\n", "\n", "- `make_moons`, 200 samples.\n", "- `normalize=\"minmax\"`, fit on train only.\n", "- `AngleEmbedding` prescaling maps it to [0, pi] in `setup()`." ] }, { "cell_type": "code", "id": "13433722", "metadata": { "ExecuteTime": { "end_time": "2026-09-19T10:45:44.483984752Z", "start_time": "2026-09-19T10:45:41.504380992Z" } }, "source": [ "from sklearn.datasets import make_moons\n", "\n", "import pyqit\n", "from pyqit import DataModule, Trainer\n", "from pyqit.ansatzes import SELAnsatz\n", "from pyqit.core import AngleEmbedding, EarlyStopping, ModelCheckpoint\n", "from pyqit.models import VQCClassifier\n", "\n", "pyqit.set_seed(42)\n", "\n", "X, y = make_moons(n_samples=200, noise=0.1, random_state=0)\n", "dm = DataModule(X, y, normalize=\"minmax\", batch_size=16, seed=42)" ], "outputs": [], "execution_count": 1 }, { "cell_type": "markdown", "id": "b936bced", "metadata": {}, "source": [ "### 2. Trainer with both callbacks\n", "\n", "- `EarlyStopping` on `val_loss`, patience 5. Stops once it stops improving.\n", "- `ModelCheckpoint` saves the best and last epoch, restores best weights after training.\n", "- Kept as named variables, not inlined, to read their state after `fit`.\n", "- `loss_fn=\"cross_entropy\"`." ] }, { "cell_type": "code", "id": "0dc28b7b", "metadata": { "ExecuteTime": { "end_time": "2026-09-19T10:45:51.552221241Z", "start_time": "2026-09-19T10:45:44.492386434Z" } }, "source": [ "model = VQCClassifier(n_qubits=4, n_layers=3, ansatz=SELAnsatz, encoder=AngleEmbedding)\n", "\n", "early_stop = EarlyStopping(monitor=\"val_loss\", patience=5)\n", "checkpoint = ModelCheckpoint(dirpath=\"ckpts\", save_best=True, save_last=True)\n", "\n", "trainer = Trainer(\n", " max_epochs=60,\n", " loss_fn=\"cross_entropy\",\n", " callbacks=[early_stop, checkpoint],\n", " verbose=1,\n", ")\n", "history = trainer.fit(model, dm)" ], "outputs": [ { "data": { "text/plain": [ "\u001B[1;36m[\u001B[0m\u001B[1;36mTrainer\u001B[0m\u001B[1;36m]\u001B[0m Starting \u001B[32mpennylane\u001B[0m backend | \u001B[1;36m60\u001B[0m epochs | \u001B[33mlr\u001B[0m=\u001B[1;36m0\u001B[0m\u001B[1;36m.01\u001B[0m\n", "\n" ], "text/html": [ "
[Trainer] Starting pennylane backend | 60 epochs | lr=0.01\n", "\n", "\n" ] }, "metadata": {}, "output_type": "display_data" }, { "data": { "text/plain": [ "Output()" ], "application/vnd.jupyter.widget-view+json": { "version_major": 2, "version_minor": 0, "model_id": "ba89de78f1b844f4a3999bce3fc3c7d5" } }, "metadata": {}, "output_type": "display_data" }, { "data": { "text/plain": [ "\u001B[1;33m[\u001B[0m\u001B[1;33mEarlyStopping\u001B[0m\u001B[1;33m]\u001B[0m Stopped at epoch \u001B[1;36m39\u001B[0m -- val_loss did not improve for \u001B[1;36m5\u001B[0m \u001B[1;35mepoch\u001B[0m\u001B[1m(\u001B[0ms\u001B[1m)\u001B[0m \u001B[1m(\u001B[0mbest val_loss: \u001B[1;36m0.3256\u001B[0m\u001B[1m)\u001B[0m\n" ], "text/html": [ "
[EarlyStopping] Stopped at epoch 39 -- val_loss did not improve for 5 epoch(s) (best val_loss: 0.3256)\n", "\n" ] }, "metadata": {}, "output_type": "display_data" }, { "data": { "text/plain": [], "text/html": [ "\n" ] }, "metadata": {}, "output_type": "display_data" }, { "data": { "text/plain": [ "\u001B[1;32m[\u001B[0m\u001B[1;32mCheckpoint\u001B[0m\u001B[1;32m]\u001B[0m Last epoch -> ckpts/last.npz\n" ], "text/html": [ "
[Checkpoint] Last epoch -> ckpts/last.npz\n",
"\n"
]
},
"metadata": {},
"output_type": "display_data"
},
{
"data": {
"text/plain": [
"\u001B[1;32m[\u001B[0m\u001B[1;32mCheckpoint\u001B[0m\u001B[1;32m]\u001B[0m Restored best weights from epoch \u001B[1;36m34\u001B[0m \u001B[1m(\u001B[0mval_loss: \u001B[1;36m0.3256\u001B[0m\u001B[1m)\u001B[0m\n"
],
"text/html": [
"[Checkpoint] Restored best weights from epoch 34 (val_loss: 0.3256)\n", "\n" ] }, "metadata": {}, "output_type": "display_data" }, { "data": { "text/plain": [ "\u001B[1;32m[\u001B[0m\u001B[1;32mTrainer\u001B[0m\u001B[1;32m]\u001B[0m Training complete.\n" ], "text/html": [ "
[Trainer] Training complete.\n",
"\n"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"execution_count": 2
},
{
"cell_type": "markdown",
"id": "fd565fbe",
"metadata": {},
"source": [
"### 3. What each callback recorded\n",
"\n",
"- `early_stop.stopped_epoch` and `early_stop.stopping_reason`, set once training stops early.\n",
"- `checkpoint.best_epoch`, `checkpoint.best_path`, `checkpoint.last_path`, the epoch and files `ModelCheckpoint` wrote.\n",
"- `history.best_epoch`, `history.best_score`, `history.best_metric`, the same best epoch from the run's own record."
]
},
{
"cell_type": "code",
"id": "5630cfc2",
"metadata": {
"ExecuteTime": {
"end_time": "2026-09-19T10:45:51.569949072Z",
"start_time": "2026-09-19T10:45:51.555130093Z"
}
},
"source": [
"print(\"stopped at:\", early_stop.stopped_epoch, \"|\", early_stop.stopping_reason)\n",
"print(\"best checkpoint:\", checkpoint.best_epoch, \"->\", checkpoint.best_path)\n",
"print(\"last checkpoint:\", checkpoint.last_path)\n",
"print(\"history best:\", history.best_epoch, history.best_score, history.best_metric)"
],
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"stopped at: 39 | val_loss did not improve for 5 epoch(s) (best val_loss: 0.3256)\n",
"best checkpoint: 34 -> ckpts/best.npz\n",
"last checkpoint: ckpts/last.npz\n",
"history best: 34 0.32560367617711666 val_loss\n"
]
}
],
"execution_count": 3
},
{
"cell_type": "markdown",
"id": "5cebe8c4",
"metadata": {},
"source": [
"### 4. Loss curve"
]
},
{
"cell_type": "code",
"id": "cf8a90fd",
"metadata": {
"ExecuteTime": {
"end_time": "2026-09-19T10:45:51.674453514Z",
"start_time": "2026-09-19T10:45:51.571820426Z"
}
},
"source": [
"import matplotlib.pyplot as plt\n",
"\n",
"plt.plot(history.train_loss, label=\"train_loss\")\n",
"plt.plot(history.val_loss, label=\"val_loss\")\n",
"plt.xlabel(\"epoch\")\n",
"plt.ylabel(\"loss\")\n",
"plt.legend()\n",
"plt.show()"
],
"outputs": [
{
"data": {
"text/plain": [
"[Trainer] Starting pennylane backend | 45 epochs | lr=0.01\n", "\n", "\n" ] }, "metadata": {}, "output_type": "display_data" }, { "data": { "text/plain": [ "\u001B[1;32m[\u001B[0m\u001B[1;32mCheckpoint\u001B[0m\u001B[1;32m]\u001B[0m Resumed from ckpts/last.npz at epoch \u001B[1;36m40\u001B[0m\n" ], "text/html": [ "
[Checkpoint] Resumed from ckpts/last.npz at epoch 40\n", "\n" ] }, "metadata": {}, "output_type": "display_data" }, { "data": { "text/plain": [ "Output()" ], "application/vnd.jupyter.widget-view+json": { "version_major": 2, "version_minor": 0, "model_id": "3bef087e6cde47f1a48d61e291d6130b" } }, "metadata": {}, "output_type": "display_data" }, { "data": { "text/plain": [], "text/html": [ "\n" ] }, "metadata": {}, "output_type": "display_data" }, { "data": { "text/plain": [ "\u001B[1;32m[\u001B[0m\u001B[1;32mTrainer\u001B[0m\u001B[1;32m]\u001B[0m Training complete.\n" ], "text/html": [ "
[Trainer] Training complete.\n",
"\n"
]
},
"metadata": {},
"output_type": "display_data"
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"40 epochs from the file, 5 trained now\n"
]
}
],
"execution_count": 5
},
{
"cell_type": "markdown",
"id": "8dd0d8092fe74a7c96281538738b07e2",
"metadata": {},
"source": [
"### 6. Saving the DataModule\n",
"\n",
"- The checkpoint holds the model. The fitted normalizer is the DataModule's, saved on request with `dm.save`.\n",
"- `DataModule.load(path, X_new)` attaches new rows to the saved settings, like `dm.for_prediction`, so predicting needs neither the training data nor a refit."
]
},
{
"cell_type": "code",
"id": "72eea5119410473aa328ad9291626812",
"metadata": {
"ExecuteTime": {
"end_time": "2026-09-19T10:45:52.325905550Z",
"start_time": "2026-09-19T10:45:52.303215238Z"
}
},
"source": [
"X_new, _ = make_moons(n_samples=5, noise=0.1, random_state=1)\n",
"\n",
"dm.save(\"ckpts/datamodule.pkl\")\n",
"dm_new = DataModule.load(\"ckpts/datamodule.pkl\", X_new)\n",
"\n",
"print(trainer.predict(fresh_model, dm_new))"
],
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"[0 0 1 1 1]\n"
]
}
],
"execution_count": 6
}
],
"metadata": {
"kernelspec": {
"display_name": ".venv (3.12.3)",
"language": "python",
"name": "python3"
},
"language_info": {
"name": "python",
"version": "3.12.3"
}
},
"nbformat": 4,
"nbformat_minor": 5
}