tensorless 0.5.0__tar.gz → 0.7.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.
- {tensorless-0.5.0/tensorless.egg-info → tensorless-0.7.0}/PKG-INFO +1 -1
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/checkpointing.md +1 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/training.md +22 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/pyproject.toml +1 -1
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/__init__.py +2 -1
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/api.py +26 -1
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/auto/config.py +1 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/backends/jax_backend.py +11 -2
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/backends/mlx_backend.py +11 -2
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/config.py +2 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/engine.py +48 -7
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/models/transformer.py +65 -196
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/serialization/tl_format.py +1 -0
- tensorless-0.7.0/tensorless/training/trainer.py +148 -0
- {tensorless-0.5.0 → tensorless-0.7.0/tensorless.egg-info}/PKG-INFO +1 -1
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_train_tabular.py +11 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_train_text_generation.py +15 -0
- tensorless-0.5.0/tensorless/training/trainer.py +0 -369
- {tensorless-0.5.0 → tensorless-0.7.0}/LICENSE +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/MANIFEST.in +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/README.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/api_reference.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/architecture.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/automatic_mode.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/cli.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/configuration.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/contributing.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/examples.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/inference.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/installation.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/limitations.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/quickstart.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/roadmap.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/tl_format.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/troubleshooting.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/docs/tutorial.md +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/examples/tabular_classification_example.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/examples/tabular_regression_example.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/examples/text_classification_example.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/examples/text_generation_example.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/setup.cfg +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/_version.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/auto/__init__.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/auto/detector.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/backends/__init__.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/checkpoint/__init__.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/checkpoint/manager.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/cli/__init__.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/cli/main.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/data/__init__.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/data/english_grammar.txt +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/data/fingerprint.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/data/inspector.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/data/loader.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/data/tabular.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/devices/__init__.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/devices/device.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/devices/memory.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/errors.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/models/__init__.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/models/mlp.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/models/registry.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/runtime.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/serialization/__init__.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/tokenization/__init__.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/tokenization/bpe_tokenizer.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/tokenization/char_tokenizer.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/training/__init__.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/training/data_prep.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless/training/early_stopping.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless.egg-info/SOURCES.txt +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless.egg-info/dependency_links.txt +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless.egg-info/entry_points.txt +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless.egg-info/requires.txt +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tensorless.egg-info/top_level.txt +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_auto_detection.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_checkpoint_resume.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_cli.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_data_loading.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_end_to_end.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_fingerprint.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_jax_backend.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_mlx_backend.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_serialization.py +0 -0
- {tensorless-0.5.0 → tensorless-0.7.0}/tests/test_train_text_classification.py +0 -0
|
@@ -16,6 +16,7 @@ create or manage this directory yourself.
|
|
|
16
16
|
- `optimizer_state_dict` — optimizer momentum/variance buffers
|
|
17
17
|
- `scheduler_state_dict` — learning rate schedule position
|
|
18
18
|
- `epoch`, `global_step` — where training left off
|
|
19
|
+
- `train_loader_epoch` — deterministic shuffle position for resumed training
|
|
19
20
|
- `early_stopping_best`, `early_stopping_bad_checks` — early stopping state
|
|
20
21
|
- `config` — the fully resolved training configuration used
|
|
21
22
|
- `meta` — task-specific sizing info (vocab size, number of classes, etc.)
|
|
@@ -12,6 +12,28 @@ Returns a `LoadedModel` (see [inference.md](inference.md)) ready for
|
|
|
12
12
|
predictions, and writes `model.tl` (plus a `model.tl.ckpt/` checkpoint
|
|
13
13
|
directory) to the current directory.
|
|
14
14
|
|
|
15
|
+
## Pretraining and fine-tuning
|
|
16
|
+
|
|
17
|
+
Use the built-in corpus (or any text corpus) to create a base model, then
|
|
18
|
+
continue training its learned weights on your own data:
|
|
19
|
+
|
|
20
|
+
```python
|
|
21
|
+
base = tl.pretrain(out="english_pretrained.tl", epochs=10)
|
|
22
|
+
model = tl.train(
|
|
23
|
+
"./my_text.txt",
|
|
24
|
+
pretrained="english_pretrained.tl",
|
|
25
|
+
out="my_model.tl",
|
|
26
|
+
epochs=5,
|
|
27
|
+
learning_rate=1e-4,
|
|
28
|
+
)
|
|
29
|
+
```
|
|
30
|
+
|
|
31
|
+
The source tokenizer and compatible architecture are reused automatically.
|
|
32
|
+
The fine-tuning run starts a fresh optimizer and scheduler; interrupted
|
|
33
|
+
fine-tuning still resumes from its own checkpoint. Use
|
|
34
|
+
`tl.load_pretrained("english_pretrained.tl")` when you want to load the base
|
|
35
|
+
model directly for inference.
|
|
36
|
+
|
|
15
37
|
## Supported data formats
|
|
16
38
|
|
|
17
39
|
| Format | Notes |
|
|
@@ -12,7 +12,7 @@ ML with maximum automation and minimum setup.
|
|
|
12
12
|
See https://github.com/tensorless/tensorless for full documentation.
|
|
13
13
|
"""
|
|
14
14
|
|
|
15
|
-
from .api import train, pretrain, run, load, inspect
|
|
15
|
+
from .api import train, pretrain, run, load, load_pretrained, inspect
|
|
16
16
|
from .config import TrainConfig
|
|
17
17
|
from .errors import (
|
|
18
18
|
TensorlessError,
|
|
@@ -29,6 +29,7 @@ __all__ = [
|
|
|
29
29
|
"pretrain",
|
|
30
30
|
"run",
|
|
31
31
|
"load",
|
|
32
|
+
"load_pretrained",
|
|
32
33
|
"inspect",
|
|
33
34
|
"TrainConfig",
|
|
34
35
|
"TensorlessError",
|
|
@@ -55,7 +55,8 @@ def inspect(path: str) -> InspectionReport:
|
|
|
55
55
|
def train(path: str, **kwargs: Any) -> LoadedModel:
|
|
56
56
|
"""Train a model on the dataset at `path`, fully automatically by
|
|
57
57
|
default. Any field of `TrainConfig` can be overridden via keyword
|
|
58
|
-
argument, e.g. `tl.train("./data", d_model=512, layers=6)`.
|
|
58
|
+
argument, e.g. `tl.train("./data", d_model=512, layers=6)`. Pass
|
|
59
|
+
`pretrained="base.tl"` to fine-tune an existing Tensorless model.
|
|
59
60
|
|
|
60
61
|
Implements the "Smart Auto Check":
|
|
61
62
|
- if an up-to-date trained model already exists for this exact
|
|
@@ -65,6 +66,14 @@ def train(path: str, **kwargs: Any) -> LoadedModel:
|
|
|
65
66
|
`ask_on_data_change=True`), unless `force=True` is passed
|
|
66
67
|
"""
|
|
67
68
|
user_cfg = _build_train_config(**kwargs)
|
|
69
|
+
pretrained_state = None
|
|
70
|
+
if user_cfg.pretrained:
|
|
71
|
+
try:
|
|
72
|
+
pretrained_state = load_tl(user_cfg.pretrained)
|
|
73
|
+
except Exception as exc:
|
|
74
|
+
raise ModelError(
|
|
75
|
+
f"Could not load pretrained model '{user_cfg.pretrained}': {exc}"
|
|
76
|
+
) from exc
|
|
68
77
|
out = user_cfg.out or "model.tl"
|
|
69
78
|
checkpoint_dir = user_cfg.checkpoint_dir or (out + ".ckpt")
|
|
70
79
|
checkpoint_mgr = CheckpointManager(checkpoint_dir)
|
|
@@ -130,6 +139,16 @@ def train(path: str, **kwargs: Any) -> LoadedModel:
|
|
|
130
139
|
ds = load_dataset(path)
|
|
131
140
|
resolved = resolve_config(ds, user_cfg)
|
|
132
141
|
cfg = resolved.to_dict()
|
|
142
|
+
if pretrained_state is not None:
|
|
143
|
+
source_cfg = pretrained_state["config"]
|
|
144
|
+
for field in ("d_model", "layers", "heads", "ff_mult", "max_seq_len", "tokenizer", "bpe_vocab_size"):
|
|
145
|
+
if field not in user_cfg.overrides() and field in source_cfg:
|
|
146
|
+
cfg[field] = source_cfg[field]
|
|
147
|
+
if pretrained_state["task"] != cfg["task"] or pretrained_state["model_type"] != cfg["model_type"]:
|
|
148
|
+
raise ModelError(
|
|
149
|
+
"Pretrained model and target training data must use the same "
|
|
150
|
+
"task and model_type."
|
|
151
|
+
)
|
|
133
152
|
|
|
134
153
|
if resume_state is not None:
|
|
135
154
|
# Resumed runs must keep the exact architecture/config used
|
|
@@ -143,6 +162,7 @@ def train(path: str, **kwargs: Any) -> LoadedModel:
|
|
|
143
162
|
checkpoint_mgr=checkpoint_mgr,
|
|
144
163
|
dataset_fingerprint=fingerprint,
|
|
145
164
|
resume_state=resume_state,
|
|
165
|
+
pretrained_state=pretrained_state if resume_state is None else None,
|
|
146
166
|
log_fn=print if cfg.get("verbose", True) else (lambda *a, **k: None),
|
|
147
167
|
)
|
|
148
168
|
|
|
@@ -205,6 +225,11 @@ def load(path: str, device: Optional[str] = None) -> LoadedModel:
|
|
|
205
225
|
return load_model(path, device=device)
|
|
206
226
|
|
|
207
227
|
|
|
228
|
+
def load_pretrained(path: str, device: Optional[str] = None) -> LoadedModel:
|
|
229
|
+
"""Load a portable model intended to be used as a fine-tuning base."""
|
|
230
|
+
return load_model(path, device=device)
|
|
231
|
+
|
|
232
|
+
|
|
208
233
|
def run(path: str, prompt: Optional[str] = None) -> Any:
|
|
209
234
|
"""Run a trained `.tl` model.
|
|
210
235
|
|
|
@@ -101,6 +101,7 @@ def resolve_config(ds: Dataset, user: TrainConfig) -> ResolvedConfig:
|
|
|
101
101
|
max_seq_len = user.max_seq_len or (256 if ds.kind in ("text", "text_labeled") else 1)
|
|
102
102
|
resolved = ResolvedConfig(
|
|
103
103
|
out=out,
|
|
104
|
+
pretrained=user.pretrained,
|
|
104
105
|
force=bool(user.force),
|
|
105
106
|
resume=user.resume,
|
|
106
107
|
ask_on_data_change=bool(user.ask_on_data_change),
|
|
@@ -58,13 +58,20 @@ class JaxTinyTransformer(TinyTransformer):
|
|
|
58
58
|
|
|
59
59
|
def _forward_jax(self, params, input_ids, attention_mask=None):
|
|
60
60
|
_, jnp = _jax()
|
|
61
|
+
|
|
62
|
+
def layer_norm(x, gain, bias, eps=1e-5):
|
|
63
|
+
mean = jnp.mean(x, axis=-1, keepdims=True)
|
|
64
|
+
var = jnp.var(x, axis=-1, keepdims=True)
|
|
65
|
+
return (x - mean) / jnp.sqrt(var + eps) * gain + bias
|
|
66
|
+
|
|
61
67
|
input_ids = jnp.asarray(input_ids, dtype=jnp.int32)
|
|
62
68
|
batch, length = input_ids.shape
|
|
63
69
|
hidden = params["tok_emb"][input_ids] + params["pos_emb"][jnp.arange(length)]
|
|
64
70
|
causal = jnp.tril(jnp.ones((length, length), dtype=bool))
|
|
65
71
|
for index in range(len(self.blocks)):
|
|
66
72
|
prefix = f"blocks.{index}"
|
|
67
|
-
|
|
73
|
+
normed1 = layer_norm(hidden, params[f"{prefix}.ln1_gain"], params[f"{prefix}.ln1_bias"])
|
|
74
|
+
qkv = normed1 @ params[f"{prefix}.qkv_weight"] + params[f"{prefix}.qkv_bias"]
|
|
68
75
|
q, k, v = jnp.split(qkv, 3, axis=-1)
|
|
69
76
|
q = q.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
|
|
70
77
|
k = k.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
|
|
@@ -80,10 +87,12 @@ class JaxTinyTransformer(TinyTransformer):
|
|
|
80
87
|
context = context.transpose(0, 2, 1, 3).reshape(batch, length, self.d_model)
|
|
81
88
|
attention = context @ params[f"{prefix}.out_weight"] + params[f"{prefix}.out_bias"]
|
|
82
89
|
residual = hidden + attention
|
|
83
|
-
|
|
90
|
+
normed2 = layer_norm(residual, params[f"{prefix}.ln2_gain"], params[f"{prefix}.ln2_bias"])
|
|
91
|
+
ff_pre = normed2 @ params[f"{prefix}.ff1_weight"] + params[f"{prefix}.ff1_bias"]
|
|
84
92
|
ff_hidden = jnp.maximum(ff_pre, 0)
|
|
85
93
|
ff = ff_hidden @ params[f"{prefix}.ff2_weight"] + params[f"{prefix}.ff2_bias"]
|
|
86
94
|
hidden = residual + ff
|
|
95
|
+
hidden = layer_norm(hidden, params["ln_f_gain"], params["ln_f_bias"])
|
|
87
96
|
if self.task == "text-generation":
|
|
88
97
|
return hidden @ params["tok_emb"].T + params["head_bias"]
|
|
89
98
|
mask = jnp.ones((batch, length), dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask)
|
|
@@ -49,13 +49,20 @@ class MlxTinyTransformer(TinyTransformer):
|
|
|
49
49
|
|
|
50
50
|
def _forward_mlx(self, params, input_ids, attention_mask=None):
|
|
51
51
|
mx = _mlx()
|
|
52
|
+
|
|
53
|
+
def layer_norm(x, gain, bias, eps=1e-5):
|
|
54
|
+
mean = mx.mean(x, axis=-1, keepdims=True)
|
|
55
|
+
var = mx.var(x, axis=-1, keepdims=True)
|
|
56
|
+
return (x - mean) / mx.sqrt(var + eps) * gain + bias
|
|
57
|
+
|
|
52
58
|
input_ids = mx.array(input_ids, dtype=mx.int32)
|
|
53
59
|
batch, length = input_ids.shape
|
|
54
60
|
hidden = params["tok_emb"][input_ids] + params["pos_emb"][mx.arange(length)]
|
|
55
61
|
causal = mx.tril(mx.ones((length, length), dtype=mx.bool_))
|
|
56
62
|
for index in range(len(self.blocks)):
|
|
57
63
|
prefix = f"blocks.{index}"
|
|
58
|
-
|
|
64
|
+
normed1 = layer_norm(hidden, params[f"{prefix}.ln1_gain"], params[f"{prefix}.ln1_bias"])
|
|
65
|
+
qkv = normed1 @ params[f"{prefix}.qkv_weight"] + params[f"{prefix}.qkv_bias"]
|
|
59
66
|
q, k, v = mx.split(qkv, 3, axis=-1)
|
|
60
67
|
q = q.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
|
|
61
68
|
k = k.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
|
|
@@ -70,10 +77,12 @@ class MlxTinyTransformer(TinyTransformer):
|
|
|
70
77
|
context = context.transpose(0, 2, 1, 3).reshape(batch, length, self.d_model)
|
|
71
78
|
attention = context @ params[f"{prefix}.out_weight"] + params[f"{prefix}.out_bias"]
|
|
72
79
|
residual = hidden + attention
|
|
73
|
-
|
|
80
|
+
normed2 = layer_norm(residual, params[f"{prefix}.ln2_gain"], params[f"{prefix}.ln2_bias"])
|
|
81
|
+
ff_pre = normed2 @ params[f"{prefix}.ff1_weight"] + params[f"{prefix}.ff1_bias"]
|
|
74
82
|
ff_hidden = mx.maximum(ff_pre, 0)
|
|
75
83
|
ff = ff_hidden @ params[f"{prefix}.ff2_weight"] + params[f"{prefix}.ff2_bias"]
|
|
76
84
|
hidden = residual + ff
|
|
85
|
+
hidden = layer_norm(hidden, params["ln_f_gain"], params["ln_f_bias"])
|
|
77
86
|
if self.task == "text-generation":
|
|
78
87
|
return hidden @ params["tok_emb"].T + params["head_bias"]
|
|
79
88
|
mask = mx.ones((batch, length)) if attention_mask is None else mx.array(attention_mask)
|
|
@@ -20,6 +20,7 @@ class TrainConfig:
|
|
|
20
20
|
|
|
21
21
|
# --- output / lifecycle ---
|
|
22
22
|
out: Optional[str] = None # output .tl path, default "model.tl"
|
|
23
|
+
pretrained: Optional[str] = None # optional .tl weights to fine-tune
|
|
23
24
|
force: bool = False # force retraining even if unchanged
|
|
24
25
|
resume: Optional[bool] = None # force/forbid resume (None = auto)
|
|
25
26
|
ask_on_data_change: bool = False # raise instead of auto-retrain on data change
|
|
@@ -81,6 +82,7 @@ class ResolvedConfig:
|
|
|
81
82
|
"""
|
|
82
83
|
|
|
83
84
|
out: str
|
|
85
|
+
pretrained: Optional[str]
|
|
84
86
|
force: bool
|
|
85
87
|
resume: Optional[bool]
|
|
86
88
|
ask_on_data_change: bool
|
|
@@ -3,6 +3,7 @@
|
|
|
3
3
|
from __future__ import annotations
|
|
4
4
|
|
|
5
5
|
from typing import Dict, Iterable, Iterator
|
|
6
|
+
import math
|
|
6
7
|
import numpy as np
|
|
7
8
|
|
|
8
9
|
|
|
@@ -83,6 +84,39 @@ class Module:
|
|
|
83
84
|
return self.train(False)
|
|
84
85
|
|
|
85
86
|
|
|
87
|
+
def layer_norm_forward(x: np.ndarray, gain: np.ndarray, bias: np.ndarray, eps: float = 1e-5):
|
|
88
|
+
"""Normalize over the last axis, then apply a learned scale/shift.
|
|
89
|
+
|
|
90
|
+
Returns the normalized output plus a cache used by `layer_norm_backward`.
|
|
91
|
+
"""
|
|
92
|
+
mean = x.mean(axis=-1, keepdims=True)
|
|
93
|
+
var = x.var(axis=-1, keepdims=True)
|
|
94
|
+
inv_std = 1.0 / np.sqrt(var + eps)
|
|
95
|
+
x_hat = (x - mean) * inv_std
|
|
96
|
+
out = x_hat * gain + bias
|
|
97
|
+
return out.astype(np.float32), (x_hat, inv_std, gain)
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def layer_norm_backward(grad_out: np.ndarray, cache):
|
|
101
|
+
"""Backward pass matching `layer_norm_forward`.
|
|
102
|
+
|
|
103
|
+
Returns (grad_x, grad_gain, grad_bias).
|
|
104
|
+
"""
|
|
105
|
+
x_hat, inv_std, gain = cache
|
|
106
|
+
d = grad_out.shape[-1]
|
|
107
|
+
flat_grad_out = grad_out.reshape(-1, d)
|
|
108
|
+
flat_x_hat = x_hat.reshape(-1, d)
|
|
109
|
+
grad_gain = (flat_grad_out * flat_x_hat).sum(axis=0)
|
|
110
|
+
grad_bias = flat_grad_out.sum(axis=0)
|
|
111
|
+
grad_x_hat = grad_out * gain
|
|
112
|
+
grad_x = inv_std / d * (
|
|
113
|
+
d * grad_x_hat
|
|
114
|
+
- grad_x_hat.sum(axis=-1, keepdims=True)
|
|
115
|
+
- x_hat * (grad_x_hat * x_hat).sum(axis=-1, keepdims=True)
|
|
116
|
+
)
|
|
117
|
+
return grad_x.astype(np.float32), grad_gain.astype(np.float32), grad_bias.astype(np.float32)
|
|
118
|
+
|
|
119
|
+
|
|
86
120
|
def softmax_cross_entropy(logits: np.ndarray, targets: np.ndarray, ignore_index=None):
|
|
87
121
|
flat = logits.reshape(-1, logits.shape[-1]).astype(np.float64)
|
|
88
122
|
target = targets.reshape(-1).astype(np.int64)
|
|
@@ -188,18 +222,25 @@ class LambdaScheduler:
|
|
|
188
222
|
self.total_steps = total_steps
|
|
189
223
|
self.step_count = 0
|
|
190
224
|
self.base_lr = optimizer.lr
|
|
225
|
+
self.optimizer.lr = self.base_lr * self._factor(0)
|
|
226
|
+
|
|
227
|
+
def _factor(self, step):
|
|
228
|
+
if self.warmup_steps and step < self.warmup_steps:
|
|
229
|
+
return (step + 1) / self.warmup_steps
|
|
230
|
+
progress = min(
|
|
231
|
+
1.0,
|
|
232
|
+
max(0.0, (step - self.warmup_steps) /
|
|
233
|
+
max(1, self.total_steps - self.warmup_steps)),
|
|
234
|
+
)
|
|
235
|
+
return 0.1 + 0.9 * 0.5 * (1.0 + math.cos(math.pi * progress))
|
|
191
236
|
|
|
192
237
|
def step(self):
|
|
193
238
|
self.step_count += 1
|
|
194
|
-
|
|
195
|
-
if self.warmup_steps and step < self.warmup_steps:
|
|
196
|
-
factor = (step + 1) / self.warmup_steps
|
|
197
|
-
else:
|
|
198
|
-
factor = max(0.1, 1.0 - (step - self.warmup_steps) / max(1, self.total_steps - self.warmup_steps))
|
|
199
|
-
self.optimizer.lr = self.base_lr * factor
|
|
239
|
+
self.optimizer.lr = self.base_lr * self._factor(self.step_count)
|
|
200
240
|
|
|
201
241
|
def state_dict(self):
|
|
202
242
|
return {"step_count": self.step_count}
|
|
203
243
|
|
|
204
244
|
def load_state_dict(self, state):
|
|
205
|
-
self.step_count = int(state.get("step_count", 0))
|
|
245
|
+
self.step_count = int(state.get("step_count", 0))
|
|
246
|
+
self.optimizer.lr = self.base_lr * self._factor(self.step_count)
|
|
@@ -1,11 +1,19 @@
|
|
|
1
|
-
"""Compact NumPy text model retaining Tensorless transformer behavior.
|
|
1
|
+
"""Compact NumPy text model retaining Tensorless transformer behavior.
|
|
2
|
+
|
|
3
|
+
Architecture: a standard pre-norm decoder-only transformer -- LayerNorm
|
|
4
|
+
before attention and before the feed-forward block, plus a final LayerNorm
|
|
5
|
+
before the output head. Normalization is what keeps deep stacks of residual
|
|
6
|
+
blocks trainable; without it, gradients and activations drift as layers are
|
|
7
|
+
stacked and generation quality suffers noticeably once you go beyond one or
|
|
8
|
+
two blocks.
|
|
9
|
+
"""
|
|
2
10
|
|
|
3
11
|
from __future__ import annotations
|
|
4
12
|
|
|
5
13
|
from typing import Optional
|
|
6
14
|
import numpy as np
|
|
7
15
|
|
|
8
|
-
from ..engine import Module, Parameter, softmax_cross_entropy
|
|
16
|
+
from ..engine import Module, Parameter, softmax_cross_entropy, layer_norm_forward, layer_norm_backward
|
|
9
17
|
|
|
10
18
|
|
|
11
19
|
def _dropout(value, probability, training):
|
|
@@ -20,10 +28,14 @@ class TransformerBlock(Module):
|
|
|
20
28
|
if d_model % heads:
|
|
21
29
|
raise ValueError(f"d_model ({d_model}) must be divisible by heads ({heads})")
|
|
22
30
|
self.heads, self.head_dim, self.dropout = heads, d_model // heads, dropout
|
|
31
|
+
self.ln1_gain = Parameter(np.ones(d_model, dtype=np.float32))
|
|
32
|
+
self.ln1_bias = Parameter(np.zeros(d_model, dtype=np.float32))
|
|
23
33
|
self.qkv_weight = Parameter(np.random.normal(0, .02, (d_model, 3 * d_model)).astype(np.float32))
|
|
24
34
|
self.qkv_bias = Parameter(np.zeros(3 * d_model, dtype=np.float32))
|
|
25
35
|
self.out_weight = Parameter(np.random.normal(0, .02, (d_model, d_model)).astype(np.float32))
|
|
26
36
|
self.out_bias = Parameter(np.zeros(d_model, dtype=np.float32))
|
|
37
|
+
self.ln2_gain = Parameter(np.ones(d_model, dtype=np.float32))
|
|
38
|
+
self.ln2_bias = Parameter(np.zeros(d_model, dtype=np.float32))
|
|
27
39
|
hidden = d_model * ff_mult
|
|
28
40
|
self.ff1_weight = Parameter(np.random.normal(0, .02, (d_model, hidden)).astype(np.float32))
|
|
29
41
|
self.ff1_bias = Parameter(np.zeros(hidden, dtype=np.float32))
|
|
@@ -32,7 +44,8 @@ class TransformerBlock(Module):
|
|
|
32
44
|
|
|
33
45
|
def forward(self, inputs, attention_mask=None, cache=False):
|
|
34
46
|
batch, length, d_model = inputs.shape
|
|
35
|
-
|
|
47
|
+
normed1, ln1_cache = layer_norm_forward(inputs, self.ln1_gain.data, self.ln1_bias.data)
|
|
48
|
+
qkv = normed1 @ self.qkv_weight.data + self.qkv_bias.data
|
|
36
49
|
q, k, v = np.split(qkv, 3, axis=-1)
|
|
37
50
|
q, k, v = (value.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
|
|
38
51
|
for value in (q, k, v))
|
|
@@ -50,26 +63,31 @@ class TransformerBlock(Module):
|
|
|
50
63
|
context @ self.out_weight.data + self.out_bias.data, self.dropout, self.training
|
|
51
64
|
)
|
|
52
65
|
residual = inputs + attention
|
|
53
|
-
|
|
66
|
+
normed2, ln2_cache = layer_norm_forward(residual, self.ln2_gain.data, self.ln2_bias.data)
|
|
67
|
+
ff_pre = normed2 @ self.ff1_weight.data + self.ff1_bias.data
|
|
54
68
|
ff_hidden = np.maximum(ff_pre, 0)
|
|
55
69
|
ff, ff_keep = _dropout(ff_hidden @ self.ff2_weight.data + self.ff2_bias.data, self.dropout, self.training)
|
|
56
70
|
output = residual + ff
|
|
57
71
|
if not cache:
|
|
58
72
|
return output
|
|
59
|
-
return output, (inputs, q, k, v, probabilities, dropped_attention, attention_keep,
|
|
60
|
-
attention_keep_output, residual, ff_pre, ff_hidden, ff_keep)
|
|
73
|
+
return output, (inputs, normed1, ln1_cache, q, k, v, probabilities, dropped_attention, attention_keep,
|
|
74
|
+
context, attention_keep_output, residual, normed2, ln2_cache, ff_pre, ff_hidden, ff_keep)
|
|
61
75
|
|
|
62
76
|
def backward(self, gradient, cache):
|
|
63
|
-
(inputs, q, k, v, probabilities, dropped_attention, attention_keep, context,
|
|
64
|
-
attention_keep_output, residual, ff_pre, ff_hidden, ff_keep) = cache
|
|
77
|
+
(inputs, normed1, ln1_cache, q, k, v, probabilities, dropped_attention, attention_keep, context,
|
|
78
|
+
attention_keep_output, residual, normed2, ln2_cache, ff_pre, ff_hidden, ff_keep) = cache
|
|
65
79
|
gradient_ff = gradient if ff_keep is None else gradient * ff_keep / (1.0 - self.dropout)
|
|
66
80
|
self.ff2_weight.grad[...] = ff_hidden.reshape(-1, ff_hidden.shape[-1]).T @ gradient_ff.reshape(-1, gradient_ff.shape[-1])
|
|
67
81
|
self.ff2_bias.grad[...] = gradient_ff.sum(axis=(0, 1))
|
|
68
82
|
gradient_hidden = gradient_ff @ self.ff2_weight.data.T
|
|
69
83
|
gradient_hidden *= ff_pre > 0
|
|
70
|
-
self.ff1_weight.grad[...] =
|
|
84
|
+
self.ff1_weight.grad[...] = normed2.reshape(-1, normed2.shape[-1]).T @ gradient_hidden.reshape(-1, gradient_hidden.shape[-1])
|
|
71
85
|
self.ff1_bias.grad[...] = gradient_hidden.sum(axis=(0, 1))
|
|
72
|
-
|
|
86
|
+
gradient_normed2 = gradient_hidden @ self.ff1_weight.data.T
|
|
87
|
+
gradient_residual_from_ff, ln2_gain_grad, ln2_bias_grad = layer_norm_backward(gradient_normed2, ln2_cache)
|
|
88
|
+
self.ln2_gain.grad[...] = ln2_gain_grad
|
|
89
|
+
self.ln2_bias.grad[...] = ln2_bias_grad
|
|
90
|
+
gradient_residual = gradient + gradient_residual_from_ff
|
|
73
91
|
gradient_attention = gradient_residual if attention_keep_output is None else gradient_residual * attention_keep_output / (1.0 - self.dropout)
|
|
74
92
|
gradient_context = gradient_attention @ self.out_weight.data.T
|
|
75
93
|
self.out_weight.grad[...] = context.reshape(-1, context.shape[-1]).T @ gradient_attention.reshape(-1, gradient_attention.shape[-1])
|
|
@@ -83,10 +101,28 @@ class TransformerBlock(Module):
|
|
|
83
101
|
scale = 1.0 / np.sqrt(self.head_dim)
|
|
84
102
|
gradient_q = gradient_scores @ k * scale
|
|
85
103
|
gradient_k = gradient_scores.transpose(0, 1, 3, 2) @ q * scale
|
|
86
|
-
|
|
87
|
-
|
|
104
|
+
# Merge each of q/k/v back from (batch, heads, length, head_dim) to
|
|
105
|
+
# (batch, length, d_model) *before* concatenating them, matching the
|
|
106
|
+
# forward pass's layout ([q_block | k_block | v_block] along the last
|
|
107
|
+
# axis, each block itself head-major). Concatenating on the head-dim
|
|
108
|
+
# axis first (as a prior version of this code did) interleaves heads
|
|
109
|
+
# and q/k/v in the wrong order and silently corrupts the qkv_weight /
|
|
110
|
+
# qkv_bias gradients for any heads > 1 configuration.
|
|
111
|
+
batch, length = inputs.shape[0], inputs.shape[1]
|
|
112
|
+
|
|
113
|
+
def _merge_heads(gradient):
|
|
114
|
+
return gradient.transpose(0, 2, 1, 3).reshape(batch, length, -1)
|
|
115
|
+
|
|
116
|
+
gradient_qkv = np.concatenate(
|
|
117
|
+
(_merge_heads(gradient_q), _merge_heads(gradient_k), _merge_heads(gradient_v)), axis=-1
|
|
118
|
+
)
|
|
119
|
+
self.qkv_weight.grad[...] = normed1.reshape(-1, normed1.shape[-1]).T @ gradient_qkv.reshape(-1, gradient_qkv.shape[-1])
|
|
88
120
|
self.qkv_bias.grad[...] = gradient_qkv.sum(axis=(0, 1))
|
|
89
|
-
|
|
121
|
+
gradient_normed1 = gradient_qkv @ self.qkv_weight.data.T
|
|
122
|
+
gradient_inputs_from_attn, ln1_gain_grad, ln1_bias_grad = layer_norm_backward(gradient_normed1, ln1_cache)
|
|
123
|
+
self.ln1_gain.grad[...] = ln1_gain_grad
|
|
124
|
+
self.ln1_bias.grad[...] = ln1_bias_grad
|
|
125
|
+
return gradient_residual + gradient_inputs_from_attn
|
|
90
126
|
|
|
91
127
|
|
|
92
128
|
class TinyTransformer(Module):
|
|
@@ -100,6 +136,8 @@ class TinyTransformer(Module):
|
|
|
100
136
|
self.tok_emb = Parameter(np.random.normal(0, .02, (vocab_size, d_model)).astype(np.float32))
|
|
101
137
|
self.pos_emb = Parameter(np.random.normal(0, .02, (max_seq_len, d_model)).astype(np.float32))
|
|
102
138
|
self.blocks = [TransformerBlock(d_model, heads, ff_mult, dropout) for _ in range(layers)]
|
|
139
|
+
self.ln_f_gain = Parameter(np.ones(d_model, dtype=np.float32))
|
|
140
|
+
self.ln_f_bias = Parameter(np.zeros(d_model, dtype=np.float32))
|
|
103
141
|
if task == "text-generation":
|
|
104
142
|
self.head_bias = Parameter(np.zeros(vocab_size, dtype=np.float32))
|
|
105
143
|
self.head_weight = self.tok_emb
|
|
@@ -127,31 +165,35 @@ class TinyTransformer(Module):
|
|
|
127
165
|
caches.append(block_cache)
|
|
128
166
|
else:
|
|
129
167
|
hidden = block.forward(hidden, attention_mask)
|
|
168
|
+
hidden_normed, lnf_cache = layer_norm_forward(hidden, self.ln_f_gain.data, self.ln_f_bias.data)
|
|
130
169
|
if self.task == "text-generation":
|
|
131
|
-
out =
|
|
170
|
+
out = hidden_normed @ self.tok_emb.data.T + self.head_bias.data
|
|
132
171
|
else:
|
|
133
172
|
mask = np.ones((batch, length), dtype=np.float32) if attention_mask is None else attention_mask
|
|
134
|
-
pooled = (
|
|
173
|
+
pooled = (hidden_normed * mask[:, :, None]).sum(1) / np.maximum(mask.sum(1, keepdims=True), 1)
|
|
135
174
|
out = pooled @ self.head_weight.data + self.head_bias.data
|
|
136
|
-
return (out, (input_ids, hidden, attention_mask, caches)) if cache else out
|
|
175
|
+
return (out, (input_ids, hidden, hidden_normed, lnf_cache, attention_mask, caches)) if cache else out
|
|
137
176
|
|
|
138
177
|
def loss_and_backward(self, input_ids, target, attention_mask=None):
|
|
139
|
-
logits, (ids, hidden, mask, caches) = self.forward(input_ids, attention_mask, cache=True)
|
|
178
|
+
logits, (ids, hidden, hidden_normed, lnf_cache, mask, caches) = self.forward(input_ids, attention_mask, cache=True)
|
|
140
179
|
loss, grad = softmax_cross_entropy(logits, target, self.pad_id if self.task == "text-generation" else None)
|
|
141
180
|
for parameter in self.parameters():
|
|
142
181
|
parameter.grad.fill(0)
|
|
143
182
|
self.head_bias.grad[...] = grad.reshape(-1, grad.shape[-1]).sum(axis=0)
|
|
144
183
|
if self.task == "text-generation":
|
|
145
|
-
|
|
146
|
-
self.tok_emb.grad[...] += flat_g.T @
|
|
147
|
-
|
|
184
|
+
flat_hn, flat_g = hidden_normed.reshape(-1, hidden_normed.shape[-1]), grad.reshape(-1, grad.shape[-1])
|
|
185
|
+
self.tok_emb.grad[...] += flat_g.T @ flat_hn
|
|
186
|
+
d_hidden_normed = (flat_g @ self.tok_emb.data).reshape(hidden_normed.shape)
|
|
148
187
|
else:
|
|
149
188
|
valid_mask = np.ones(ids.shape, dtype=np.float32) if mask is None else mask
|
|
150
|
-
pooled = (
|
|
189
|
+
pooled = (hidden_normed * valid_mask[:, :, None]).sum(1) / np.maximum(valid_mask.sum(1, keepdims=True), 1)
|
|
151
190
|
self.head_weight.grad[...] = pooled.T @ grad
|
|
152
191
|
self.head_bias.grad[...] = grad.sum(0)
|
|
153
|
-
|
|
154
|
-
|
|
192
|
+
d_hidden_normed = (grad @ self.head_weight.data.T)[:, None, :] * valid_mask[:, :, None]
|
|
193
|
+
d_hidden_normed /= np.maximum(valid_mask.sum(1, keepdims=True)[:, :, None], 1)
|
|
194
|
+
dh, lnf_gain_grad, lnf_bias_grad = layer_norm_backward(d_hidden_normed, lnf_cache)
|
|
195
|
+
self.ln_f_gain.grad[...] = lnf_gain_grad
|
|
196
|
+
self.ln_f_bias.grad[...] = lnf_bias_grad
|
|
155
197
|
for block, block_cache in zip(self.blocks[::-1], caches[::-1]):
|
|
156
198
|
if self.gradient_checkpointing:
|
|
157
199
|
block_input, block_mask, rng_state = block_cache
|
|
@@ -178,176 +220,3 @@ class TinyTransformer(Module):
|
|
|
178
220
|
ids = np.concatenate([ids, next_ids], axis=1)
|
|
179
221
|
if eos_id is not None and np.all(next_ids == eos_id): break
|
|
180
222
|
return ids
|
|
181
|
-
"""A small, dependency-free legacy implementation retained only as inert text.
|
|
182
|
-
|
|
183
|
-
Used for:
|
|
184
|
-
- "text-generation": next-token prediction over the char vocabulary
|
|
185
|
-
- "text-classification": same backbone, with a classification head on
|
|
186
|
-
the final token's hidden state instead of a language-modeling head
|
|
187
|
-
|
|
188
|
-
Kept intentionally compact -- this is not meant to compete with
|
|
189
|
-
production LLM training frameworks, it's meant to give Tensorless a real,
|
|
190
|
-
working, from-scratch model that trains fast enough on CPU for the
|
|
191
|
-
"zero setup" experience to actually be pleasant.
|
|
192
|
-
|
|
193
|
-
from __future__ import annotations
|
|
194
|
-
|
|
195
|
-
import math
|
|
196
|
-
from typing import Optional
|
|
197
|
-
|
|
198
|
-
import numpy as np
|
|
199
|
-
|
|
200
|
-
|
|
201
|
-
class CausalSelfAttention(nn.Module):
|
|
202
|
-
def __init__(self, d_model: int, heads: int, dropout: float):
|
|
203
|
-
super().__init__()
|
|
204
|
-
assert d_model % heads == 0, "d_model must be divisible by heads"
|
|
205
|
-
self.heads = heads
|
|
206
|
-
self.head_dim = d_model // heads
|
|
207
|
-
self.qkv = nn.Linear(d_model, 3 * d_model)
|
|
208
|
-
self.proj = nn.Linear(d_model, d_model)
|
|
209
|
-
self.dropout = dropout
|
|
210
|
-
self.resid_drop = nn.Dropout(dropout)
|
|
211
|
-
|
|
212
|
-
def forward(self, x: np.ndarray, attn_mask: Optional[np.ndarray] = None) -> np.ndarray:
|
|
213
|
-
B, T, C = x.shape
|
|
214
|
-
qkv = self.qkv(x)
|
|
215
|
-
q, k, v = qkv.split(C, dim=2)
|
|
216
|
-
q = q.view(B, T, self.heads, self.head_dim).transpose(1, 2)
|
|
217
|
-
k = k.view(B, T, self.heads, self.head_dim).transpose(1, 2)
|
|
218
|
-
v = v.view(B, T, self.heads, self.head_dim).transpose(1, 2)
|
|
219
|
-
|
|
220
|
-
out = F.scaled_dot_product_attention(
|
|
221
|
-
q, k, v, attn_mask=attn_mask, dropout_p=self.dropout if self.training else 0.0, is_causal=attn_mask is None
|
|
222
|
-
)
|
|
223
|
-
out = out.transpose(1, 2).contiguous().view(B, T, C)
|
|
224
|
-
return self.resid_drop(self.proj(out))
|
|
225
|
-
|
|
226
|
-
|
|
227
|
-
class MLP(nn.Module):
|
|
228
|
-
def __init__(self, d_model: int, ff_mult: int, dropout: float):
|
|
229
|
-
super().__init__()
|
|
230
|
-
self.fc1 = nn.Linear(d_model, d_model * ff_mult)
|
|
231
|
-
self.fc2 = nn.Linear(d_model * ff_mult, d_model)
|
|
232
|
-
self.drop = nn.Dropout(dropout)
|
|
233
|
-
|
|
234
|
-
def forward(self, x: np.ndarray) -> np.ndarray:
|
|
235
|
-
return self.drop(self.fc2(F.gelu(self.fc1(x))))
|
|
236
|
-
|
|
237
|
-
|
|
238
|
-
class Block(nn.Module):
|
|
239
|
-
def __init__(self, d_model: int, heads: int, ff_mult: int, dropout: float):
|
|
240
|
-
super().__init__()
|
|
241
|
-
self.ln1 = nn.LayerNorm(d_model)
|
|
242
|
-
self.attn = CausalSelfAttention(d_model, heads, dropout)
|
|
243
|
-
self.ln2 = nn.LayerNorm(d_model)
|
|
244
|
-
self.mlp = MLP(d_model, ff_mult, dropout)
|
|
245
|
-
|
|
246
|
-
def forward(self, x: np.ndarray, attn_mask: Optional[np.ndarray] = None) -> np.ndarray:
|
|
247
|
-
x = x + self.attn(self.ln1(x), attn_mask=attn_mask)
|
|
248
|
-
x = x + self.mlp(self.ln2(x))
|
|
249
|
-
return x
|
|
250
|
-
|
|
251
|
-
|
|
252
|
-
class TinyTransformer(nn.Module):
|
|
253
|
-
Legacy decoder-only transformer usable for LM or sequence classification.
|
|
254
|
-
|
|
255
|
-
def __init__(
|
|
256
|
-
self,
|
|
257
|
-
vocab_size: int,
|
|
258
|
-
d_model: int,
|
|
259
|
-
layers: int,
|
|
260
|
-
heads: int,
|
|
261
|
-
ff_mult: int,
|
|
262
|
-
dropout: float,
|
|
263
|
-
max_seq_len: int,
|
|
264
|
-
task: str = "text-generation",
|
|
265
|
-
n_classes: int = 0,
|
|
266
|
-
pad_id: int = 0,
|
|
267
|
-
):
|
|
268
|
-
super().__init__()
|
|
269
|
-
self.task = task
|
|
270
|
-
self.max_seq_len = max_seq_len
|
|
271
|
-
self.pad_id = pad_id
|
|
272
|
-
|
|
273
|
-
self.tok_emb = nn.Embedding(vocab_size, d_model)
|
|
274
|
-
self.pos_emb = nn.Embedding(max_seq_len, d_model)
|
|
275
|
-
self.drop = nn.Dropout(dropout)
|
|
276
|
-
self.blocks = nn.ModuleList(
|
|
277
|
-
[Block(d_model, heads, ff_mult, dropout) for _ in range(layers)]
|
|
278
|
-
)
|
|
279
|
-
self.ln_f = nn.LayerNorm(d_model)
|
|
280
|
-
|
|
281
|
-
if task == "text-generation":
|
|
282
|
-
self.head = nn.Linear(d_model, vocab_size, bias=False)
|
|
283
|
-
self.head.weight = self.tok_emb.weight # weight tying
|
|
284
|
-
elif task == "text-classification":
|
|
285
|
-
assert n_classes > 0, "n_classes must be set for text-classification"
|
|
286
|
-
self.head = nn.Linear(d_model, n_classes)
|
|
287
|
-
else:
|
|
288
|
-
raise ValueError(f"Unsupported task for TinyTransformer: {task}")
|
|
289
|
-
|
|
290
|
-
self.apply(self._init_weights)
|
|
291
|
-
|
|
292
|
-
def _init_weights(self, module: nn.Module) -> None:
|
|
293
|
-
if isinstance(module, nn.Linear):
|
|
294
|
-
nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
|
295
|
-
if module.bias is not None:
|
|
296
|
-
nn.init.zeros_(module.bias)
|
|
297
|
-
elif isinstance(module, nn.Embedding):
|
|
298
|
-
nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
|
299
|
-
|
|
300
|
-
def forward(self, input_ids: torch.Tensor, attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
|
301
|
-
B, T = input_ids.shape
|
|
302
|
-
assert T <= self.max_seq_len, (
|
|
303
|
-
f"Sequence length {T} exceeds max_seq_len {self.max_seq_len}"
|
|
304
|
-
)
|
|
305
|
-
pos = torch.arange(T, device=input_ids.device).unsqueeze(0)
|
|
306
|
-
x = self.tok_emb(input_ids) + self.pos_emb(pos)
|
|
307
|
-
x = self.drop(x)
|
|
308
|
-
|
|
309
|
-
attn_mask = None
|
|
310
|
-
if attention_mask is not None:
|
|
311
|
-
# Combine causal mask with padding mask.
|
|
312
|
-
causal = torch.tril(torch.ones(T, T, device=input_ids.device, dtype=torch.bool))
|
|
313
|
-
pad = attention_mask.bool().unsqueeze(1).unsqueeze(1) # B,1,1,T
|
|
314
|
-
attn_mask = (causal.unsqueeze(0).unsqueeze(0) & pad)
|
|
315
|
-
|
|
316
|
-
for block in self.blocks:
|
|
317
|
-
x = block(x, attn_mask=attn_mask)
|
|
318
|
-
x = self.ln_f(x)
|
|
319
|
-
|
|
320
|
-
if self.task == "text-generation":
|
|
321
|
-
return self.head(x) # B, T, vocab_size
|
|
322
|
-
else:
|
|
323
|
-
if attention_mask is not None:
|
|
324
|
-
lengths = attention_mask.sum(dim=1).clamp(min=1) - 1
|
|
325
|
-
else:
|
|
326
|
-
lengths = torch.full((B,), T - 1, device=input_ids.device)
|
|
327
|
-
pooled = x[torch.arange(B, device=input_ids.device), lengths]
|
|
328
|
-
return self.head(pooled) # B, n_classes
|
|
329
|
-
|
|
330
|
-
@torch.no_grad()
|
|
331
|
-
def generate(
|
|
332
|
-
self,
|
|
333
|
-
input_ids: np.ndarray,
|
|
334
|
-
max_new_tokens: int,
|
|
335
|
-
temperature: float = 0.8,
|
|
336
|
-
top_k: Optional[int] = 40,
|
|
337
|
-
eos_id: Optional[int] = None,
|
|
338
|
-
) -> np.ndarray:
|
|
339
|
-
self.eval()
|
|
340
|
-
for _ in range(max_new_tokens):
|
|
341
|
-
cond = input_ids[:, -self.max_seq_len:]
|
|
342
|
-
logits = self(cond)
|
|
343
|
-
logits = logits[:, -1, :] / max(temperature, 1e-5)
|
|
344
|
-
if top_k is not None:
|
|
345
|
-
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
|
|
346
|
-
logits[logits < v[:, [-1]]] = -float("inf")
|
|
347
|
-
probs = F.softmax(logits, dim=-1)
|
|
348
|
-
next_id = torch.multinomial(probs, num_samples=1)
|
|
349
|
-
input_ids = torch.cat([input_ids, next_id], dim=1)
|
|
350
|
-
if eos_id is not None and (next_id == eos_id).all():
|
|
351
|
-
break
|
|
352
|
-
return input_ids
|
|
353
|
-
"""
|