tensorless 0.1.0__py3-none-any.whl
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.
- tensorless/__init__.py +40 -0
- tensorless/_version.py +6 -0
- tensorless/api.py +230 -0
- tensorless/auto/__init__.py +4 -0
- tensorless/auto/config.py +113 -0
- tensorless/auto/detector.py +95 -0
- tensorless/checkpoint/__init__.py +3 -0
- tensorless/checkpoint/manager.py +66 -0
- tensorless/cli/__init__.py +3 -0
- tensorless/cli/main.py +121 -0
- tensorless/config.py +121 -0
- tensorless/data/__init__.py +11 -0
- tensorless/data/fingerprint.py +71 -0
- tensorless/data/inspector.py +161 -0
- tensorless/data/loader.py +255 -0
- tensorless/data/tabular.py +179 -0
- tensorless/devices/__init__.py +3 -0
- tensorless/devices/device.py +107 -0
- tensorless/errors.py +36 -0
- tensorless/models/__init__.py +5 -0
- tensorless/models/mlp.py +64 -0
- tensorless/models/registry.py +53 -0
- tensorless/models/transformer.py +175 -0
- tensorless/runtime.py +158 -0
- tensorless/serialization/__init__.py +3 -0
- tensorless/serialization/tl_format.py +89 -0
- tensorless/tokenization/__init__.py +3 -0
- tensorless/tokenization/char_tokenizer.py +84 -0
- tensorless/training/__init__.py +4 -0
- tensorless/training/data_prep.py +215 -0
- tensorless/training/early_stopping.py +31 -0
- tensorless/training/trainer.py +234 -0
- tensorless-0.1.0.dist-info/METADATA +111 -0
- tensorless-0.1.0.dist-info/RECORD +38 -0
- tensorless-0.1.0.dist-info/WHEEL +5 -0
- tensorless-0.1.0.dist-info/entry_points.txt +2 -0
- tensorless-0.1.0.dist-info/licenses/LICENSE +21 -0
- tensorless-0.1.0.dist-info/top_level.txt +1 -0
tensorless/runtime.py
ADDED
|
@@ -0,0 +1,158 @@
|
|
|
1
|
+
"""Runtime inference wrapper.
|
|
2
|
+
|
|
3
|
+
`tl.load("model.tl")` returns a `LoadedModel`, which knows how to rebuild
|
|
4
|
+
the exact architecture used at training time, load the weights, and
|
|
5
|
+
expose a task-appropriate prediction API:
|
|
6
|
+
|
|
7
|
+
- text-generation -> `.generate(prompt)` and `.chat()`
|
|
8
|
+
- text-classification -> `.predict(text)`
|
|
9
|
+
- classification/regression (tabular) -> `.predict(record_or_records)`
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
from typing import Any, Dict, List, Union
|
|
15
|
+
|
|
16
|
+
import torch
|
|
17
|
+
|
|
18
|
+
from .models.registry import build_model
|
|
19
|
+
from .tokenization.char_tokenizer import CharTokenizer
|
|
20
|
+
from .data.tabular import TabularPreprocessor
|
|
21
|
+
from .devices.device import get_torch_device
|
|
22
|
+
from .errors import ModelError
|
|
23
|
+
from .serialization.tl_format import load_tl
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class LoadedModel:
|
|
27
|
+
def __init__(self, payload: Dict[str, Any], device: str = None):
|
|
28
|
+
self.payload = payload
|
|
29
|
+
self.task: str = payload["task"]
|
|
30
|
+
self.model_type: str = payload["model_type"]
|
|
31
|
+
self.config: Dict[str, Any] = payload["config"]
|
|
32
|
+
self.meta: Dict[str, Any] = payload["meta"]
|
|
33
|
+
self.metrics: Dict[str, Any] = payload.get("metrics", {})
|
|
34
|
+
self.dataset_fingerprint: str = payload.get("dataset_fingerprint")
|
|
35
|
+
|
|
36
|
+
self.tokenizer = None
|
|
37
|
+
if payload.get("tokenizer_state") is not None:
|
|
38
|
+
self.tokenizer = CharTokenizer.from_state_dict(payload["tokenizer_state"])
|
|
39
|
+
|
|
40
|
+
self.preprocessor = None
|
|
41
|
+
if payload.get("preprocessor_state") is not None:
|
|
42
|
+
self.preprocessor = TabularPreprocessor.from_state_dict(payload["preprocessor_state"])
|
|
43
|
+
|
|
44
|
+
device_name = device or self.config.get("device", "cpu")
|
|
45
|
+
self.device = get_torch_device(device_name)
|
|
46
|
+
|
|
47
|
+
self.model = build_model(self.task, self.model_type, self.config, self.meta)
|
|
48
|
+
self.model.load_state_dict(payload["model_state_dict"])
|
|
49
|
+
self.model.to(self.device)
|
|
50
|
+
self.model.eval()
|
|
51
|
+
|
|
52
|
+
# ------------------------------------------------------------------
|
|
53
|
+
# Text generation
|
|
54
|
+
# ------------------------------------------------------------------
|
|
55
|
+
def generate(self, prompt: str = "", max_new_tokens: int = 200, temperature: float = 0.8, top_k: int = 40) -> str:
|
|
56
|
+
if self.task != "text-generation":
|
|
57
|
+
raise ModelError(f"generate() is only available for text-generation models, not '{self.task}'.")
|
|
58
|
+
ids = self.tokenizer.encode(prompt, add_special_tokens=True)[:-1] # drop trailing eos
|
|
59
|
+
if not ids:
|
|
60
|
+
ids = [self.tokenizer.bos_id]
|
|
61
|
+
input_ids = torch.tensor([ids], dtype=torch.long, device=self.device)
|
|
62
|
+
out = self.model.generate(
|
|
63
|
+
input_ids,
|
|
64
|
+
max_new_tokens=max_new_tokens,
|
|
65
|
+
temperature=temperature,
|
|
66
|
+
top_k=top_k,
|
|
67
|
+
eos_id=self.tokenizer.eos_id,
|
|
68
|
+
)
|
|
69
|
+
return self.tokenizer.decode(out[0].tolist())
|
|
70
|
+
|
|
71
|
+
def chat(self) -> None:
|
|
72
|
+
"""Interactive terminal chat loop for text-generation models."""
|
|
73
|
+
if self.task != "text-generation":
|
|
74
|
+
raise ModelError(f"chat() is only available for text-generation models, not '{self.task}'.")
|
|
75
|
+
print("Tensorless interactive chat. Type 'exit' or Ctrl+C to quit.")
|
|
76
|
+
while True:
|
|
77
|
+
try:
|
|
78
|
+
prompt = input("> ")
|
|
79
|
+
except (EOFError, KeyboardInterrupt):
|
|
80
|
+
print()
|
|
81
|
+
break
|
|
82
|
+
if prompt.strip().lower() in ("exit", "quit"):
|
|
83
|
+
break
|
|
84
|
+
reply = self.generate(prompt, max_new_tokens=200)
|
|
85
|
+
print(reply)
|
|
86
|
+
|
|
87
|
+
# ------------------------------------------------------------------
|
|
88
|
+
# Text classification
|
|
89
|
+
# ------------------------------------------------------------------
|
|
90
|
+
def _predict_text_classification(self, texts: List[str]) -> List[str]:
|
|
91
|
+
block_size = self.config["max_seq_len"]
|
|
92
|
+
batch_ids, batch_mask = [], []
|
|
93
|
+
for t in texts:
|
|
94
|
+
ids = self.tokenizer.encode(t, add_special_tokens=True)[:block_size]
|
|
95
|
+
mask = [1] * len(ids)
|
|
96
|
+
if len(ids) < block_size:
|
|
97
|
+
pad_len = block_size - len(ids)
|
|
98
|
+
ids = ids + [self.tokenizer.pad_id] * pad_len
|
|
99
|
+
mask = mask + [0] * pad_len
|
|
100
|
+
batch_ids.append(ids)
|
|
101
|
+
batch_mask.append(mask)
|
|
102
|
+
input_ids = torch.tensor(batch_ids, dtype=torch.long, device=self.device)
|
|
103
|
+
attn_mask = torch.tensor(batch_mask, dtype=torch.long, device=self.device)
|
|
104
|
+
with torch.no_grad():
|
|
105
|
+
logits = self.model(input_ids, attention_mask=attn_mask)
|
|
106
|
+
pred_idx = logits.argmax(dim=-1).tolist()
|
|
107
|
+
classes = self.meta["classes"]
|
|
108
|
+
return [classes[i] for i in pred_idx]
|
|
109
|
+
|
|
110
|
+
# ------------------------------------------------------------------
|
|
111
|
+
# Tabular classification / regression
|
|
112
|
+
# ------------------------------------------------------------------
|
|
113
|
+
def _predict_tabular(self, records: List[Dict[str, Any]]) -> List[Any]:
|
|
114
|
+
transformed = self.preprocessor.transform(records, with_target=False)
|
|
115
|
+
numeric = transformed["numeric"].to(self.device)
|
|
116
|
+
categorical = transformed["categorical"].to(self.device)
|
|
117
|
+
with torch.no_grad():
|
|
118
|
+
out = self.model(numeric, categorical)
|
|
119
|
+
if self.task == "classification":
|
|
120
|
+
pred_idx = out.argmax(dim=-1)
|
|
121
|
+
return self.preprocessor.inverse_target(pred_idx)
|
|
122
|
+
else:
|
|
123
|
+
return self.preprocessor.inverse_target(out)
|
|
124
|
+
|
|
125
|
+
# ------------------------------------------------------------------
|
|
126
|
+
# Unified predict()
|
|
127
|
+
# ------------------------------------------------------------------
|
|
128
|
+
def predict(self, x: Union[str, Dict[str, Any], List[Any]]) -> Any:
|
|
129
|
+
single = not isinstance(x, list)
|
|
130
|
+
items = [x] if single else x
|
|
131
|
+
|
|
132
|
+
if self.task == "text-classification":
|
|
133
|
+
preds = self._predict_text_classification(items)
|
|
134
|
+
elif self.task in ("classification", "regression"):
|
|
135
|
+
preds = self._predict_tabular(items)
|
|
136
|
+
elif self.task == "text-generation":
|
|
137
|
+
preds = [self.generate(prompt=str(i)) for i in items]
|
|
138
|
+
else:
|
|
139
|
+
raise ModelError(f"predict() not supported for task '{self.task}'.")
|
|
140
|
+
|
|
141
|
+
return preds[0] if single else preds
|
|
142
|
+
|
|
143
|
+
def info(self) -> Dict[str, Any]:
|
|
144
|
+
return {
|
|
145
|
+
"task": self.task,
|
|
146
|
+
"model_type": self.model_type,
|
|
147
|
+
"tensorless_version": self.payload.get("tensorless_version"),
|
|
148
|
+
"tl_format_version": self.payload.get("tl_format_version"),
|
|
149
|
+
"config": self.config,
|
|
150
|
+
"metrics": self.metrics,
|
|
151
|
+
"training_complete": self.payload.get("training_complete"),
|
|
152
|
+
"n_parameters": sum(p.numel() for p in self.model.parameters()),
|
|
153
|
+
}
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
def load_model(path: str, device: str = None) -> LoadedModel:
|
|
157
|
+
payload = load_tl(path)
|
|
158
|
+
return LoadedModel(payload, device=device)
|
|
@@ -0,0 +1,89 @@
|
|
|
1
|
+
"""The `.tl` file format.
|
|
2
|
+
|
|
3
|
+
A `.tl` file is a single portable file (a torch pickle archive under the
|
|
4
|
+
hood) containing everything needed for inference on a *different*
|
|
5
|
+
machine, with no access to the original dataset or training code:
|
|
6
|
+
|
|
7
|
+
{
|
|
8
|
+
"tl_format_version": int,
|
|
9
|
+
"tensorless_version": str,
|
|
10
|
+
"task": str, # e.g. "text-generation"
|
|
11
|
+
"model_type": str, # e.g. "transformer"
|
|
12
|
+
"config": {...}, # resolved training config
|
|
13
|
+
"meta": {...}, # vocab_size / n_classes / column info
|
|
14
|
+
"model_state_dict": {...},
|
|
15
|
+
"tokenizer_state": {...} | None,
|
|
16
|
+
"preprocessor_state": {...} | None,
|
|
17
|
+
"dataset_fingerprint": str,
|
|
18
|
+
"training_complete": bool,
|
|
19
|
+
"metrics": {...},
|
|
20
|
+
}
|
|
21
|
+
|
|
22
|
+
We deliberately use a single file (rather than a directory/zip of many
|
|
23
|
+
files) so users can `scp`/email/upload one `model.tl` and have it just
|
|
24
|
+
work elsewhere, per the framework's "portable single file" requirement.
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
from __future__ import annotations
|
|
28
|
+
|
|
29
|
+
import os
|
|
30
|
+
from typing import Any, Dict
|
|
31
|
+
|
|
32
|
+
import torch
|
|
33
|
+
|
|
34
|
+
from .._version import __version__, TL_FORMAT_VERSION
|
|
35
|
+
from ..errors import SerializationError
|
|
36
|
+
|
|
37
|
+
REQUIRED_KEYS = (
|
|
38
|
+
"tl_format_version",
|
|
39
|
+
"task",
|
|
40
|
+
"model_type",
|
|
41
|
+
"config",
|
|
42
|
+
"meta",
|
|
43
|
+
"model_state_dict",
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def save_tl(path: str, payload: Dict[str, Any]) -> None:
|
|
48
|
+
payload = dict(payload)
|
|
49
|
+
payload.setdefault("tl_format_version", TL_FORMAT_VERSION)
|
|
50
|
+
payload.setdefault("tensorless_version", __version__)
|
|
51
|
+
|
|
52
|
+
parent = os.path.dirname(os.path.abspath(path))
|
|
53
|
+
if parent:
|
|
54
|
+
os.makedirs(parent, exist_ok=True)
|
|
55
|
+
|
|
56
|
+
tmp_path = path + ".tmp"
|
|
57
|
+
try:
|
|
58
|
+
torch.save(payload, tmp_path)
|
|
59
|
+
os.replace(tmp_path, path)
|
|
60
|
+
except Exception as e:
|
|
61
|
+
if os.path.exists(tmp_path):
|
|
62
|
+
os.remove(tmp_path)
|
|
63
|
+
raise SerializationError(f"Failed to write .tl file to '{path}': {e}") from e
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def load_tl(path: str, map_location: str = "cpu") -> Dict[str, Any]:
|
|
67
|
+
if not os.path.isfile(path):
|
|
68
|
+
raise SerializationError(f"'{path}' does not exist.")
|
|
69
|
+
try:
|
|
70
|
+
payload = torch.load(path, map_location=map_location, weights_only=False)
|
|
71
|
+
except Exception as e:
|
|
72
|
+
raise SerializationError(f"Failed to read .tl file '{path}': {e}") from e
|
|
73
|
+
|
|
74
|
+
missing = [k for k in REQUIRED_KEYS if k not in payload]
|
|
75
|
+
if missing:
|
|
76
|
+
raise SerializationError(
|
|
77
|
+
f"'{path}' is missing required field(s) {missing}. It may be "
|
|
78
|
+
f"corrupt or not a valid Tensorless .tl file."
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
file_version = payload.get("tl_format_version")
|
|
82
|
+
if file_version > TL_FORMAT_VERSION:
|
|
83
|
+
raise SerializationError(
|
|
84
|
+
f"'{path}' was created with a newer .tl format (v{file_version}) "
|
|
85
|
+
f"than this installed version of Tensorless supports "
|
|
86
|
+
f"(v{TL_FORMAT_VERSION}). Please upgrade Tensorless."
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
return payload
|
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
"""Character-level tokenizer.
|
|
2
|
+
|
|
3
|
+
Tensorless defaults to a character-level tokenizer because it:
|
|
4
|
+
- requires no external vocabulary files or training corpus assumptions
|
|
5
|
+
- works on any UTF-8 text out of the box (any language, code, etc.)
|
|
6
|
+
- has a small, easily-portable vocabulary that fits directly inside a
|
|
7
|
+
`.tl` file
|
|
8
|
+
|
|
9
|
+
This is deliberately simple rather than a full BPE tokenizer -- Tensorless
|
|
10
|
+
optimizes for "it just works with zero setup" over maximal efficiency.
|
|
11
|
+
Advanced users can plug in their own tokenizer via `model_type=` /
|
|
12
|
+
extension points documented in the developer docs.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
from __future__ import annotations
|
|
16
|
+
|
|
17
|
+
from typing import Dict, List, Sequence
|
|
18
|
+
|
|
19
|
+
PAD = "<pad>"
|
|
20
|
+
UNK = "<unk>"
|
|
21
|
+
BOS = "<bos>"
|
|
22
|
+
EOS = "<eos>"
|
|
23
|
+
SPECIAL_TOKENS = [PAD, UNK, BOS, EOS]
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class CharTokenizer:
|
|
27
|
+
def __init__(self, vocab: List[str] = None):
|
|
28
|
+
self.vocab: List[str] = vocab if vocab is not None else list(SPECIAL_TOKENS)
|
|
29
|
+
self._stoi: Dict[str, int] = {c: i for i, c in enumerate(self.vocab)}
|
|
30
|
+
|
|
31
|
+
@property
|
|
32
|
+
def vocab_size(self) -> int:
|
|
33
|
+
return len(self.vocab)
|
|
34
|
+
|
|
35
|
+
@property
|
|
36
|
+
def pad_id(self) -> int:
|
|
37
|
+
return self._stoi[PAD]
|
|
38
|
+
|
|
39
|
+
@property
|
|
40
|
+
def unk_id(self) -> int:
|
|
41
|
+
return self._stoi[UNK]
|
|
42
|
+
|
|
43
|
+
@property
|
|
44
|
+
def bos_id(self) -> int:
|
|
45
|
+
return self._stoi[BOS]
|
|
46
|
+
|
|
47
|
+
@property
|
|
48
|
+
def eos_id(self) -> int:
|
|
49
|
+
return self._stoi[EOS]
|
|
50
|
+
|
|
51
|
+
@classmethod
|
|
52
|
+
def build(cls, texts: Sequence[str]) -> "CharTokenizer":
|
|
53
|
+
chars = set()
|
|
54
|
+
for t in texts:
|
|
55
|
+
chars.update(t)
|
|
56
|
+
vocab = list(SPECIAL_TOKENS) + sorted(chars)
|
|
57
|
+
return cls(vocab)
|
|
58
|
+
|
|
59
|
+
def encode(self, text: str, add_special_tokens: bool = True) -> List[int]:
|
|
60
|
+
ids = [self._stoi.get(c, self.unk_id) for c in text]
|
|
61
|
+
if add_special_tokens:
|
|
62
|
+
ids = [self.bos_id] + ids + [self.eos_id]
|
|
63
|
+
return ids
|
|
64
|
+
|
|
65
|
+
def decode(self, ids: Sequence[int], skip_special_tokens: bool = True) -> str:
|
|
66
|
+
chars = []
|
|
67
|
+
special_ids = {self.pad_id, self.bos_id, self.eos_id} if skip_special_tokens else set()
|
|
68
|
+
for i in ids:
|
|
69
|
+
if i in special_ids:
|
|
70
|
+
continue
|
|
71
|
+
if 0 <= i < len(self.vocab):
|
|
72
|
+
tok = self.vocab[i]
|
|
73
|
+
if tok == UNK:
|
|
74
|
+
chars.append("\ufffd")
|
|
75
|
+
elif tok not in SPECIAL_TOKENS:
|
|
76
|
+
chars.append(tok)
|
|
77
|
+
return "".join(chars)
|
|
78
|
+
|
|
79
|
+
def state_dict(self) -> Dict:
|
|
80
|
+
return {"vocab": self.vocab}
|
|
81
|
+
|
|
82
|
+
@classmethod
|
|
83
|
+
def from_state_dict(cls, state: Dict) -> "CharTokenizer":
|
|
84
|
+
return cls(vocab=list(state["vocab"]))
|
|
@@ -0,0 +1,215 @@
|
|
|
1
|
+
"""Turns a loaded `Dataset` + resolved config into PyTorch-ready tensors,
|
|
2
|
+
train/val splits, and the tokenizer/preprocessor that produced them.
|
|
3
|
+
|
|
4
|
+
Kept separate from `trainer.py` so the "how do I turn data into tensors"
|
|
5
|
+
logic can be tested and extended independently of the training loop.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import random
|
|
11
|
+
from dataclasses import dataclass
|
|
12
|
+
from typing import Any, Dict, List, Optional, Tuple
|
|
13
|
+
|
|
14
|
+
import torch
|
|
15
|
+
from torch.utils.data import DataLoader, Dataset as TorchDataset
|
|
16
|
+
|
|
17
|
+
from ..data.loader import Dataset
|
|
18
|
+
from ..data.tabular import TabularPreprocessor
|
|
19
|
+
from ..tokenization.char_tokenizer import CharTokenizer
|
|
20
|
+
from ..auto.detector import target_column
|
|
21
|
+
from ..errors import DataError
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
@dataclass
|
|
25
|
+
class PreparedData:
|
|
26
|
+
train_loader: DataLoader
|
|
27
|
+
val_loader: Optional[DataLoader]
|
|
28
|
+
meta: Dict[str, Any]
|
|
29
|
+
tokenizer: Optional[CharTokenizer] = None
|
|
30
|
+
preprocessor: Optional[TabularPreprocessor] = None
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class _LMChunkDataset(TorchDataset):
|
|
34
|
+
def __init__(self, token_ids: List[int], block_size: int):
|
|
35
|
+
self.ids = token_ids
|
|
36
|
+
self.block_size = block_size
|
|
37
|
+
|
|
38
|
+
def __len__(self) -> int:
|
|
39
|
+
return max(0, len(self.ids) - self.block_size)
|
|
40
|
+
|
|
41
|
+
def __getitem__(self, idx: int):
|
|
42
|
+
chunk = self.ids[idx: idx + self.block_size + 1]
|
|
43
|
+
x = torch.tensor(chunk[:-1], dtype=torch.long)
|
|
44
|
+
y = torch.tensor(chunk[1:], dtype=torch.long)
|
|
45
|
+
return x, y
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
class _ClsTextDataset(TorchDataset):
|
|
49
|
+
def __init__(self, input_ids: List[List[int]], attn_masks: List[List[int]], labels: List[int]):
|
|
50
|
+
self.input_ids = input_ids
|
|
51
|
+
self.attn_masks = attn_masks
|
|
52
|
+
self.labels = labels
|
|
53
|
+
|
|
54
|
+
def __len__(self) -> int:
|
|
55
|
+
return len(self.labels)
|
|
56
|
+
|
|
57
|
+
def __getitem__(self, idx: int):
|
|
58
|
+
return (
|
|
59
|
+
torch.tensor(self.input_ids[idx], dtype=torch.long),
|
|
60
|
+
torch.tensor(self.attn_masks[idx], dtype=torch.long),
|
|
61
|
+
torch.tensor(self.labels[idx], dtype=torch.long),
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
class _TabularDataset(TorchDataset):
|
|
66
|
+
def __init__(self, numeric: torch.Tensor, categorical: torch.Tensor, target: torch.Tensor):
|
|
67
|
+
self.numeric = numeric
|
|
68
|
+
self.categorical = categorical
|
|
69
|
+
self.target = target
|
|
70
|
+
|
|
71
|
+
def __len__(self) -> int:
|
|
72
|
+
return self.numeric.shape[0]
|
|
73
|
+
|
|
74
|
+
def __getitem__(self, idx: int):
|
|
75
|
+
return self.numeric[idx], self.categorical[idx], self.target[idx]
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def _split_indices(n: int, val_split: float, seed: int) -> Tuple[List[int], List[int]]:
|
|
79
|
+
idx = list(range(n))
|
|
80
|
+
random.Random(seed).shuffle(idx)
|
|
81
|
+
n_val = int(n * val_split)
|
|
82
|
+
val_idx = idx[:n_val]
|
|
83
|
+
train_idx = idx[n_val:]
|
|
84
|
+
if not train_idx:
|
|
85
|
+
train_idx = idx
|
|
86
|
+
val_idx = []
|
|
87
|
+
return train_idx, val_idx
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
def prepare_text_generation(
|
|
91
|
+
ds: Dataset, cfg: Dict[str, Any], tokenizer: Optional[CharTokenizer] = None
|
|
92
|
+
) -> PreparedData:
|
|
93
|
+
tokenizer = tokenizer or CharTokenizer.build(ds.texts)
|
|
94
|
+
all_ids: List[int] = []
|
|
95
|
+
for t in ds.texts:
|
|
96
|
+
all_ids.extend(tokenizer.encode(t, add_special_tokens=True))
|
|
97
|
+
|
|
98
|
+
block_size = cfg["max_seq_len"]
|
|
99
|
+
if len(all_ids) <= block_size:
|
|
100
|
+
# Pad tiny corpora so we have at least one training example.
|
|
101
|
+
all_ids = all_ids + [tokenizer.pad_id] * (block_size + 1 - len(all_ids))
|
|
102
|
+
|
|
103
|
+
n_val_tokens = int(len(all_ids) * cfg["val_split"])
|
|
104
|
+
if n_val_tokens > block_size + 1:
|
|
105
|
+
split_point = len(all_ids) - n_val_tokens
|
|
106
|
+
train_ids, val_ids = all_ids[:split_point], all_ids[split_point:]
|
|
107
|
+
else:
|
|
108
|
+
train_ids, val_ids = all_ids, None
|
|
109
|
+
|
|
110
|
+
train_ds = _LMChunkDataset(train_ids, block_size)
|
|
111
|
+
train_loader = DataLoader(
|
|
112
|
+
train_ds, batch_size=cfg["batch_size"], shuffle=True, drop_last=False
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
val_loader = None
|
|
116
|
+
if val_ids is not None:
|
|
117
|
+
val_ds = _LMChunkDataset(val_ids, block_size)
|
|
118
|
+
if len(val_ds) > 0:
|
|
119
|
+
val_loader = DataLoader(val_ds, batch_size=cfg["batch_size"], shuffle=False)
|
|
120
|
+
|
|
121
|
+
meta = {
|
|
122
|
+
"vocab_size": tokenizer.vocab_size,
|
|
123
|
+
"pad_id": tokenizer.pad_id,
|
|
124
|
+
"n_classes": 0,
|
|
125
|
+
}
|
|
126
|
+
return PreparedData(train_loader=train_loader, val_loader=val_loader, meta=meta, tokenizer=tokenizer)
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
def prepare_text_classification(
|
|
130
|
+
ds: Dataset,
|
|
131
|
+
cfg: Dict[str, Any],
|
|
132
|
+
tokenizer: Optional[CharTokenizer] = None,
|
|
133
|
+
classes: Optional[List[str]] = None,
|
|
134
|
+
) -> PreparedData:
|
|
135
|
+
tokenizer = tokenizer or CharTokenizer.build(ds.texts)
|
|
136
|
+
classes = classes or sorted(set(ds.labels))
|
|
137
|
+
label2id = {c: i for i, c in enumerate(classes)}
|
|
138
|
+
|
|
139
|
+
block_size = cfg["max_seq_len"]
|
|
140
|
+
input_ids, attn_masks, labels = [], [], []
|
|
141
|
+
for text, label in zip(ds.texts, ds.labels):
|
|
142
|
+
ids = tokenizer.encode(text, add_special_tokens=True)[:block_size]
|
|
143
|
+
mask = [1] * len(ids)
|
|
144
|
+
if len(ids) < block_size:
|
|
145
|
+
pad_len = block_size - len(ids)
|
|
146
|
+
ids = ids + [tokenizer.pad_id] * pad_len
|
|
147
|
+
mask = mask + [0] * pad_len
|
|
148
|
+
input_ids.append(ids)
|
|
149
|
+
attn_masks.append(mask)
|
|
150
|
+
labels.append(label2id.get(label, 0))
|
|
151
|
+
|
|
152
|
+
n = len(labels)
|
|
153
|
+
train_idx, val_idx = _split_indices(n, cfg["val_split"], cfg["seed"])
|
|
154
|
+
|
|
155
|
+
def subset(indices):
|
|
156
|
+
return _ClsTextDataset(
|
|
157
|
+
[input_ids[i] for i in indices],
|
|
158
|
+
[attn_masks[i] for i in indices],
|
|
159
|
+
[labels[i] for i in indices],
|
|
160
|
+
)
|
|
161
|
+
|
|
162
|
+
train_loader = DataLoader(subset(train_idx), batch_size=cfg["batch_size"], shuffle=True)
|
|
163
|
+
val_loader = (
|
|
164
|
+
DataLoader(subset(val_idx), batch_size=cfg["batch_size"], shuffle=False) if val_idx else None
|
|
165
|
+
)
|
|
166
|
+
|
|
167
|
+
meta = {
|
|
168
|
+
"vocab_size": tokenizer.vocab_size,
|
|
169
|
+
"pad_id": tokenizer.pad_id,
|
|
170
|
+
"n_classes": len(classes),
|
|
171
|
+
"classes": classes,
|
|
172
|
+
}
|
|
173
|
+
return PreparedData(train_loader=train_loader, val_loader=val_loader, meta=meta, tokenizer=tokenizer)
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
def prepare_tabular(
|
|
177
|
+
ds: Dataset,
|
|
178
|
+
cfg: Dict[str, Any],
|
|
179
|
+
task: str,
|
|
180
|
+
preprocessor: Optional[TabularPreprocessor] = None,
|
|
181
|
+
) -> PreparedData:
|
|
182
|
+
target_col = target_column(ds)
|
|
183
|
+
if target_col is None:
|
|
184
|
+
raise DataError("Could not determine a target column for tabular training.")
|
|
185
|
+
|
|
186
|
+
prep = preprocessor or TabularPreprocessor().fit(ds.records, ds.columns, target_col, task)
|
|
187
|
+
transformed = prep.transform(ds.records, with_target=True)
|
|
188
|
+
|
|
189
|
+
n = transformed["numeric"].shape[0]
|
|
190
|
+
train_idx, val_idx = _split_indices(n, cfg["val_split"], cfg["seed"])
|
|
191
|
+
train_idx_t = torch.tensor(train_idx, dtype=torch.long)
|
|
192
|
+
|
|
193
|
+
train_ds = _TabularDataset(
|
|
194
|
+
transformed["numeric"][train_idx_t],
|
|
195
|
+
transformed["categorical"][train_idx_t],
|
|
196
|
+
transformed["target"][train_idx_t],
|
|
197
|
+
)
|
|
198
|
+
train_loader = DataLoader(train_ds, batch_size=cfg["batch_size"], shuffle=True)
|
|
199
|
+
|
|
200
|
+
val_loader = None
|
|
201
|
+
if val_idx:
|
|
202
|
+
val_idx_t = torch.tensor(val_idx, dtype=torch.long)
|
|
203
|
+
val_ds = _TabularDataset(
|
|
204
|
+
transformed["numeric"][val_idx_t],
|
|
205
|
+
transformed["categorical"][val_idx_t],
|
|
206
|
+
transformed["target"][val_idx_t],
|
|
207
|
+
)
|
|
208
|
+
val_loader = DataLoader(val_ds, batch_size=cfg["batch_size"], shuffle=False)
|
|
209
|
+
|
|
210
|
+
meta = {
|
|
211
|
+
"n_numeric": len(prep.numeric_columns),
|
|
212
|
+
"categorical_vocab_sizes": prep.categorical_vocab_sizes(),
|
|
213
|
+
"n_classes": len(prep.classes) if task == "classification" else 0,
|
|
214
|
+
}
|
|
215
|
+
return PreparedData(train_loader=train_loader, val_loader=val_loader, meta=meta, preprocessor=prep)
|
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
"""Early stopping.
|
|
2
|
+
|
|
3
|
+
Tracks a validation metric (lower is better, e.g. loss) and reports when
|
|
4
|
+
training should stop because it hasn't improved by at least `min_delta`
|
|
5
|
+
for `patience` consecutive checks.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class EarlyStopping:
|
|
12
|
+
def __init__(self, patience: int = 5, min_delta: float = 1e-4):
|
|
13
|
+
self.patience = patience
|
|
14
|
+
self.min_delta = min_delta
|
|
15
|
+
self.best: float = float("inf")
|
|
16
|
+
self.best_state: dict = None
|
|
17
|
+
self.num_bad_checks = 0
|
|
18
|
+
self.should_stop = False
|
|
19
|
+
|
|
20
|
+
def step(self, value: float, state: dict = None) -> bool:
|
|
21
|
+
"""Returns True if `value` is a new best."""
|
|
22
|
+
if value < self.best - self.min_delta:
|
|
23
|
+
self.best = value
|
|
24
|
+
self.best_state = state
|
|
25
|
+
self.num_bad_checks = 0
|
|
26
|
+
return True
|
|
27
|
+
else:
|
|
28
|
+
self.num_bad_checks += 1
|
|
29
|
+
if self.num_bad_checks >= self.patience:
|
|
30
|
+
self.should_stop = True
|
|
31
|
+
return False
|