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
|
@@ -0,0 +1,179 @@
|
|
|
1
|
+
"""Tabular preprocessing.
|
|
2
|
+
|
|
3
|
+
Fits simple, fully-reversible preprocessing on tabular records:
|
|
4
|
+
- numeric columns -> standardized (mean/std), missing values imputed
|
|
5
|
+
with the training-set mean
|
|
6
|
+
- categorical columns -> integer-indexed vocabulary (+ <unk>/<missing>),
|
|
7
|
+
fed into per-column embeddings by the MLP model
|
|
8
|
+
|
|
9
|
+
The fitted state is small and JSON-serializable, so it can be embedded
|
|
10
|
+
directly inside a `.tl` file and reproduced exactly at inference time.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
from dataclasses import dataclass, field
|
|
16
|
+
from typing import Any, Dict, List, Optional, Tuple
|
|
17
|
+
|
|
18
|
+
import torch
|
|
19
|
+
|
|
20
|
+
_MISSING = "<missing>"
|
|
21
|
+
_UNK = "<unk>"
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def _try_float(v: Any) -> Optional[float]:
|
|
25
|
+
if v is None:
|
|
26
|
+
return None
|
|
27
|
+
if isinstance(v, str) and v.strip() == "":
|
|
28
|
+
return None
|
|
29
|
+
try:
|
|
30
|
+
return float(v)
|
|
31
|
+
except (TypeError, ValueError):
|
|
32
|
+
return None
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
@dataclass
|
|
36
|
+
class ColumnStats:
|
|
37
|
+
kind: str # "numeric" | "categorical"
|
|
38
|
+
mean: float = 0.0
|
|
39
|
+
std: float = 1.0
|
|
40
|
+
vocab: List[str] = field(default_factory=list)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
@dataclass
|
|
44
|
+
class TabularPreprocessor:
|
|
45
|
+
feature_columns: List[str] = field(default_factory=list)
|
|
46
|
+
target_column: Optional[str] = None
|
|
47
|
+
task: str = "classification"
|
|
48
|
+
column_stats: Dict[str, ColumnStats] = field(default_factory=dict)
|
|
49
|
+
classes: List[str] = field(default_factory=list) # for classification targets
|
|
50
|
+
target_mean: float = 0.0
|
|
51
|
+
target_std: float = 1.0
|
|
52
|
+
|
|
53
|
+
@property
|
|
54
|
+
def numeric_columns(self) -> List[str]:
|
|
55
|
+
return [c for c in self.feature_columns if self.column_stats[c].kind == "numeric"]
|
|
56
|
+
|
|
57
|
+
@property
|
|
58
|
+
def categorical_columns(self) -> List[str]:
|
|
59
|
+
return [c for c in self.feature_columns if self.column_stats[c].kind == "categorical"]
|
|
60
|
+
|
|
61
|
+
def fit(
|
|
62
|
+
self,
|
|
63
|
+
records: List[Dict[str, Any]],
|
|
64
|
+
columns: List[str],
|
|
65
|
+
target_column: str,
|
|
66
|
+
task: str,
|
|
67
|
+
) -> "TabularPreprocessor":
|
|
68
|
+
self.target_column = target_column
|
|
69
|
+
self.task = task
|
|
70
|
+
self.feature_columns = [c for c in columns if c != target_column]
|
|
71
|
+
|
|
72
|
+
for col in self.feature_columns:
|
|
73
|
+
values = [r.get(col) for r in records]
|
|
74
|
+
numeric_vals = [_try_float(v) for v in values]
|
|
75
|
+
n_present = sum(1 for v in values if v not in (None, ""))
|
|
76
|
+
n_numeric = sum(1 for v in numeric_vals if v is not None)
|
|
77
|
+
if n_present > 0 and n_numeric / n_present > 0.95:
|
|
78
|
+
nums = [v for v in numeric_vals if v is not None]
|
|
79
|
+
mean = sum(nums) / len(nums) if nums else 0.0
|
|
80
|
+
var = sum((x - mean) ** 2 for x in nums) / len(nums) if nums else 1.0
|
|
81
|
+
std = max(var ** 0.5, 1e-6)
|
|
82
|
+
self.column_stats[col] = ColumnStats(kind="numeric", mean=mean, std=std)
|
|
83
|
+
else:
|
|
84
|
+
cats = sorted({str(v) for v in values if v not in (None, "")})
|
|
85
|
+
vocab = [_MISSING, _UNK] + cats
|
|
86
|
+
self.column_stats[col] = ColumnStats(kind="categorical", vocab=vocab)
|
|
87
|
+
|
|
88
|
+
target_vals = [r.get(target_column) for r in records]
|
|
89
|
+
if task == "regression":
|
|
90
|
+
nums = [v for v in (_try_float(v) for v in target_vals) if v is not None]
|
|
91
|
+
self.target_mean = sum(nums) / len(nums) if nums else 0.0
|
|
92
|
+
var = sum((x - self.target_mean) ** 2 for x in nums) / len(nums) if nums else 1.0
|
|
93
|
+
self.target_std = max(var ** 0.5, 1e-6)
|
|
94
|
+
else:
|
|
95
|
+
self.classes = sorted({str(v) for v in target_vals if v not in (None, "")})
|
|
96
|
+
|
|
97
|
+
return self
|
|
98
|
+
|
|
99
|
+
def transform(
|
|
100
|
+
self, records: List[Dict[str, Any]], with_target: bool = True
|
|
101
|
+
) -> Dict[str, torch.Tensor]:
|
|
102
|
+
n = len(records)
|
|
103
|
+
num_cols = self.numeric_columns
|
|
104
|
+
cat_cols = self.categorical_columns
|
|
105
|
+
|
|
106
|
+
numeric = torch.zeros(n, max(len(num_cols), 1), dtype=torch.float32)
|
|
107
|
+
categorical = torch.zeros(n, max(len(cat_cols), 1), dtype=torch.long)
|
|
108
|
+
|
|
109
|
+
for i, r in enumerate(records):
|
|
110
|
+
for j, col in enumerate(num_cols):
|
|
111
|
+
stats = self.column_stats[col]
|
|
112
|
+
v = _try_float(r.get(col))
|
|
113
|
+
if v is None:
|
|
114
|
+
v = stats.mean
|
|
115
|
+
numeric[i, j] = (v - stats.mean) / stats.std
|
|
116
|
+
for j, col in enumerate(cat_cols):
|
|
117
|
+
stats = self.column_stats[col]
|
|
118
|
+
raw = r.get(col)
|
|
119
|
+
key = _MISSING if raw in (None, "") else str(raw)
|
|
120
|
+
idx = stats.vocab.index(key) if key in stats.vocab else stats.vocab.index(_UNK)
|
|
121
|
+
categorical[i, j] = idx
|
|
122
|
+
|
|
123
|
+
out = {"numeric": numeric, "categorical": categorical}
|
|
124
|
+
|
|
125
|
+
if with_target and self.target_column is not None:
|
|
126
|
+
if self.task == "regression":
|
|
127
|
+
target = torch.zeros(n, dtype=torch.float32)
|
|
128
|
+
for i, r in enumerate(records):
|
|
129
|
+
v = _try_float(r.get(self.target_column))
|
|
130
|
+
v = self.target_mean if v is None else v
|
|
131
|
+
target[i] = (v - self.target_mean) / self.target_std
|
|
132
|
+
out["target"] = target
|
|
133
|
+
else:
|
|
134
|
+
target = torch.zeros(n, dtype=torch.long)
|
|
135
|
+
for i, r in enumerate(records):
|
|
136
|
+
raw = str(r.get(self.target_column))
|
|
137
|
+
idx = self.classes.index(raw) if raw in self.classes else 0
|
|
138
|
+
target[i] = idx
|
|
139
|
+
out["target"] = target
|
|
140
|
+
|
|
141
|
+
return out
|
|
142
|
+
|
|
143
|
+
def inverse_target(self, values: torch.Tensor) -> List[Any]:
|
|
144
|
+
if self.task == "regression":
|
|
145
|
+
return [(v.item() * self.target_std + self.target_mean) for v in values]
|
|
146
|
+
return [self.classes[int(v.item())] for v in values]
|
|
147
|
+
|
|
148
|
+
def categorical_vocab_sizes(self) -> List[int]:
|
|
149
|
+
return [len(self.column_stats[c].vocab) for c in self.categorical_columns]
|
|
150
|
+
|
|
151
|
+
def state_dict(self) -> Dict:
|
|
152
|
+
return {
|
|
153
|
+
"feature_columns": self.feature_columns,
|
|
154
|
+
"target_column": self.target_column,
|
|
155
|
+
"task": self.task,
|
|
156
|
+
"column_stats": {
|
|
157
|
+
c: {"kind": s.kind, "mean": s.mean, "std": s.std, "vocab": s.vocab}
|
|
158
|
+
for c, s in self.column_stats.items()
|
|
159
|
+
},
|
|
160
|
+
"classes": self.classes,
|
|
161
|
+
"target_mean": self.target_mean,
|
|
162
|
+
"target_std": self.target_std,
|
|
163
|
+
}
|
|
164
|
+
|
|
165
|
+
@classmethod
|
|
166
|
+
def from_state_dict(cls, state: Dict) -> "TabularPreprocessor":
|
|
167
|
+
prep = cls(
|
|
168
|
+
feature_columns=list(state["feature_columns"]),
|
|
169
|
+
target_column=state["target_column"],
|
|
170
|
+
task=state["task"],
|
|
171
|
+
classes=list(state.get("classes", [])),
|
|
172
|
+
target_mean=state.get("target_mean", 0.0),
|
|
173
|
+
target_std=state.get("target_std", 1.0),
|
|
174
|
+
)
|
|
175
|
+
prep.column_stats = {
|
|
176
|
+
c: ColumnStats(kind=v["kind"], mean=v.get("mean", 0.0), std=v.get("std", 1.0), vocab=list(v.get("vocab", [])))
|
|
177
|
+
for c, v in state["column_stats"].items()
|
|
178
|
+
}
|
|
179
|
+
return prep
|
|
@@ -0,0 +1,107 @@
|
|
|
1
|
+
"""Hardware auto-detection.
|
|
2
|
+
|
|
3
|
+
Preference order is TPU -> GPU -> CPU, but selection is "intelligent"
|
|
4
|
+
rather than blind: we verify each backend is actually usable (not just
|
|
5
|
+
importable) before choosing it, and we fall back gracefully -- including
|
|
6
|
+
at *runtime*, if a chosen device turns out to error out mid-training, the
|
|
7
|
+
trainer (see `training/trainer.py`) will catch that and fall back too.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
from typing import Optional, Tuple
|
|
13
|
+
|
|
14
|
+
import torch
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _tpu_available() -> bool:
|
|
18
|
+
try:
|
|
19
|
+
import torch_xla.core.xla_model as xm # noqa: F401
|
|
20
|
+
|
|
21
|
+
return True
|
|
22
|
+
except Exception:
|
|
23
|
+
return False
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def _cuda_available() -> bool:
|
|
27
|
+
try:
|
|
28
|
+
return torch.cuda.is_available() and torch.cuda.device_count() > 0
|
|
29
|
+
except Exception:
|
|
30
|
+
return False
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def _mps_available() -> bool:
|
|
34
|
+
try:
|
|
35
|
+
return torch.backends.mps.is_available()
|
|
36
|
+
except Exception:
|
|
37
|
+
return False
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def _cuda_supports_bf16() -> bool:
|
|
41
|
+
try:
|
|
42
|
+
return torch.cuda.is_bf16_supported()
|
|
43
|
+
except Exception:
|
|
44
|
+
return False
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def auto_select_device(user_device: Optional[str], user_precision: Optional[str]) -> Tuple[str, str]:
|
|
48
|
+
"""Resolve the device and precision to use.
|
|
49
|
+
|
|
50
|
+
`user_device` / `user_precision` are honored if given (with a
|
|
51
|
+
graceful downgrade if the requested device isn't actually available).
|
|
52
|
+
Otherwise we pick automatically: tpu > cuda > mps > cpu.
|
|
53
|
+
"""
|
|
54
|
+
if user_device is not None:
|
|
55
|
+
device = user_device
|
|
56
|
+
if device == "tpu" and not _tpu_available():
|
|
57
|
+
device = "cuda" if _cuda_available() else "cpu"
|
|
58
|
+
elif device == "cuda" and not _cuda_available():
|
|
59
|
+
device = "cpu"
|
|
60
|
+
elif device == "mps" and not _mps_available():
|
|
61
|
+
device = "cpu"
|
|
62
|
+
else:
|
|
63
|
+
if _tpu_available():
|
|
64
|
+
device = "tpu"
|
|
65
|
+
elif _cuda_available():
|
|
66
|
+
device = "cuda"
|
|
67
|
+
elif _mps_available():
|
|
68
|
+
device = "mps"
|
|
69
|
+
else:
|
|
70
|
+
device = "cpu"
|
|
71
|
+
|
|
72
|
+
if user_precision is not None:
|
|
73
|
+
precision = user_precision
|
|
74
|
+
else:
|
|
75
|
+
if device == "cuda" and _cuda_supports_bf16():
|
|
76
|
+
precision = "bf16"
|
|
77
|
+
elif device == "cuda":
|
|
78
|
+
precision = "fp16"
|
|
79
|
+
elif device == "tpu":
|
|
80
|
+
precision = "bf16"
|
|
81
|
+
else:
|
|
82
|
+
# CPU and MPS: stick to fp32 for correctness/stability by default.
|
|
83
|
+
precision = "fp32"
|
|
84
|
+
|
|
85
|
+
return device, precision
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def get_torch_device(device: str) -> torch.device:
|
|
89
|
+
"""Convert our string device name into a torch.device, with a
|
|
90
|
+
runtime fallback to CPU if the requested backend is unavailable.
|
|
91
|
+
"""
|
|
92
|
+
try:
|
|
93
|
+
if device == "tpu":
|
|
94
|
+
import torch_xla.core.xla_model as xm
|
|
95
|
+
|
|
96
|
+
return xm.xla_device()
|
|
97
|
+
if device == "cuda":
|
|
98
|
+
if not _cuda_available():
|
|
99
|
+
return torch.device("cpu")
|
|
100
|
+
return torch.device("cuda")
|
|
101
|
+
if device == "mps":
|
|
102
|
+
if not _mps_available():
|
|
103
|
+
return torch.device("cpu")
|
|
104
|
+
return torch.device("mps")
|
|
105
|
+
return torch.device("cpu")
|
|
106
|
+
except Exception:
|
|
107
|
+
return torch.device("cpu")
|
tensorless/errors.py
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
1
|
+
"""Tensorless error hierarchy.
|
|
2
|
+
|
|
3
|
+
All errors raised by Tensorless inherit from :class:`TensorlessError`, so
|
|
4
|
+
callers can do::
|
|
5
|
+
|
|
6
|
+
try:
|
|
7
|
+
tl.train("./data")
|
|
8
|
+
except tl.TensorlessError as e:
|
|
9
|
+
...
|
|
10
|
+
|
|
11
|
+
instead of catching bare Exception.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class TensorlessError(Exception):
|
|
16
|
+
"""Base class for all Tensorless errors."""
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class DataError(TensorlessError):
|
|
20
|
+
"""Raised when a dataset cannot be read, is empty, or is malformed."""
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class ConfigError(TensorlessError):
|
|
24
|
+
"""Raised when user-supplied configuration is invalid or contradictory."""
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
class ModelError(TensorlessError):
|
|
28
|
+
"""Raised for unsupported model/task combinations or model build failures."""
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class CheckpointError(TensorlessError):
|
|
32
|
+
"""Raised when a checkpoint is missing, corrupt, or incompatible."""
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class SerializationError(TensorlessError):
|
|
36
|
+
"""Raised when a `.tl` file cannot be written or read."""
|
tensorless/models/mlp.py
ADDED
|
@@ -0,0 +1,64 @@
|
|
|
1
|
+
"""MLP model for tabular data (classification / regression).
|
|
2
|
+
|
|
3
|
+
Numeric columns feed directly into the network; each categorical column
|
|
4
|
+
gets its own small embedding table, and all features are concatenated
|
|
5
|
+
before the hidden layers.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
from typing import List
|
|
11
|
+
|
|
12
|
+
import torch
|
|
13
|
+
import torch.nn as nn
|
|
14
|
+
import torch.nn.functional as F
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _embedding_dim(vocab_size: int) -> int:
|
|
18
|
+
# Common rule of thumb, capped to keep tiny-data models tiny.
|
|
19
|
+
return max(2, min(32, round(1.6 * (vocab_size ** 0.56))))
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class TabularMLP(nn.Module):
|
|
23
|
+
def __init__(
|
|
24
|
+
self,
|
|
25
|
+
n_numeric: int,
|
|
26
|
+
categorical_vocab_sizes: List[int],
|
|
27
|
+
d_model: int,
|
|
28
|
+
layers: int,
|
|
29
|
+
dropout: float,
|
|
30
|
+
task: str,
|
|
31
|
+
n_classes: int = 0,
|
|
32
|
+
):
|
|
33
|
+
super().__init__()
|
|
34
|
+
self.task = task
|
|
35
|
+
self.n_numeric = n_numeric
|
|
36
|
+
self.categorical_vocab_sizes = categorical_vocab_sizes
|
|
37
|
+
|
|
38
|
+
self.embeddings = nn.ModuleList(
|
|
39
|
+
[nn.Embedding(v, _embedding_dim(v)) for v in categorical_vocab_sizes]
|
|
40
|
+
)
|
|
41
|
+
cat_dim = sum(_embedding_dim(v) for v in categorical_vocab_sizes)
|
|
42
|
+
in_dim = n_numeric + cat_dim
|
|
43
|
+
|
|
44
|
+
dims = [in_dim] + [d_model] * layers
|
|
45
|
+
modules = []
|
|
46
|
+
for i in range(len(dims) - 1):
|
|
47
|
+
modules += [nn.Linear(dims[i], dims[i + 1]), nn.ReLU(), nn.Dropout(dropout)]
|
|
48
|
+
self.backbone = nn.Sequential(*modules)
|
|
49
|
+
|
|
50
|
+
out_dim = n_classes if task == "classification" else 1
|
|
51
|
+
self.head = nn.Linear(dims[-1], out_dim)
|
|
52
|
+
|
|
53
|
+
def forward(self, numeric: torch.Tensor, categorical: torch.Tensor) -> torch.Tensor:
|
|
54
|
+
parts = []
|
|
55
|
+
if self.n_numeric > 0:
|
|
56
|
+
parts.append(numeric)
|
|
57
|
+
for i, emb in enumerate(self.embeddings):
|
|
58
|
+
parts.append(emb(categorical[:, i]))
|
|
59
|
+
x = torch.cat(parts, dim=1) if parts else numeric
|
|
60
|
+
x = self.backbone(x)
|
|
61
|
+
out = self.head(x)
|
|
62
|
+
if self.task == "regression":
|
|
63
|
+
return out.squeeze(-1)
|
|
64
|
+
return out
|
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
"""Model registry.
|
|
2
|
+
|
|
3
|
+
Central place that knows how to build a model given a task + resolved
|
|
4
|
+
config + the dataset-derived metadata (vocab size, n_classes, etc). Kept
|
|
5
|
+
separate from `auto/config.py` so new model types/backends can be added
|
|
6
|
+
here without touching auto-configuration logic, per the "extensible
|
|
7
|
+
architecture" requirement.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
from typing import Any, Dict
|
|
13
|
+
|
|
14
|
+
import torch.nn as nn
|
|
15
|
+
|
|
16
|
+
from .transformer import TinyTransformer
|
|
17
|
+
from .mlp import TabularMLP
|
|
18
|
+
from ..errors import ModelError
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def build_model(task: str, model_type: str, cfg: Dict[str, Any], meta: Dict[str, Any]) -> nn.Module:
|
|
22
|
+
"""Build a fresh, randomly-initialized model.
|
|
23
|
+
|
|
24
|
+
`cfg` is the resolved training config (dict). `meta` carries
|
|
25
|
+
task-specific sizing info produced during data prep, e.g.:
|
|
26
|
+
- text tasks: {"vocab_size": int, "pad_id": int, "n_classes": int}
|
|
27
|
+
- tabular tasks: {"n_numeric": int, "categorical_vocab_sizes": [...], "n_classes": int}
|
|
28
|
+
"""
|
|
29
|
+
if model_type == "transformer":
|
|
30
|
+
return TinyTransformer(
|
|
31
|
+
vocab_size=meta["vocab_size"],
|
|
32
|
+
d_model=cfg["d_model"],
|
|
33
|
+
layers=cfg["layers"],
|
|
34
|
+
heads=cfg["heads"],
|
|
35
|
+
ff_mult=cfg["ff_mult"],
|
|
36
|
+
dropout=cfg["dropout"],
|
|
37
|
+
max_seq_len=cfg["max_seq_len"],
|
|
38
|
+
task=task,
|
|
39
|
+
n_classes=meta.get("n_classes", 0),
|
|
40
|
+
pad_id=meta.get("pad_id", 0),
|
|
41
|
+
)
|
|
42
|
+
elif model_type == "mlp":
|
|
43
|
+
return TabularMLP(
|
|
44
|
+
n_numeric=meta["n_numeric"],
|
|
45
|
+
categorical_vocab_sizes=meta["categorical_vocab_sizes"],
|
|
46
|
+
d_model=cfg["d_model"],
|
|
47
|
+
layers=cfg["layers"],
|
|
48
|
+
dropout=cfg["dropout"],
|
|
49
|
+
task=task,
|
|
50
|
+
n_classes=meta.get("n_classes", 0),
|
|
51
|
+
)
|
|
52
|
+
else:
|
|
53
|
+
raise ModelError(f"Unknown model_type '{model_type}'.")
|
|
@@ -0,0 +1,175 @@
|
|
|
1
|
+
"""A small, dependency-free (beyond PyTorch) GPT-style decoder transformer.
|
|
2
|
+
|
|
3
|
+
Used for:
|
|
4
|
+
- "text-generation": next-token prediction over the char vocabulary
|
|
5
|
+
- "text-classification": same backbone, with a classification head on
|
|
6
|
+
the final token's hidden state instead of a language-modeling head
|
|
7
|
+
|
|
8
|
+
Kept intentionally compact -- this is not meant to compete with
|
|
9
|
+
production LLM training frameworks, it's meant to give Tensorless a real,
|
|
10
|
+
working, from-scratch model that trains fast enough on CPU for the
|
|
11
|
+
"zero setup" experience to actually be pleasant.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import math
|
|
17
|
+
from typing import Optional
|
|
18
|
+
|
|
19
|
+
import torch
|
|
20
|
+
import torch.nn as nn
|
|
21
|
+
import torch.nn.functional as F
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class CausalSelfAttention(nn.Module):
|
|
25
|
+
def __init__(self, d_model: int, heads: int, dropout: float):
|
|
26
|
+
super().__init__()
|
|
27
|
+
assert d_model % heads == 0, "d_model must be divisible by heads"
|
|
28
|
+
self.heads = heads
|
|
29
|
+
self.head_dim = d_model // heads
|
|
30
|
+
self.qkv = nn.Linear(d_model, 3 * d_model)
|
|
31
|
+
self.proj = nn.Linear(d_model, d_model)
|
|
32
|
+
self.dropout = dropout
|
|
33
|
+
self.resid_drop = nn.Dropout(dropout)
|
|
34
|
+
|
|
35
|
+
def forward(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
|
36
|
+
B, T, C = x.shape
|
|
37
|
+
qkv = self.qkv(x)
|
|
38
|
+
q, k, v = qkv.split(C, dim=2)
|
|
39
|
+
q = q.view(B, T, self.heads, self.head_dim).transpose(1, 2)
|
|
40
|
+
k = k.view(B, T, self.heads, self.head_dim).transpose(1, 2)
|
|
41
|
+
v = v.view(B, T, self.heads, self.head_dim).transpose(1, 2)
|
|
42
|
+
|
|
43
|
+
out = F.scaled_dot_product_attention(
|
|
44
|
+
q, k, v, attn_mask=attn_mask, dropout_p=self.dropout if self.training else 0.0, is_causal=attn_mask is None
|
|
45
|
+
)
|
|
46
|
+
out = out.transpose(1, 2).contiguous().view(B, T, C)
|
|
47
|
+
return self.resid_drop(self.proj(out))
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class MLP(nn.Module):
|
|
51
|
+
def __init__(self, d_model: int, ff_mult: int, dropout: float):
|
|
52
|
+
super().__init__()
|
|
53
|
+
self.fc1 = nn.Linear(d_model, d_model * ff_mult)
|
|
54
|
+
self.fc2 = nn.Linear(d_model * ff_mult, d_model)
|
|
55
|
+
self.drop = nn.Dropout(dropout)
|
|
56
|
+
|
|
57
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
58
|
+
return self.drop(self.fc2(F.gelu(self.fc1(x))))
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
class Block(nn.Module):
|
|
62
|
+
def __init__(self, d_model: int, heads: int, ff_mult: int, dropout: float):
|
|
63
|
+
super().__init__()
|
|
64
|
+
self.ln1 = nn.LayerNorm(d_model)
|
|
65
|
+
self.attn = CausalSelfAttention(d_model, heads, dropout)
|
|
66
|
+
self.ln2 = nn.LayerNorm(d_model)
|
|
67
|
+
self.mlp = MLP(d_model, ff_mult, dropout)
|
|
68
|
+
|
|
69
|
+
def forward(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
|
70
|
+
x = x + self.attn(self.ln1(x), attn_mask=attn_mask)
|
|
71
|
+
x = x + self.mlp(self.ln2(x))
|
|
72
|
+
return x
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
class TinyTransformer(nn.Module):
|
|
76
|
+
"""Decoder-only transformer usable for LM or sequence classification."""
|
|
77
|
+
|
|
78
|
+
def __init__(
|
|
79
|
+
self,
|
|
80
|
+
vocab_size: int,
|
|
81
|
+
d_model: int,
|
|
82
|
+
layers: int,
|
|
83
|
+
heads: int,
|
|
84
|
+
ff_mult: int,
|
|
85
|
+
dropout: float,
|
|
86
|
+
max_seq_len: int,
|
|
87
|
+
task: str = "text-generation",
|
|
88
|
+
n_classes: int = 0,
|
|
89
|
+
pad_id: int = 0,
|
|
90
|
+
):
|
|
91
|
+
super().__init__()
|
|
92
|
+
self.task = task
|
|
93
|
+
self.max_seq_len = max_seq_len
|
|
94
|
+
self.pad_id = pad_id
|
|
95
|
+
|
|
96
|
+
self.tok_emb = nn.Embedding(vocab_size, d_model)
|
|
97
|
+
self.pos_emb = nn.Embedding(max_seq_len, d_model)
|
|
98
|
+
self.drop = nn.Dropout(dropout)
|
|
99
|
+
self.blocks = nn.ModuleList(
|
|
100
|
+
[Block(d_model, heads, ff_mult, dropout) for _ in range(layers)]
|
|
101
|
+
)
|
|
102
|
+
self.ln_f = nn.LayerNorm(d_model)
|
|
103
|
+
|
|
104
|
+
if task == "text-generation":
|
|
105
|
+
self.head = nn.Linear(d_model, vocab_size, bias=False)
|
|
106
|
+
self.head.weight = self.tok_emb.weight # weight tying
|
|
107
|
+
elif task == "text-classification":
|
|
108
|
+
assert n_classes > 0, "n_classes must be set for text-classification"
|
|
109
|
+
self.head = nn.Linear(d_model, n_classes)
|
|
110
|
+
else:
|
|
111
|
+
raise ValueError(f"Unsupported task for TinyTransformer: {task}")
|
|
112
|
+
|
|
113
|
+
self.apply(self._init_weights)
|
|
114
|
+
|
|
115
|
+
def _init_weights(self, module: nn.Module) -> None:
|
|
116
|
+
if isinstance(module, nn.Linear):
|
|
117
|
+
nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
|
118
|
+
if module.bias is not None:
|
|
119
|
+
nn.init.zeros_(module.bias)
|
|
120
|
+
elif isinstance(module, nn.Embedding):
|
|
121
|
+
nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
|
122
|
+
|
|
123
|
+
def forward(self, input_ids: torch.Tensor, attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
|
124
|
+
B, T = input_ids.shape
|
|
125
|
+
assert T <= self.max_seq_len, (
|
|
126
|
+
f"Sequence length {T} exceeds max_seq_len {self.max_seq_len}"
|
|
127
|
+
)
|
|
128
|
+
pos = torch.arange(T, device=input_ids.device).unsqueeze(0)
|
|
129
|
+
x = self.tok_emb(input_ids) + self.pos_emb(pos)
|
|
130
|
+
x = self.drop(x)
|
|
131
|
+
|
|
132
|
+
attn_mask = None
|
|
133
|
+
if attention_mask is not None:
|
|
134
|
+
# Combine causal mask with padding mask.
|
|
135
|
+
causal = torch.tril(torch.ones(T, T, device=input_ids.device, dtype=torch.bool))
|
|
136
|
+
pad = attention_mask.bool().unsqueeze(1).unsqueeze(1) # B,1,1,T
|
|
137
|
+
attn_mask = (causal.unsqueeze(0).unsqueeze(0) & pad)
|
|
138
|
+
|
|
139
|
+
for block in self.blocks:
|
|
140
|
+
x = block(x, attn_mask=attn_mask)
|
|
141
|
+
x = self.ln_f(x)
|
|
142
|
+
|
|
143
|
+
if self.task == "text-generation":
|
|
144
|
+
return self.head(x) # B, T, vocab_size
|
|
145
|
+
else:
|
|
146
|
+
if attention_mask is not None:
|
|
147
|
+
lengths = attention_mask.sum(dim=1).clamp(min=1) - 1
|
|
148
|
+
else:
|
|
149
|
+
lengths = torch.full((B,), T - 1, device=input_ids.device)
|
|
150
|
+
pooled = x[torch.arange(B, device=input_ids.device), lengths]
|
|
151
|
+
return self.head(pooled) # B, n_classes
|
|
152
|
+
|
|
153
|
+
@torch.no_grad()
|
|
154
|
+
def generate(
|
|
155
|
+
self,
|
|
156
|
+
input_ids: torch.Tensor,
|
|
157
|
+
max_new_tokens: int,
|
|
158
|
+
temperature: float = 0.8,
|
|
159
|
+
top_k: Optional[int] = 40,
|
|
160
|
+
eos_id: Optional[int] = None,
|
|
161
|
+
) -> torch.Tensor:
|
|
162
|
+
self.eval()
|
|
163
|
+
for _ in range(max_new_tokens):
|
|
164
|
+
cond = input_ids[:, -self.max_seq_len:]
|
|
165
|
+
logits = self(cond)
|
|
166
|
+
logits = logits[:, -1, :] / max(temperature, 1e-5)
|
|
167
|
+
if top_k is not None:
|
|
168
|
+
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
|
|
169
|
+
logits[logits < v[:, [-1]]] = -float("inf")
|
|
170
|
+
probs = F.softmax(logits, dim=-1)
|
|
171
|
+
next_id = torch.multinomial(probs, num_samples=1)
|
|
172
|
+
input_ids = torch.cat([input_ids, next_id], dim=1)
|
|
173
|
+
if eos_id is not None and (next_id == eos_id).all():
|
|
174
|
+
break
|
|
175
|
+
return input_ids
|