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.
- dashai_frankenstein-0.1.0/.gitignore +36 -0
- dashai_frankenstein-0.1.0/PKG-INFO +61 -0
- dashai_frankenstein-0.1.0/README.md +43 -0
- dashai_frankenstein-0.1.0/pyproject.toml +38 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/__init__.py +20 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/adapters/__init__.py +4 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/adapters/dataset.py +281 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/adapters/io.py +70 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/adapters/metrics.py +74 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/config.py +126 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/engine.py +218 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/models/__init__.py +10 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/models/base.py +267 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/models/decoder.py +185 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/models/mlm.py +202 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/models/vit_classifier.py +200 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/models/vit_segmenter.py +178 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/presets.py +88 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/tasks/__init__.py +4 -0
- dashai_frankenstein-0.1.0/src/dashai_frankenstein/tasks/segmentation.py +109 -0
|
@@ -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,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)
|