dabsn 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.
dabsn/__init__.py ADDED
@@ -0,0 +1,86 @@
1
+ """Canonical DABSN framework."""
2
+
3
+ from .config import (
4
+ DABSN_ARCH,
5
+ DABSNConfig,
6
+ DABSNLayerSpec,
7
+ DABSNPretrainConfig,
8
+ parse_dabsn_layer_specs,
9
+ resolve_dabsn_layers,
10
+ resolve_layer_geometries,
11
+ )
12
+ from .core import DABSNCore
13
+ from .pretrain import pretrain_next_token
14
+ from .read import DABSNRead
15
+ from .checkpoint import dabsn_config_dict, inspect_dabsn, load_dabsn, save_dabsn
16
+ from .model import (
17
+ DABSNBackbone,
18
+ DABSNBlock,
19
+ DABSNModel,
20
+ DABSNSequenceLM,
21
+ DABSNTaskModel,
22
+ build_dabsn,
23
+ build_dabsn_from_config,
24
+ dabsn_adamw_param_groups,
25
+ )
26
+ from .runtime import (
27
+ DABSNSequenceModule,
28
+ DistributedState,
29
+ cleanup_distributed,
30
+ load_distributed_optimizer,
31
+ load_sharded_model_checkpoint,
32
+ load_sharded_training_checkpoint,
33
+ prepare_distributed_model,
34
+ save_distributed_dabsn,
35
+ save_sharded_training_checkpoint,
36
+ setup_distributed,
37
+ train,
38
+ train_step,
39
+ evaluate,
40
+ export_dabsn,
41
+ infer,
42
+ verify_gradients,
43
+ )
44
+
45
+ __all__ = [
46
+ "DABSN_ARCH",
47
+ "DABSNConfig",
48
+ "DABSNCore",
49
+ "DABSNBackbone",
50
+ "DABSNBlock",
51
+ "DABSNLayerSpec",
52
+ "DABSNPretrainConfig",
53
+ "DABSNRead",
54
+ "DABSNModel",
55
+ "DABSNSequenceLM",
56
+ "DABSNTaskModel",
57
+ "DABSNSequenceModule",
58
+ "DistributedState",
59
+ "build_dabsn",
60
+ "build_dabsn_from_config",
61
+ "dabsn_adamw_param_groups",
62
+ "dabsn_config_dict",
63
+ "cleanup_distributed",
64
+ "evaluate",
65
+ "export_dabsn",
66
+ "infer",
67
+ "load_dabsn",
68
+ "load_distributed_optimizer",
69
+ "load_sharded_model_checkpoint",
70
+ "load_sharded_training_checkpoint",
71
+ "inspect_dabsn",
72
+ "parse_dabsn_layer_specs",
73
+ "resolve_dabsn_layers",
74
+ "resolve_layer_geometries",
75
+ "prepare_distributed_model",
76
+ "pretrain_next_token",
77
+ "save_dabsn",
78
+ "save_distributed_dabsn",
79
+ "save_sharded_training_checkpoint",
80
+ "setup_distributed",
81
+ "train",
82
+ "train_step",
83
+ "verify_gradients",
84
+ ]
85
+
86
+ __version__ = "0.1.0"
dabsn/adapters.py ADDED
@@ -0,0 +1,118 @@
1
+ """Built-in DABSN adapters and registration APIs for task-specific modules."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Callable
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+ from torch import Tensor
10
+
11
+
12
+ class IdentityInputAdapter(nn.Module):
13
+ """Already-vectorized sequence or field input."""
14
+
15
+ def __init__(self, dim: int) -> None:
16
+ super().__init__()
17
+ self.output_dim = dim
18
+
19
+ def forward(self, inputs: Tensor) -> Tensor:
20
+ return inputs.float()
21
+
22
+
23
+ class LinearInputAdapter(nn.Module):
24
+ """Continuous features projected to the DABSN input width."""
25
+
26
+ def __init__(self, input_dim: int, output_dim: int) -> None:
27
+ super().__init__()
28
+ self.output_dim = output_dim
29
+ self.proj = nn.Linear(input_dim, output_dim)
30
+
31
+ def forward(self, inputs: Tensor) -> Tensor:
32
+ return self.proj(inputs.float())
33
+
34
+
35
+ class ByteInputAdapter(nn.Module):
36
+ """Byte ids ``[B,T]`` to vectors ``[B,T,D]``."""
37
+
38
+ def __init__(self, output_dim: int, vocab_size: int = 256) -> None:
39
+ super().__init__()
40
+ self.output_dim = output_dim
41
+ self.emb = nn.Embedding(vocab_size, output_dim)
42
+
43
+ def forward(self, inputs: Tensor) -> Tensor:
44
+ return self.emb(inputs.long())
45
+
46
+
47
+ class LinearOutputHead(nn.Module):
48
+ """Per-position linear output head."""
49
+
50
+ def __init__(self, input_dim: int, out_dim: int) -> None:
51
+ super().__init__()
52
+ self.proj = nn.Linear(input_dim, out_dim)
53
+
54
+ def forward(self, hidden: Tensor) -> Tensor:
55
+ return self.proj(hidden)
56
+
57
+
58
+ class ByteOutputHead(LinearOutputHead):
59
+ def __init__(self, input_dim: int, out_dim: int = 256) -> None:
60
+ del out_dim
61
+ super().__init__(input_dim, 256)
62
+
63
+
64
+ InputBuilder = Callable[[int, int | None], nn.Module]
65
+ OutputBuilder = Callable[[int, int], nn.Module]
66
+
67
+ _INPUT_ADAPTERS: dict[str, InputBuilder] = {
68
+ "identity": lambda input_dim, output_dim: IdentityInputAdapter(
69
+ input_dim if output_dim is None else output_dim
70
+ ),
71
+ "tensor": lambda input_dim, output_dim: IdentityInputAdapter(
72
+ input_dim if output_dim is None else output_dim
73
+ ),
74
+ "linear": lambda input_dim, output_dim: LinearInputAdapter(
75
+ input_dim,
76
+ input_dim if output_dim is None else output_dim,
77
+ ),
78
+ "byte": lambda input_dim, output_dim: ByteInputAdapter(
79
+ input_dim if output_dim is None else output_dim
80
+ ),
81
+ }
82
+
83
+ _OUTPUT_HEADS: dict[str, OutputBuilder] = {
84
+ "linear": LinearOutputHead,
85
+ "field": LinearOutputHead,
86
+ "token": LinearOutputHead,
87
+ "byte": ByteOutputHead,
88
+ }
89
+
90
+
91
+ def register_input_adapter(kind: str, builder: InputBuilder) -> None:
92
+ """Register an input-adapter builder under a case-insensitive name."""
93
+ _INPUT_ADAPTERS[kind.lower()] = builder
94
+
95
+
96
+ def register_output_head(kind: str, builder: OutputBuilder) -> None:
97
+ """Register an output-head builder under a case-insensitive name."""
98
+ _OUTPUT_HEADS[kind.lower()] = builder
99
+
100
+
101
+ def build_input_adapter(
102
+ kind: str,
103
+ input_dim: int,
104
+ output_dim: int | None = None,
105
+ ) -> nn.Module:
106
+ """Construct a registered input adapter."""
107
+ normalized = kind.lower()
108
+ if normalized not in _INPUT_ADAPTERS:
109
+ raise ValueError(f"unknown DABSN input adapter: {normalized}")
110
+ return _INPUT_ADAPTERS[normalized](input_dim, output_dim)
111
+
112
+
113
+ def build_output_head(kind: str, input_dim: int, out_dim: int) -> nn.Module:
114
+ """Construct a registered output head."""
115
+ normalized = kind.lower()
116
+ if normalized not in _OUTPUT_HEADS:
117
+ raise ValueError(f"unknown DABSN output head: {normalized}")
118
+ return _OUTPUT_HEADS[normalized](input_dim, out_dim)
dabsn/checkpoint.py ADDED
@@ -0,0 +1,231 @@
1
+ """Self-describing, non-pickle DABSN model checkpoints."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import os
7
+ from pathlib import Path
8
+ from typing import Mapping
9
+
10
+ import torch
11
+ from torch import Tensor
12
+
13
+ from .config import DABSNConfig, DABSNLayerSpec
14
+ from .model import (
15
+ DABSNModel,
16
+ DABSNSequenceLM,
17
+ DABSNTaskModel,
18
+ build_dabsn_from_config,
19
+ )
20
+
21
+ Model = DABSNModel | DABSNTaskModel | DABSNSequenceLM
22
+
23
+ CHECKPOINT_FORMAT = "dabsn-model"
24
+ CHECKPOINT_VERSION = 1
25
+
26
+
27
+ def _introspect_config(model: DABSNModel | DABSNTaskModel) -> DABSNConfig:
28
+ if isinstance(model, DABSNTaskModel):
29
+ body = model.body
30
+ backbone = body.backbone
31
+ return DABSNConfig(
32
+ input_dim=int(model.raw_input_dim),
33
+ out_dim=int(model.out_dim),
34
+ hidden_dim=int(model.model_input_dim),
35
+ input_adapter=model.input_adapter_kind,
36
+ output_adapter=body.output_adapter_kind,
37
+ layers=[spec.to_metadata() for spec in backbone.layer_specs],
38
+ )
39
+ return DABSNConfig(
40
+ input_dim=int(model.input_dim),
41
+ out_dim=int(model.out_dim),
42
+ hidden_dim=int(model.input_dim),
43
+ input_adapter="identity",
44
+ output_adapter=model.output_adapter_kind,
45
+ layers=[spec.to_metadata() for spec in model.backbone.layer_specs],
46
+ )
47
+
48
+
49
+ def dabsn_config_dict(model: Model) -> dict[str, object]:
50
+ """Return the complete architecture needed to reconstruct ``model``."""
51
+
52
+ if isinstance(model, DABSNSequenceLM):
53
+ return {
54
+ "model_kind": "sequence_lm",
55
+ "vocab": int(model.vocab),
56
+ "hidden_dim": int(model.hidden_dim),
57
+ "depth": int(model.depth),
58
+ "state_dim": model.state_dim,
59
+ "layers": [spec.to_metadata() for spec in model.layers],
60
+ "tie_embeddings": bool(model.tie_embeddings),
61
+ "grad_checkpoint": bool(model.backbone.grad_checkpoint),
62
+ }
63
+ config = getattr(model, "_dabsn_config", None) or _introspect_config(model)
64
+ data = {name: getattr(config, name) for name in DABSNConfig.__dataclass_fields__}
65
+ data["layers"] = [
66
+ spec.to_metadata() if isinstance(spec, DABSNLayerSpec) else dict(spec)
67
+ for spec in config.layer_specs()
68
+ ]
69
+ data["model_kind"] = "task" if isinstance(model, DABSNTaskModel) else "model"
70
+ data["grad_checkpoint"] = bool(model.backbone.grad_checkpoint)
71
+ return data
72
+
73
+
74
+ def build_dabsn_from_checkpoint_config(config: Mapping[str, object]) -> Model:
75
+ """Construct a model from clean checkpoint metadata without loading weights."""
76
+
77
+ if config.get("model_kind") == "sequence_lm":
78
+ return DABSNSequenceLM(
79
+ vocab=int(config["vocab"]),
80
+ hidden_dim=int(config["hidden_dim"]),
81
+ depth=int(config["depth"]),
82
+ layers=config["layers"],
83
+ state_dim=config.get("state_dim"),
84
+ tie_embeddings=bool(config.get("tie_embeddings", False)),
85
+ grad_checkpoint=bool(config.get("grad_checkpoint", False)),
86
+ )
87
+ config_fields = {
88
+ name: value
89
+ for name, value in config.items()
90
+ if name in DABSNConfig.__dataclass_fields__
91
+ }
92
+ return build_dabsn_from_config(
93
+ DABSNConfig(**config_fields),
94
+ grad_checkpoint=bool(config.get("grad_checkpoint", False)),
95
+ )
96
+
97
+
98
+ def _storage_identity(tensor: Tensor) -> tuple[int, int, tuple[int, ...], tuple[int, ...]]:
99
+ try:
100
+ pointer = tensor.untyped_storage().data_ptr()
101
+ except AttributeError:
102
+ pointer = tensor.storage().data_ptr()
103
+ return pointer, tensor.storage_offset(), tuple(tensor.shape), tuple(tensor.stride())
104
+
105
+
106
+ def _deduplicate_state_dict(
107
+ state_dict: Mapping[str, Tensor],
108
+ ) -> tuple[dict[str, Tensor], dict[str, str]]:
109
+ saved: dict[str, Tensor] = {}
110
+ shared: dict[str, str] = {}
111
+ seen: dict[tuple[int, int, tuple[int, ...], tuple[int, ...]], str] = {}
112
+ for name, tensor in state_dict.items():
113
+ identity = _storage_identity(tensor)
114
+ if identity in seen:
115
+ shared[name] = seen[identity]
116
+ else:
117
+ seen[identity] = name
118
+ saved[name] = tensor.detach().cpu().contiguous()
119
+ return saved, shared
120
+
121
+
122
+ def save_dabsn_state(
123
+ state_dict: Mapping[str, Tensor],
124
+ config: Mapping[str, object],
125
+ path: str | Path,
126
+ *,
127
+ extra: Mapping[str, object] | None = None,
128
+ ) -> None:
129
+ """Atomically write a clean DABSN SafeTensors checkpoint."""
130
+
131
+ from safetensors.torch import save_file
132
+
133
+ destination = Path(path)
134
+ destination.parent.mkdir(parents=True, exist_ok=True)
135
+ tensors, shared = _deduplicate_state_dict(state_dict)
136
+ metadata = {
137
+ "format": CHECKPOINT_FORMAT,
138
+ "version": str(CHECKPOINT_VERSION),
139
+ "config": json.dumps(dict(config), sort_keys=True, separators=(",", ":")),
140
+ "extra": json.dumps(dict(extra or {}), sort_keys=True, separators=(",", ":")),
141
+ "shared": json.dumps(shared, sort_keys=True, separators=(",", ":")),
142
+ }
143
+ temporary = destination.with_name(destination.name + ".tmp")
144
+ try:
145
+ save_file(tensors, str(temporary), metadata=metadata)
146
+ os.replace(temporary, destination)
147
+ finally:
148
+ temporary.unlink(missing_ok=True)
149
+
150
+
151
+ def save_dabsn(
152
+ model: Model,
153
+ path: str | Path,
154
+ *,
155
+ extra: Mapping[str, object] | None = None,
156
+ ) -> None:
157
+ """Save ``model`` as a self-describing SafeTensors checkpoint."""
158
+
159
+ save_dabsn_state(model.state_dict(), dabsn_config_dict(model), path, extra=extra)
160
+
161
+
162
+ def inspect_dabsn(path: str | Path) -> dict[str, object]:
163
+ """Read clean checkpoint metadata without allocating model tensors."""
164
+
165
+ from safetensors import safe_open
166
+
167
+ source = Path(path)
168
+ try:
169
+ with safe_open(str(source), framework="pt") as checkpoint:
170
+ metadata = checkpoint.metadata() or {}
171
+ except Exception as exc:
172
+ raise ValueError(f"{source} is not a readable SafeTensors checkpoint") from exc
173
+ if metadata.get("format") != CHECKPOINT_FORMAT:
174
+ raise ValueError(f"{source} is not a {CHECKPOINT_FORMAT} checkpoint")
175
+ version = int(metadata.get("version", "0"))
176
+ if version != CHECKPOINT_VERSION:
177
+ raise ValueError(
178
+ f"unsupported {CHECKPOINT_FORMAT} version {version}; expected {CHECKPOINT_VERSION}"
179
+ )
180
+ try:
181
+ config = json.loads(metadata["config"])
182
+ extra = json.loads(metadata.get("extra", "{}"))
183
+ shared = json.loads(metadata.get("shared", "{}"))
184
+ except (KeyError, TypeError, json.JSONDecodeError) as exc:
185
+ raise ValueError(f"{source} has invalid DABSN metadata") from exc
186
+ if not isinstance(config, dict) or not isinstance(extra, dict) or not isinstance(shared, dict):
187
+ raise ValueError(f"{source} has invalid DABSN metadata objects")
188
+ return {
189
+ "format": CHECKPOINT_FORMAT,
190
+ "version": version,
191
+ "config": config,
192
+ "extra": extra,
193
+ "shared": shared,
194
+ }
195
+
196
+
197
+ def load_dabsn(
198
+ path: str | Path,
199
+ *,
200
+ map_location: str | torch.device | None = None,
201
+ strict: bool = True,
202
+ ) -> Model:
203
+ """Load only the clean DABSN SafeTensors format; no legacy pickle fallback."""
204
+
205
+ from safetensors.torch import load_file
206
+
207
+ source = Path(path)
208
+ metadata = inspect_dabsn(source)
209
+ device = torch.device("cpu" if map_location is None else map_location)
210
+ model = build_dabsn_from_checkpoint_config(metadata["config"]).to(device)
211
+ state_dict = load_file(str(source), device=str(device))
212
+ for duplicate, original in metadata["shared"].items():
213
+ if original not in state_dict:
214
+ raise ValueError(
215
+ f"checkpoint shared tensor {duplicate!r} refers to missing {original!r}"
216
+ )
217
+ state_dict[duplicate] = state_dict[original]
218
+ model.load_state_dict(state_dict, strict=strict)
219
+ return model
220
+
221
+
222
+ __all__ = [
223
+ "CHECKPOINT_FORMAT",
224
+ "CHECKPOINT_VERSION",
225
+ "build_dabsn_from_checkpoint_config",
226
+ "dabsn_config_dict",
227
+ "inspect_dabsn",
228
+ "load_dabsn",
229
+ "save_dabsn",
230
+ "save_dabsn_state",
231
+ ]