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 +86 -0
- dabsn/adapters.py +118 -0
- dabsn/checkpoint.py +231 -0
- dabsn/cli.py +568 -0
- dabsn/config.py +293 -0
- dabsn/core.py +305 -0
- dabsn/distributed.py +653 -0
- dabsn/kernels/__init__.py +171 -0
- dabsn/kernels/cpu.py +543 -0
- dabsn/kernels/cpu_runtime.cpp +1261 -0
- dabsn/kernels/local_field.py +231 -0
- dabsn/kernels/long.py +422 -0
- dabsn/kernels/permanent.py +383 -0
- dabsn/kernels/portable_math.py +164 -0
- dabsn/kernels/triton.py +242 -0
- dabsn/kernels/triton_runtime.py +3144 -0
- dabsn/model.py +468 -0
- dabsn/pretrain.py +528 -0
- dabsn/read.py +711 -0
- dabsn/reproductions/__init__.py +1 -0
- dabsn/reproductions/a5.py +8 -0
- dabsn/reproductions/common.py +236 -0
- dabsn/reproductions/copy.py +8 -0
- dabsn/reproductions/keyvalue.py +8 -0
- dabsn/reproductions/mqar.py +8 -0
- dabsn/runtime/__init__.py +58 -0
- dabsn/runtime/api.py +203 -0
- dabsn-0.1.0.dist-info/METADATA +535 -0
- dabsn-0.1.0.dist-info/RECORD +33 -0
- dabsn-0.1.0.dist-info/WHEEL +5 -0
- dabsn-0.1.0.dist-info/entry_points.txt +6 -0
- dabsn-0.1.0.dist-info/licenses/LICENSE +201 -0
- dabsn-0.1.0.dist-info/top_level.txt +1 -0
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
|
+
]
|