dashai-frankenstein 0.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,36 @@
1
+ .venv/
2
+ uv.lock
3
+ wandb/
4
+ output/
5
+ dist/
6
+ **/__pycache__/
7
+ **/*.pyc
8
+ **/temp_data/
9
+ **/logs/
10
+
11
+ **/checkpoints
12
+ **/es_redpajama_50k.model
13
+ **/es_redpajama_50k.vocab
14
+ **/nohup.out**
15
+ **/training_metrics.csv
16
+
17
+ docs/**/*.aux
18
+ docs/**/*.log
19
+ docs/**/*.out
20
+ docs/**/*.toc
21
+ docs/**/*.bbl
22
+ docs/**/*.blg
23
+ docs/_build/
24
+ docs/pdoc/
25
+
26
+ *.csv
27
+
28
+ .opencode/node_modules/
29
+
30
+ # Exhaustive end-to-end harness workspace (generated tokenizers, checkpoints, deploys)
31
+ full_tests/tmp/
32
+ full_tests/runs/
33
+ full_tests/**/__pycache__/
34
+ full_tests/**/*.pyc
35
+ full_tests/**/*.model
36
+ full_tests/**/*.vocab
@@ -0,0 +1,61 @@
1
+ Metadata-Version: 2.5
2
+ Name: dashai-frankenstein
3
+ Version: 0.1.0
4
+ Summary: DashAI plugin that registers Frankenstein Transformer model classes (MLM encoder, causal decoder, ViT) as DashAI components.
5
+ Author-email: Erick Merino <erickfmm@gmail.com>
6
+ License: MIT
7
+ Requires-Python: >=3.10
8
+ Requires-Dist: datasets>=3.0.0
9
+ Requires-Dist: frankenstein-transformer>=1.1.0
10
+ Requires-Dist: numpy<3.0,>=1.26.0
11
+ Requires-Dist: pyyaml>=6.0.1
12
+ Requires-Dist: torch>=2.0
13
+ Requires-Dist: transformers<5.0.0,>=4.45.0
14
+ Provides-Extra: sbert
15
+ Requires-Dist: sentence-transformers>=3.3.0; extra == 'sbert'
16
+ Requires-Dist: sentencepiece>=0.2.0; extra == 'sbert'
17
+ Description-Content-Type: text/markdown
18
+
19
+ # dashai-frankenstein
20
+
21
+ A [DashAI](https://docs.dash-ai.com/) plugin that registers
22
+ [Frankenstein Transformer](https://github.com/erickfmm/frankenstein-transformer)
23
+ model classes as DashAI components, so end users can train, evaluate, predict,
24
+ save, and load them from the DashAI UI.
25
+
26
+ ## Components registered
27
+
28
+ | Entry point | Class | DashAI base | Binds to task |
29
+ |---|---|---|---|
30
+ | `frankenstein_mlm` | `FrankensteinMLMModel` | `BaseModel` | `TextClassificationTask` |
31
+ | `frankenstein_decoder` | `FrankensteinDecoderModel` | `BaseGenerativeModel` | `TextToTextGenerationTask` |
32
+ | `frankenstein_vit_cls` | `FrankensteinViTClassifier` | `BaseModel` | `ImageClassificationTask` |
33
+ | `frankenstein_vit_seg` | `FrankensteinViTSegmenter` | `BaseModel` | `SegmentationTask` |
34
+ | `segmentation_task` | `SegmentationTask` | `BaseTask` | (new task provided by this plugin) |
35
+
36
+ ## Schema (v1: passthrough YAML)
37
+
38
+ Each model exposes a minimal pydantic schema. The primary field,
39
+ `frankenstein_yaml`, is a string containing a full Frankenstein training YAML,
40
+ validated by Frankenstein's own config loader + JSON Schema (the Frankenstein
41
+ schema remains the single source of truth). A `preset` dropdown is populated
42
+ from Frankenstein's bundled `configs/*.yaml`; `device`, `batch_size`, and
43
+ `num_epochs` are convenience overrides merged into the YAML before validation.
44
+
45
+ ## Install
46
+
47
+ ```bash
48
+ pip install dashai-frankenstein # from PyPI once published
49
+ # or, from this repo:
50
+ pip install -e ./dashai-frankenstein
51
+ ```
52
+
53
+ DashAI discovers the plugin via the `dashai.plugins` entry-points group on
54
+ startup — no DashAI source edits required.
55
+
56
+ ## Architecture
57
+
58
+ See `docs/dashai-plugin-audit.md` in the Frankenstein repo for the full
59
+ integration design (§5 component designs, §6 phased plan, §7 Frankenstein
60
+ changes). This package is the Phase 1–3 adapter layer; it consumes the
61
+ Frankenstein engine API (`src.engine`) added in Phase 0.
@@ -0,0 +1,43 @@
1
+ # dashai-frankenstein
2
+
3
+ A [DashAI](https://docs.dash-ai.com/) plugin that registers
4
+ [Frankenstein Transformer](https://github.com/erickfmm/frankenstein-transformer)
5
+ model classes as DashAI components, so end users can train, evaluate, predict,
6
+ save, and load them from the DashAI UI.
7
+
8
+ ## Components registered
9
+
10
+ | Entry point | Class | DashAI base | Binds to task |
11
+ |---|---|---|---|
12
+ | `frankenstein_mlm` | `FrankensteinMLMModel` | `BaseModel` | `TextClassificationTask` |
13
+ | `frankenstein_decoder` | `FrankensteinDecoderModel` | `BaseGenerativeModel` | `TextToTextGenerationTask` |
14
+ | `frankenstein_vit_cls` | `FrankensteinViTClassifier` | `BaseModel` | `ImageClassificationTask` |
15
+ | `frankenstein_vit_seg` | `FrankensteinViTSegmenter` | `BaseModel` | `SegmentationTask` |
16
+ | `segmentation_task` | `SegmentationTask` | `BaseTask` | (new task provided by this plugin) |
17
+
18
+ ## Schema (v1: passthrough YAML)
19
+
20
+ Each model exposes a minimal pydantic schema. The primary field,
21
+ `frankenstein_yaml`, is a string containing a full Frankenstein training YAML,
22
+ validated by Frankenstein's own config loader + JSON Schema (the Frankenstein
23
+ schema remains the single source of truth). A `preset` dropdown is populated
24
+ from Frankenstein's bundled `configs/*.yaml`; `device`, `batch_size`, and
25
+ `num_epochs` are convenience overrides merged into the YAML before validation.
26
+
27
+ ## Install
28
+
29
+ ```bash
30
+ pip install dashai-frankenstein # from PyPI once published
31
+ # or, from this repo:
32
+ pip install -e ./dashai-frankenstein
33
+ ```
34
+
35
+ DashAI discovers the plugin via the `dashai.plugins` entry-points group on
36
+ startup — no DashAI source edits required.
37
+
38
+ ## Architecture
39
+
40
+ See `docs/dashai-plugin-audit.md` in the Frankenstein repo for the full
41
+ integration design (§5 component designs, §6 phased plan, §7 Frankenstein
42
+ changes). This package is the Phase 1–3 adapter layer; it consumes the
43
+ Frankenstein engine API (`src.engine`) added in Phase 0.
@@ -0,0 +1,38 @@
1
+ [build-system]
2
+ requires = ["hatchling"]
3
+ build-backend = "hatchling.build"
4
+
5
+ [project]
6
+ name = "dashai-frankenstein"
7
+ version = "0.1.0"
8
+ description = "DashAI plugin that registers Frankenstein Transformer model classes (MLM encoder, causal decoder, ViT) as DashAI components."
9
+ readme = "README.md"
10
+ license = { text = "MIT" }
11
+ authors = [{ name = "Erick Merino", email = "erickfmm@gmail.com" }]
12
+ requires-python = ">=3.10"
13
+ dependencies = [
14
+ # Frankenstein ships the importable ``src`` package and the engine API.
15
+ # Pin to the Phase-0 release that exposes engine.build_model / train_from_config.
16
+ "frankenstein-transformer>=1.1.0",
17
+ "torch>=2.0",
18
+ "transformers>=4.45.0,<5.0.0",
19
+ "datasets>=3.0.0",
20
+ "numpy>=1.26.0,<3.0",
21
+ "PyYAML>=6.0.1",
22
+ ]
23
+
24
+ [project.optional-dependencies]
25
+ sbert = ["sentence-transformers>=3.3.0", "sentencepiece>=0.2.0"]
26
+
27
+ # DashAI discovers plugins whose distribution name startswith("dashai").
28
+ # Each entry points to a class that inherits a DashAI base (BaseModel /
29
+ # BaseGenerativeModel / BaseTask) carrying a TYPE attribute.
30
+ [project.entry-points."dashai.plugins"]
31
+ frankenstein_mlm = "dashai_frankenstein:FrankensteinMLMModel"
32
+ frankenstein_decoder = "dashai_frankenstein:FrankensteinDecoderModel"
33
+ frankenstein_vit_cls = "dashai_frankenstein:FrankensteinViTClassifier"
34
+ frankenstein_vit_seg = "dashai_frankenstein:FrankensteinViTSegmenter"
35
+ segmentation_task = "dashai_frankenstein:SegmentationTask"
36
+
37
+ [tool.hatch.build.targets.wheel]
38
+ packages = ["src/dashai_frankenstein"]
@@ -0,0 +1,20 @@
1
+ """dashai-frankenstein — DashAI plugin for Frankenstein Transformer.
2
+
3
+ Registers Frankenstein model classes (encoder, decoder, ViT classifier/segmenter)
4
+ and a SegmentationTask as DashAI components, discovered via the
5
+ ``dashai.plugins`` entry-points group.
6
+ """
7
+ from dashai_frankenstein.models.decoder import FrankensteinDecoderModel # noqa: F401
8
+ from dashai_frankenstein.models.mlm import FrankensteinMLMModel # noqa: F401
9
+ from dashai_frankenstein.models.vit_classifier import FrankensteinViTClassifier # noqa: F401
10
+ from dashai_frankenstein.models.vit_segmenter import FrankensteinViTSegmenter # noqa: F401
11
+ from dashai_frankenstein.tasks.segmentation import SegmentationTask # noqa: F401
12
+
13
+ __all__ = [
14
+ "FrankensteinMLMModel",
15
+ "FrankensteinDecoderModel",
16
+ "FrankensteinViTClassifier",
17
+ "FrankensteinViTSegmenter",
18
+ "SegmentationTask",
19
+ ]
20
+ __version__ = "0.1.0"
@@ -0,0 +1,4 @@
1
+ """DashAI <-> Frankenstein adapters (dataset, IO, metrics)."""
2
+ from dashai_frankenstein.adapters import dataset, io, metrics # noqa: F401
3
+
4
+ __all__ = ["dataset", "io", "metrics"]
@@ -0,0 +1,281 @@
1
+ """Dataset adapter: DashAIDataset <-> Frankenstein training tensors.
2
+
3
+ A ``DashAIDataset`` is a HuggingFace ``datasets.Dataset`` wrapper. Frankenstein
4
+ models consume tensors directly. This module extracts the relevant columns,
5
+ tokenizes text (for NLP components), and yields PyTorch ``DataLoader`` batches
6
+ ready for the in-process training loop.
7
+ """
8
+ from __future__ import annotations
9
+
10
+ import logging
11
+ from typing import Any, List, Optional, Tuple
12
+
13
+ import numpy as np
14
+
15
+ log = logging.getLogger(__name__)
16
+
17
+
18
+ def find_text_column(dataset: Any, exclude: Optional[List[str]] = None) -> str:
19
+ """Return the single non-categorical text column of a DashAIDataset.
20
+
21
+ Parameters
22
+ ----------
23
+ dataset : DashAIDataset
24
+ A dataset whose ``column_names`` and ``types`` are available.
25
+ exclude : list of str, optional
26
+ Column names to ignore (e.g. the label column).
27
+
28
+ Returns
29
+ -------
30
+ str
31
+ The name of the text column.
32
+
33
+ Raises
34
+ ------
35
+ ValueError
36
+ If there is not exactly one text column.
37
+ """
38
+ from DashAI.back.types.categorical import Categorical
39
+
40
+ exclude = set(exclude or [])
41
+ try:
42
+ types = dataset.types
43
+ except AttributeError:
44
+ types = {}
45
+ text_cols = [
46
+ col
47
+ for col in dataset.column_names
48
+ if col not in exclude and not isinstance(types.get(col), Categorical)
49
+ ]
50
+ if len(text_cols) != 1:
51
+ raise ValueError(
52
+ f"Expected exactly one text column, found {text_cols} "
53
+ f"(columns={list(dataset.column_names)}, excluded={sorted(exclude)})."
54
+ )
55
+ return text_cols[0]
56
+
57
+
58
+ def extract_label_column(dataset: Any) -> str:
59
+ """Return the (single) categorical/output column name."""
60
+ from DashAI.back.types.categorical import Categorical
61
+
62
+ try:
63
+ types = dataset.types
64
+ except AttributeError:
65
+ types = {}
66
+ cat_cols = [
67
+ col for col in dataset.column_names if isinstance(types.get(col), Categorical)
68
+ ]
69
+ if len(cat_cols) != 1:
70
+ raise ValueError(
71
+ f"Expected exactly one categorical label column, found {cat_cols}."
72
+ )
73
+ return cat_cols[0]
74
+
75
+
76
+ def tokenized_dataloader(
77
+ dataset: Any,
78
+ tokenizer: Any,
79
+ text_column: str,
80
+ label_column: str,
81
+ *,
82
+ batch_size: int = 16,
83
+ max_length: int = 512,
84
+ device: str = "cpu",
85
+ shuffle: bool = True,
86
+ ) -> Any:
87
+ """Build a torch DataLoader of tokenized (input_ids, attention_mask, labels).
88
+
89
+ Parameters
90
+ ----------
91
+ dataset : DashAIDataset
92
+ Source dataset.
93
+ tokenizer : Any
94
+ HF tokenizer (or Frankenstein SPM tokenizer exposing ``__call__``
95
+ returning ``input_ids``).
96
+ text_column : str
97
+ Name of the text column.
98
+ label_column : str
99
+ Name of the integer-label column.
100
+ batch_size : int
101
+ Batch size.
102
+ max_length : int
103
+ Max token length.
104
+ device : str
105
+ Torch device (for pinning).
106
+ shuffle : bool
107
+ Whether to shuffle.
108
+
109
+ Returns
110
+ -------
111
+ torch.utils.data.DataLoader
112
+ Yields dicts with ``input_ids``, ``attention_mask``, ``labels`` tensors.
113
+ """
114
+ import torch
115
+ from torch.utils.data import DataLoader, TensorDataset
116
+
117
+ texts = list(dataset[text_column])
118
+ labels = list(dataset[label_column])
119
+
120
+ enc = tokenizer(texts, truncation=True, padding=True, max_length=max_length)
121
+ input_ids = torch.tensor(enc["input_ids"], dtype=torch.long)
122
+ attention_mask = torch.tensor(enc["attention_mask"], dtype=torch.long)
123
+ label_tensor = torch.tensor(np.asarray(labels).astype("int64"), dtype=torch.long)
124
+
125
+ ds = TensorDataset(input_ids, attention_mask, label_tensor)
126
+ return DataLoader(
127
+ ds,
128
+ batch_size=batch_size,
129
+ shuffle=shuffle,
130
+ pin_memory=str(device).startswith("cuda"),
131
+ )
132
+
133
+
134
+ def prediction_loader(
135
+ dataset: Any,
136
+ tokenizer: Any,
137
+ text_column: str,
138
+ *,
139
+ batch_size: int = 32,
140
+ max_length: int = 512,
141
+ device: str = "cpu",
142
+ ) -> Any:
143
+ """Build a DataLoader for inference (no labels required)."""
144
+ import torch
145
+ from torch.utils.data import DataLoader, TensorDataset
146
+
147
+ texts = list(dataset[text_column])
148
+ enc = tokenizer(texts, truncation=True, padding=True, max_length=max_length)
149
+ input_ids = torch.tensor(enc["input_ids"], dtype=torch.long)
150
+ attention_mask = torch.tensor(enc["attention_mask"], dtype=torch.long)
151
+ ds = TensorDataset(input_ids, attention_mask)
152
+ return DataLoader(
153
+ ds,
154
+ batch_size=batch_size,
155
+ shuffle=False,
156
+ pin_memory=str(device).startswith("cuda"),
157
+ )
158
+
159
+
160
+ # ---------------------------------------------------------------------------
161
+ # Vision: DashAI image columns -> (pixel_values, labels) tensors
162
+ # ---------------------------------------------------------------------------
163
+
164
+ def _image_column(dataset: Any) -> str:
165
+ """Return the single image column of a DashAIDataset (DashAIImage type)."""
166
+ try:
167
+ from DashAI.back.types.dashai_image import DashAIImage
168
+ except ImportError: # pragma: no cover
169
+ DashAIImage = () # type: ignore
170
+ try:
171
+ types = dataset.types
172
+ except AttributeError:
173
+ types = {}
174
+ img_cols = [
175
+ col
176
+ for col in dataset.column_names
177
+ if (DashAIImage and isinstance(types.get(col), DashAIImage))
178
+ or str(types.get(col)).lower() == "dashaiimage"
179
+ ]
180
+ if len(img_cols) != 1:
181
+ raise ValueError(
182
+ f"Expected exactly one image column, found {img_cols} "
183
+ f"(columns={list(dataset.column_names)})."
184
+ )
185
+ return img_cols[0]
186
+
187
+
188
+ def image_dataloader(
189
+ dataset: Any,
190
+ y_dataset: Any = None,
191
+ *,
192
+ image_size: int = 224,
193
+ batch_size: int = 32,
194
+ device: str = "cpu",
195
+ shuffle: bool = True,
196
+ label_column: Optional[str] = None,
197
+ ) -> Any:
198
+ """Build a DataLoader of (pixel_values, labels) or pixel_values tensors.
199
+
200
+ Images are resized to ``image_size`` x ``image_size`` and normalized with
201
+ ImageNet statistics (matching the torchvision DashAI classifiers).
202
+
203
+ Parameters
204
+ ----------
205
+ dataset : DashAIDataset
206
+ Source dataset carrying an image column.
207
+ y_dataset : DashAIDataset, optional
208
+ Label dataset (train mode). When ``None``, the loader yields images only.
209
+ image_size : int
210
+ Target square image size.
211
+ batch_size, device, shuffle
212
+ Loader options.
213
+ label_column : str, optional
214
+ Override label column name in ``y_dataset``.
215
+
216
+ Returns
217
+ -------
218
+ torch.utils.data.DataLoader
219
+ """
220
+ import torch
221
+ import torch.utils.data
222
+ from torchvision import transforms
223
+
224
+ image_col = _image_column(dataset)
225
+ label_col = label_column
226
+ if y_dataset is not None and label_col is None:
227
+ label_col = y_dataset.column_names[0]
228
+
229
+ label_to_idx: dict = {}
230
+ if y_dataset is not None and label_col is not None:
231
+ cat = (getattr(y_dataset, "types", {}) or {}).get(label_col)
232
+ if cat is not None and getattr(cat, "categories", None):
233
+ unique_labels = sorted(cat.categories)
234
+ else:
235
+ unique_labels = sorted(set(y_dataset[label_col]))
236
+ label_to_idx = {lbl: i for i, lbl in enumerate(unique_labels)}
237
+
238
+ transform = transforms.Compose(
239
+ [
240
+ transforms.Lambda(lambda img: img.convert("RGB")),
241
+ transforms.Resize((image_size, image_size)),
242
+ transforms.ToTensor(),
243
+ transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
244
+ ]
245
+ )
246
+
247
+ class _ImageDataset(torch.utils.data.Dataset):
248
+ def __init__(self):
249
+ self.x = dataset
250
+ self.y = y_dataset
251
+
252
+ def __len__(self):
253
+ return len(self.x)
254
+
255
+ def __getitem__(self, idx):
256
+ image = transform(self.x[idx][image_col].to_pil())
257
+ if self.y is None:
258
+ return image
259
+ label_str = self.y[idx][label_col]
260
+ return image, int(label_to_idx.get(label_str, -1))
261
+
262
+ def _collate_with_labels(batch):
263
+ images = torch.stack([b[0] for b in batch])
264
+ labels = torch.tensor([b[1] for b in batch], dtype=torch.long)
265
+ return images, labels
266
+
267
+ def _collate_images(batch):
268
+ return torch.stack(batch)
269
+
270
+ ds_obj = _ImageDataset()
271
+ return (
272
+ torch.utils.data.DataLoader(
273
+ ds_obj,
274
+ batch_size=batch_size,
275
+ shuffle=shuffle,
276
+ collate_fn=_collate_with_labels if y_dataset is not None else _collate_images,
277
+ pin_memory=str(device).startswith("cuda"),
278
+ ),
279
+ label_to_idx,
280
+ len(label_to_idx),
281
+ )
@@ -0,0 +1,70 @@
1
+ """IO adapter: DashAI run directory <-> Frankenstein checkpoint bundle.
2
+
3
+ Each DashAI run persists artifacts under a run-specific directory. This module
4
+ bridges the DashAI ``save(filename)``/``load(filename)`` contract to the
5
+ Frankenstein engine's :func:`save_checkpoint`/:func:`load_checkpoint`, which
6
+ write a self-contained bundle (``model.pt`` + ``config.yaml`` + ``tokenizer/``
7
+ + ``dashai_meta.json``).
8
+ """
9
+ from __future__ import annotations
10
+
11
+ import json
12
+ import os
13
+ from typing import Any, Dict, Optional, Tuple
14
+
15
+
16
+ def save_run(
17
+ filename: str,
18
+ model: Any,
19
+ loaded_config: Any,
20
+ tokenizer: Any,
21
+ *,
22
+ extra: Optional[Dict[str, Any]] = None,
23
+ ) -> str:
24
+ """Persist a Frankenstein model bundle into the DashAI run directory.
25
+
26
+ Parameters
27
+ ----------
28
+ filename : str
29
+ DashAI run directory (created if missing).
30
+ model : torch.nn.Module
31
+ The trained Frankenstein model.
32
+ loaded_config : LoadedTrainingConfig
33
+ The validated Frankenstein config (round-tripped as config.yaml).
34
+ tokenizer : Any
35
+ Tokenizer persisted alongside the weights.
36
+ extra : dict, optional
37
+ Extra metadata (e.g. ``{"num_labels": N, "label_encodings": {...}}``).
38
+
39
+ Returns
40
+ -------
41
+ str
42
+ Path to the saved ``model.pt``.
43
+ """
44
+ from dashai_frankenstein.engine import save_checkpoint
45
+
46
+ os.makedirs(filename, exist_ok=True)
47
+ path = save_checkpoint(filename, model, loaded_config, tokenizer, extra=extra)
48
+
49
+ # Drop a small DashAI-facing marker so the load path can sanity-check.
50
+ marker = {"plugin": "dashai-frankenstein", "version": "0.1.0"}
51
+ with open(os.path.join(filename, "plugin_marker.json"), "w", encoding="utf-8") as fh:
52
+ json.dump(marker, fh, indent=2)
53
+ return path
54
+
55
+
56
+ def load_run(filename: str) -> Tuple[Any, Any, Any, Dict[str, Any]]:
57
+ """Reload a bundle saved by :func:`save_run`.
58
+
59
+ Rebuilds the ``nn.Module`` from ``config.yaml`` via the Frankenstein engine
60
+ and restores weights. ``extra`` (including ``num_labels``) is applied to the
61
+ config before building so the classification head is reconstructed.
62
+
63
+ Returns
64
+ -------
65
+ tuple
66
+ ``(model, loaded_config, tokenizer, extra)``.
67
+ """
68
+ from dashai_frankenstein.engine import load_checkpoint
69
+
70
+ return load_checkpoint(filename)
@@ -0,0 +1,74 @@
1
+ """Metrics adapter: stream per-epoch metrics into DashAI's metric store.
2
+
3
+ DashAI models expose ``calculate_metrics(split, level, x_data, y_data, ...)``,
4
+ which calls ``self.predict(x_data)`` and writes scores to the database. This
5
+ module provides a tiny callback the training loop invokes at the end of each
6
+ epoch, mirroring the HuggingFace ``MetricsCallback`` pattern.
7
+ """
8
+ from __future__ import annotations
9
+
10
+ import logging
11
+ from typing import Any, Optional
12
+
13
+ from DashAI.back.core.enums.metrics import LevelEnum, SplitEnum
14
+
15
+ log = logging.getLogger(__name__)
16
+
17
+
18
+ class EpochMetricsHook:
19
+ """Invoke ``model_instance.calculate_metrics`` once per epoch.
20
+
21
+ Parameters
22
+ ----------
23
+ model_instance : BaseModel
24
+ The DashAI model component (carries ``run_id``, ``train_metrics`` ...).
25
+ x_train, y_train : DashAIDataset
26
+ Training split.
27
+ x_val, y_val : DashAIDataset, optional
28
+ Validation split. When ``None``, validation metrics are skipped.
29
+ log_every_n_epochs : int
30
+ Log frequency (epochs). ``1`` logs every epoch.
31
+ """
32
+
33
+ def __init__(
34
+ self,
35
+ model_instance: Any,
36
+ x_train: Any,
37
+ y_train: Any,
38
+ x_val: Any = None,
39
+ y_val: Any = None,
40
+ *,
41
+ log_every_n_epochs: int = 1,
42
+ ) -> None:
43
+ self.model_instance = model_instance
44
+ self.x_train = x_train
45
+ self.y_train = y_train
46
+ self.x_val = x_val
47
+ self.y_val = y_val
48
+ self.log_every_n_epochs = max(1, int(log_every_n_epochs))
49
+ self._last_epoch = -1
50
+
51
+ def __call__(self, epoch: int, **_: Any) -> None:
52
+ if epoch <= self._last_epoch:
53
+ return
54
+ self._last_epoch = epoch
55
+ if epoch % self.log_every_n_epochs != 0:
56
+ return
57
+ try:
58
+ self.model_instance.calculate_metrics(
59
+ split=SplitEnum.TRAIN,
60
+ level=LevelEnum.EPOCH,
61
+ x_data=self.x_train,
62
+ y_data=self.y_train,
63
+ log_index=epoch,
64
+ )
65
+ if self.x_val is not None and self.y_val is not None:
66
+ self.model_instance.calculate_metrics(
67
+ split=SplitEnum.VALIDATION,
68
+ level=LevelEnum.EPOCH,
69
+ x_data=self.x_val,
70
+ y_data=self.y_val,
71
+ log_index=epoch,
72
+ )
73
+ except Exception as exc: # noqa: BLE001
74
+ log.warning("calculate_metrics failed at epoch %s: %s", epoch, exc)