laya-coreml 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.
- laya_coreml/__init__.py +6 -0
- laya_coreml/__main__.py +3 -0
- laya_coreml/agent.py +91 -0
- laya_coreml/ane.py +167 -0
- laya_coreml/artifacts.py +90 -0
- laya_coreml/cli.py +73 -0
- laya_coreml/common.py +119 -0
- laya_coreml/convert.py +189 -0
- laya_coreml/hub.py +30 -0
- laya_coreml/inputs.py +38 -0
- laya_coreml/prompt.py +54 -0
- laya_coreml/result.py +65 -0
- laya_coreml/snake/__init__.py +2 -0
- laya_coreml/snake/__main__.py +4 -0
- laya_coreml/snake/benchmark.py +301 -0
- laya_coreml/snake/cli.py +300 -0
- laya_coreml/snake/game.py +157 -0
- laya_coreml/snake/policy.py +213 -0
- laya_coreml/snake/replay.py +257 -0
- laya_coreml/snake/ui.py +170 -0
- laya_coreml/tokenizer.py +27 -0
- laya_coreml/torch_model.py +237 -0
- laya_coreml-0.1.0.dist-info/METADATA +198 -0
- laya_coreml-0.1.0.dist-info/RECORD +28 -0
- laya_coreml-0.1.0.dist-info/WHEEL +4 -0
- laya_coreml-0.1.0.dist-info/entry_points.txt +3 -0
- laya_coreml-0.1.0.dist-info/licenses/LICENSE +176 -0
- laya_coreml-0.1.0.dist-info/licenses/NOTICE +22 -0
laya_coreml/__init__.py
ADDED
laya_coreml/__main__.py
ADDED
laya_coreml/agent.py
ADDED
|
@@ -0,0 +1,91 @@
|
|
|
1
|
+
"""Core ML inference. No MLX, Transformers, or PyTorch dependency at runtime."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import math
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
|
|
8
|
+
from .artifacts import package_for_coreml
|
|
9
|
+
from .hub import DEFAULT_MODEL, resolve_checkpoint
|
|
10
|
+
from .prompt import PromptMixin
|
|
11
|
+
from .result import ResultMixin
|
|
12
|
+
from .tokenizer import Tokenizer
|
|
13
|
+
|
|
14
|
+
COMPUTE_UNITS = {"all": "ALL", "cpu": "CPU_ONLY", "cpu_gpu": "CPU_AND_GPU", "cpu_ne": "CPU_AND_NE"}
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class Agent(PromptMixin, ResultMixin):
|
|
18
|
+
def __init__(
|
|
19
|
+
self,
|
|
20
|
+
model_dir,
|
|
21
|
+
*,
|
|
22
|
+
compute_units="cpu_gpu",
|
|
23
|
+
allow_unvalidated_gpu=False,
|
|
24
|
+
revision=None,
|
|
25
|
+
local_files_only=False,
|
|
26
|
+
):
|
|
27
|
+
if compute_units not in COMPUTE_UNITS:
|
|
28
|
+
raise ValueError(f"compute_units must be one of {list(COMPUTE_UNITS)}")
|
|
29
|
+
self.model_dir = resolve_checkpoint(
|
|
30
|
+
model_dir, revision=revision, local_files_only=local_files_only
|
|
31
|
+
)
|
|
32
|
+
self.manifest = json.loads((self.model_dir / "coreml_config.json").read_text())
|
|
33
|
+
if self.manifest.get("format") != "laya-coreml" or self.manifest.get("format_version") != 1:
|
|
34
|
+
raise ValueError("Unsupported Core ML export format")
|
|
35
|
+
self.shape = self.manifest["shape"]
|
|
36
|
+
if (
|
|
37
|
+
compute_units == "cpu_gpu"
|
|
38
|
+
and self.shape["flexible"]
|
|
39
|
+
and not self.shape.get("lengths")
|
|
40
|
+
and not allow_unvalidated_gpu
|
|
41
|
+
):
|
|
42
|
+
raise ValueError(
|
|
43
|
+
"RangeDim + CPU_AND_GPU failed local fidelity and repeatability checks. "
|
|
44
|
+
"Re-export with the default enumerated shapes, or use compute_units='cpu'. "
|
|
45
|
+
"allow_unvalidated_gpu=True is for reproducing the failure only."
|
|
46
|
+
)
|
|
47
|
+
self.cfg = json.loads((self.model_dir / "rl_agent_config.json").read_text())
|
|
48
|
+
self.temperature = self.cfg.get("temperature", [1.0, 1.0, 1.0])
|
|
49
|
+
self.temperature_by_options = self.cfg.get("temperature_by_options", {})
|
|
50
|
+
if len(self.temperature) != 3 or any(
|
|
51
|
+
not math.isfinite(float(t)) or float(t) <= 0
|
|
52
|
+
for t in [*self.temperature, *self.temperature_by_options.values()]
|
|
53
|
+
):
|
|
54
|
+
raise ValueError("Calibration temperatures must be finite and positive")
|
|
55
|
+
self.tok = Tokenizer(self.model_dir / "tokenizer")
|
|
56
|
+
self.batch_size = self.shape["batch_size"]
|
|
57
|
+
self.pad_to_multiple = 16
|
|
58
|
+
import coremltools as ct
|
|
59
|
+
|
|
60
|
+
self.compute_units = compute_units
|
|
61
|
+
self.model = ct.models.MLModel(
|
|
62
|
+
str(package_for_coreml(self.model_dir / "model.mlpackage")),
|
|
63
|
+
compute_units=getattr(ct.ComputeUnit, COMPUTE_UNITS[compute_units]),
|
|
64
|
+
)
|
|
65
|
+
|
|
66
|
+
def forward(self, batch):
|
|
67
|
+
outputs = self.model.predict(batch)
|
|
68
|
+
return np.asarray(outputs["logits"], np.float32), np.asarray(
|
|
69
|
+
outputs["action_logits"], np.float32
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def load(
|
|
74
|
+
model_dir=DEFAULT_MODEL,
|
|
75
|
+
*,
|
|
76
|
+
revision=None,
|
|
77
|
+
local_files_only=False,
|
|
78
|
+
compute_units=None,
|
|
79
|
+
allow_unvalidated_gpu=False,
|
|
80
|
+
):
|
|
81
|
+
directory = resolve_checkpoint(model_dir, revision=revision, local_files_only=local_files_only)
|
|
82
|
+
manifest = json.loads((directory / "coreml_config.json").read_text())
|
|
83
|
+
if manifest.get("format") == "laya-coreml-ane":
|
|
84
|
+
from .ane import ANEAgent
|
|
85
|
+
|
|
86
|
+
return ANEAgent(directory, compute_units=compute_units or "cpu_ne")
|
|
87
|
+
return Agent(
|
|
88
|
+
directory,
|
|
89
|
+
compute_units=compute_units or "cpu_gpu",
|
|
90
|
+
allow_unvalidated_gpu=allow_unvalidated_gpu,
|
|
91
|
+
)
|
laya_coreml/ane.py
ADDED
|
@@ -0,0 +1,167 @@
|
|
|
1
|
+
"""Host embedding lookup + one ANE graph + CPU action-head runtime."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import math
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
|
|
7
|
+
import coremltools as ct
|
|
8
|
+
import numpy as np
|
|
9
|
+
from safetensors import safe_open
|
|
10
|
+
|
|
11
|
+
from laya_coreml.prompt import PromptMixin
|
|
12
|
+
from laya_coreml.result import ResultMixin
|
|
13
|
+
from laya_coreml.tokenizer import Tokenizer
|
|
14
|
+
|
|
15
|
+
from .artifacts import package_for_coreml, verify_files, verify_research_manifest
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class ANEAgent(PromptMixin, ResultMixin):
|
|
19
|
+
def __init__(self, source, package=None, *, length=None, compute_units="cpu_ne"):
|
|
20
|
+
self.source = Path(source)
|
|
21
|
+
if package is None:
|
|
22
|
+
self.manifest = json.loads((self.source / "coreml_config.json").read_text())
|
|
23
|
+
if (self.manifest.get("format"), self.manifest.get("format_version")) != (
|
|
24
|
+
"laya-coreml-ane",
|
|
25
|
+
1,
|
|
26
|
+
):
|
|
27
|
+
raise ValueError("Unsupported ANE bundle format")
|
|
28
|
+
shape = self.manifest["shape"]
|
|
29
|
+
if shape["batch_size"] != 1 or shape["max_options"] != 32 or shape["flexible"]:
|
|
30
|
+
raise ValueError("ANE runtime requires a fixed B1/K32 bundle")
|
|
31
|
+
if length is not None and length != shape["max_length"]:
|
|
32
|
+
raise ValueError("Requested length does not match the ANE bundle")
|
|
33
|
+
length = shape["max_length"]
|
|
34
|
+
verify_files(self.source, self.manifest["files"])
|
|
35
|
+
package = self.source / "model.mlpackage"
|
|
36
|
+
host_weights = self.source / "host_weights.safetensors"
|
|
37
|
+
else:
|
|
38
|
+
length = 96 if length is None else length
|
|
39
|
+
self.manifest = verify_research_manifest(self.source, package, length=length)
|
|
40
|
+
host_weights = self.source / "model.safetensors"
|
|
41
|
+
self.model_dir = self.source
|
|
42
|
+
self.cfg = json.loads((self.source / "rl_agent_config.json").read_text())
|
|
43
|
+
self.temperature = self.cfg.get("temperature", [1.0, 1.0, 1.0])
|
|
44
|
+
self.temperature_by_options = self.cfg.get("temperature_by_options", {})
|
|
45
|
+
self.tok = Tokenizer(self.source / "tokenizer")
|
|
46
|
+
self.batch_size, self.pad_to_multiple = 1, 16
|
|
47
|
+
self.shape = {
|
|
48
|
+
"batch_size": 1,
|
|
49
|
+
"max_length": length,
|
|
50
|
+
"min_length": length,
|
|
51
|
+
"max_options": 32,
|
|
52
|
+
"flexible": False,
|
|
53
|
+
"lengths": None,
|
|
54
|
+
}
|
|
55
|
+
self.compute_units = compute_units
|
|
56
|
+
units = {
|
|
57
|
+
"cpu_ne": ct.ComputeUnit.CPU_AND_NE,
|
|
58
|
+
"cpu_gpu": ct.ComputeUnit.CPU_AND_GPU,
|
|
59
|
+
"all": ct.ComputeUnit.ALL,
|
|
60
|
+
"cpu": ct.ComputeUnit.CPU_ONLY,
|
|
61
|
+
}
|
|
62
|
+
if compute_units not in units:
|
|
63
|
+
raise ValueError(f"compute_units must be one of {list(units)}")
|
|
64
|
+
self.model = ct.models.MLModel(
|
|
65
|
+
str(package_for_coreml(package)), compute_units=units[compute_units]
|
|
66
|
+
)
|
|
67
|
+
self.encoder_cfg = json.loads((self.source / "encoder/config.json").read_text())
|
|
68
|
+
width = int(self.encoder_cfg["hidden_size"])
|
|
69
|
+
expected_shapes = {
|
|
70
|
+
"embeddings": (1, width, 1, length),
|
|
71
|
+
"full_mask": (1, length, 1, length),
|
|
72
|
+
"local_mask": (1, length, 1, length),
|
|
73
|
+
"type_vectors": (1, width, 1, 1),
|
|
74
|
+
"marker_map": (1, length, 1, 32),
|
|
75
|
+
}
|
|
76
|
+
actual_shapes = {
|
|
77
|
+
feature.name: tuple(feature.type.multiArrayType.shape)
|
|
78
|
+
for feature in self.model.get_spec().description.input
|
|
79
|
+
}
|
|
80
|
+
if actual_shapes != expected_shapes:
|
|
81
|
+
raise ValueError(
|
|
82
|
+
f"Package signature mismatch: expected {expected_shapes}, got {actual_shapes}"
|
|
83
|
+
)
|
|
84
|
+
with safe_open(str(host_weights), framework="numpy") as weights:
|
|
85
|
+
self.embedding = weights.get_tensor("encoder.embeddings.tok_embeddings.weight")
|
|
86
|
+
self.type_embedding = weights.get_tensor("type_emb.weight")
|
|
87
|
+
self.action = {
|
|
88
|
+
key: weights.get_tensor("act_head." + key).astype(np.float32)
|
|
89
|
+
for key in ("0.weight", "0.bias", "2.weight", "2.bias")
|
|
90
|
+
}
|
|
91
|
+
spec = self.model.get_spec()
|
|
92
|
+
self.output_names = [output.name for output in spec.description.output]
|
|
93
|
+
self._erf = np.frompyfunc(math.erf, 1, 1)
|
|
94
|
+
positions = np.arange(length)
|
|
95
|
+
self.window = (
|
|
96
|
+
np.abs(positions[:, None] - positions[None, :])
|
|
97
|
+
<= int(self.encoder_cfg.get("local_attention", 128)) // 2
|
|
98
|
+
)
|
|
99
|
+
|
|
100
|
+
def model_inputs(self, batch):
|
|
101
|
+
expected_shapes = {
|
|
102
|
+
"input_ids": (1, self.shape["max_length"]),
|
|
103
|
+
"attention_mask": (1, self.shape["max_length"]),
|
|
104
|
+
"marker_pos": (1, 32),
|
|
105
|
+
"marker_mask": (1, 32),
|
|
106
|
+
"qtype": (1,),
|
|
107
|
+
}
|
|
108
|
+
if set(batch) != set(expected_shapes):
|
|
109
|
+
raise ValueError("Prepared batch fields do not match the fixed ANE signature")
|
|
110
|
+
for name, shape in expected_shapes.items():
|
|
111
|
+
if batch[name].shape != shape or not np.issubdtype(batch[name].dtype, np.integer):
|
|
112
|
+
raise ValueError(f"{name} must have integer dtype and shape {shape}")
|
|
113
|
+
if np.any(batch["input_ids"] < 0) or np.any(batch["input_ids"] >= self.embedding.shape[0]):
|
|
114
|
+
raise ValueError("Token id outside checkpoint vocabulary")
|
|
115
|
+
for name in ("attention_mask", "marker_mask"):
|
|
116
|
+
if not np.isin(batch[name], (0, 1)).all():
|
|
117
|
+
raise ValueError(f"{name} must contain only zero or one")
|
|
118
|
+
if not batch["attention_mask"].any(axis=-1).all():
|
|
119
|
+
raise ValueError("Every batch row needs at least one valid attention key")
|
|
120
|
+
if np.any(batch["qtype"] < 0) or np.any(batch["qtype"] > 2):
|
|
121
|
+
raise ValueError("Question type must be 0, 1 or 2")
|
|
122
|
+
if np.any(batch["marker_pos"] < 0) or np.any(
|
|
123
|
+
batch["marker_pos"] >= self.shape["max_length"]
|
|
124
|
+
):
|
|
125
|
+
raise ValueError("Marker position outside exported sequence")
|
|
126
|
+
ids, valid = batch["input_ids"], batch["attention_mask"].astype(bool)
|
|
127
|
+
embeddings = self.embedding[ids].transpose(0, 2, 1)[:, :, None, :]
|
|
128
|
+
full = np.broadcast_to(valid[:, None, :], (ids.shape[0], ids.shape[1], ids.shape[1]))
|
|
129
|
+
local = (self.window[None] | ~valid[:, :, None]) & full
|
|
130
|
+
# Core ML BC1S attention scores are [B,key,1,query].
|
|
131
|
+
masks = {"full_mask": full, "local_mask": local}
|
|
132
|
+
result = {
|
|
133
|
+
name: np.where(value.transpose(0, 2, 1)[:, :, None, :], 0, -1e4).astype(np.float16)
|
|
134
|
+
for name, value in masks.items()
|
|
135
|
+
}
|
|
136
|
+
result["embeddings"] = np.ascontiguousarray(embeddings, dtype=np.float16)
|
|
137
|
+
result["type_vectors"] = np.ascontiguousarray(
|
|
138
|
+
self.type_embedding[batch["qtype"]][:, :, None, None], dtype=np.float16
|
|
139
|
+
)
|
|
140
|
+
marker_map = np.zeros((ids.shape[0], ids.shape[1], 1, 32), np.float16)
|
|
141
|
+
for row in range(ids.shape[0]):
|
|
142
|
+
marker_map[row, batch["marker_pos"][row], 0, np.arange(32)] = 1
|
|
143
|
+
result["marker_map"] = marker_map
|
|
144
|
+
return result
|
|
145
|
+
|
|
146
|
+
def forward(self, batch):
|
|
147
|
+
outputs = self.model.predict(self.model_inputs(batch))
|
|
148
|
+
# Output names are traced identifiers; shapes uniquely identify these outputs.
|
|
149
|
+
logits = (
|
|
150
|
+
next(v for v in outputs.values() if v.shape[1] == 1).reshape(1, 32).astype(np.float32)
|
|
151
|
+
)
|
|
152
|
+
pooled = (
|
|
153
|
+
next(v for v in outputs.values() if v.shape[1] != 1).reshape(1, -1).astype(np.float32)
|
|
154
|
+
)
|
|
155
|
+
logits = np.where(batch["marker_mask"].astype(bool), logits, -1e4)
|
|
156
|
+
p = np.exp(logits - logits.max(axis=-1, keepdims=True))
|
|
157
|
+
p /= p.sum(axis=-1, keepdims=True)
|
|
158
|
+
k = np.maximum(batch["marker_mask"].sum(axis=-1), 2).astype(np.float32)
|
|
159
|
+
entropy = -(p * np.log(np.maximum(p, 1e-9))).sum(axis=-1) / np.log(k)
|
|
160
|
+
top = np.sort(p, axis=-1)[:, -2:]
|
|
161
|
+
features = np.stack((top[:, 1], top[:, 1] - top[:, 0], entropy, k / 255.0), axis=-1)
|
|
162
|
+
action_input = np.concatenate((pooled, features), axis=-1)
|
|
163
|
+
hidden = action_input @ self.action["0.weight"].T + self.action["0.bias"]
|
|
164
|
+
# Only 256 host elements: exact erf GELU, no tanh/sigmoid approximation.
|
|
165
|
+
hidden = hidden * (1 + self._erf(hidden / np.sqrt(2)).astype(np.float32)) / 2
|
|
166
|
+
action = hidden @ self.action["2.weight"].T + self.action["2.bias"]
|
|
167
|
+
return logits, action
|
laya_coreml/artifacts.py
ADDED
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
"""Integrity checks for portable Core ML bundles."""
|
|
2
|
+
|
|
3
|
+
import errno
|
|
4
|
+
import hashlib
|
|
5
|
+
import os
|
|
6
|
+
import shutil
|
|
7
|
+
import tempfile
|
|
8
|
+
from pathlib import Path, PurePosixPath
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def file_digest(path):
|
|
12
|
+
digest = hashlib.sha256()
|
|
13
|
+
with Path(path).open("rb") as stream:
|
|
14
|
+
for block in iter(lambda: stream.read(8 * 1024**2), b""):
|
|
15
|
+
digest.update(block)
|
|
16
|
+
return digest.hexdigest()
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def tree_digest(path):
|
|
20
|
+
path = Path(path)
|
|
21
|
+
digest = hashlib.sha256()
|
|
22
|
+
for file in sorted(path.rglob("*")):
|
|
23
|
+
if file.is_file():
|
|
24
|
+
digest.update(str(file.relative_to(path)).encode())
|
|
25
|
+
digest.update(bytes.fromhex(file_digest(file)))
|
|
26
|
+
return digest.hexdigest()
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def package_for_coreml(package):
|
|
30
|
+
"""Materialize Hub symlinks: the native compiler can copy them into broken paths."""
|
|
31
|
+
package = Path(package)
|
|
32
|
+
if not package.is_symlink() and not any(path.is_symlink() for path in package.rglob("*")):
|
|
33
|
+
return package
|
|
34
|
+
digest = tree_digest(package)
|
|
35
|
+
root = (
|
|
36
|
+
Path(os.environ.get("LAYA_COREML_CACHE", Path.home() / ".cache/laya-coreml")) / "packages"
|
|
37
|
+
)
|
|
38
|
+
root.mkdir(parents=True, exist_ok=True)
|
|
39
|
+
destination = root / digest
|
|
40
|
+
target = destination / "model.mlpackage"
|
|
41
|
+
if destination.exists():
|
|
42
|
+
if not target.is_dir() or tree_digest(target) != digest:
|
|
43
|
+
raise ValueError(f"Core ML cache integrity failure; remove {destination} and retry")
|
|
44
|
+
return target
|
|
45
|
+
temporary = Path(tempfile.mkdtemp(prefix=".preparing-", dir=root))
|
|
46
|
+
try:
|
|
47
|
+
copied = temporary / "model.mlpackage"
|
|
48
|
+
shutil.copytree(package, copied, symlinks=False)
|
|
49
|
+
if tree_digest(copied) != digest:
|
|
50
|
+
raise ValueError("Core ML package changed while materializing cached weights")
|
|
51
|
+
try:
|
|
52
|
+
temporary.rename(destination)
|
|
53
|
+
except OSError as error:
|
|
54
|
+
if error.errno not in (errno.EEXIST, errno.ENOTEMPTY) or not target.is_dir():
|
|
55
|
+
raise
|
|
56
|
+
if tree_digest(target) != digest:
|
|
57
|
+
raise ValueError("Materialized Core ML package failed its integrity check")
|
|
58
|
+
finally:
|
|
59
|
+
shutil.rmtree(temporary, ignore_errors=True)
|
|
60
|
+
return target
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def verify_files(directory, files):
|
|
64
|
+
for name, expected in files.items():
|
|
65
|
+
relative = PurePosixPath(name)
|
|
66
|
+
if relative.is_absolute() or ".." in relative.parts or "\\" in name:
|
|
67
|
+
raise ValueError("Bundle file names must be relative paths inside the model directory")
|
|
68
|
+
if file_digest(Path(directory) / name) != expected["sha256"]:
|
|
69
|
+
raise ValueError(f"Bundle file does not match its manifest: {name}")
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def verify_research_manifest(source, package, *, length):
|
|
73
|
+
import json
|
|
74
|
+
|
|
75
|
+
source, package = Path(source), Path(package)
|
|
76
|
+
manifest = json.loads((package.parent / "manifest.json").read_text())
|
|
77
|
+
if manifest.get("format") != "laya-ane-research" or manifest.get("version") != 1:
|
|
78
|
+
raise ValueError("Unsupported ANE research artifact manifest")
|
|
79
|
+
if manifest.get("kind") != "body" or manifest.get("shape") != {
|
|
80
|
+
"batch": 1,
|
|
81
|
+
"length": length,
|
|
82
|
+
"options": 32,
|
|
83
|
+
}:
|
|
84
|
+
raise ValueError("Requested runtime shape does not match artifact manifest")
|
|
85
|
+
verify_files(
|
|
86
|
+
source, {name: {"sha256": sha} for name, sha in manifest["source_files_sha256"].items()}
|
|
87
|
+
)
|
|
88
|
+
if tree_digest(package) != manifest["package_sha256"]:
|
|
89
|
+
raise ValueError("ANE package content does not match its manifest")
|
|
90
|
+
return manifest
|
laya_coreml/cli.py
ADDED
|
@@ -0,0 +1,73 @@
|
|
|
1
|
+
"""Convert and run Laya Core ML packages."""
|
|
2
|
+
|
|
3
|
+
import argparse
|
|
4
|
+
import json
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def main():
|
|
8
|
+
parser = argparse.ArgumentParser(description=__doc__)
|
|
9
|
+
commands = parser.add_subparsers(dest="command", required=True)
|
|
10
|
+
export = commands.add_parser("convert", help="Export original Laya weights to Core ML")
|
|
11
|
+
export.add_argument("source", help="Local original checkpoint, or pinned Laya model name")
|
|
12
|
+
export.add_argument("output")
|
|
13
|
+
export.add_argument("--max-length", type=int)
|
|
14
|
+
export.add_argument("--batch-size", type=int, default=1)
|
|
15
|
+
export.add_argument("--max-options", type=int, default=32)
|
|
16
|
+
export.add_argument("--fixed", action="store_true")
|
|
17
|
+
export.add_argument("--precision", choices=["float16", "float32"], default="float16")
|
|
18
|
+
export.add_argument("--revision")
|
|
19
|
+
export.add_argument("--attention", choices=["explicit", "sdpa"], default="sdpa")
|
|
20
|
+
export.add_argument("--shape-mode", choices=["enumerated", "range"], default="enumerated")
|
|
21
|
+
predict = commands.add_parser("predict")
|
|
22
|
+
predict.add_argument("model_dir")
|
|
23
|
+
predict.add_argument("--state", required=True, help="Literal state text")
|
|
24
|
+
predict.add_argument(
|
|
25
|
+
"--questions", required=True, help="JSON file containing question definitions"
|
|
26
|
+
)
|
|
27
|
+
predict.add_argument(
|
|
28
|
+
"--compute-units",
|
|
29
|
+
choices=["all", "cpu", "cpu_gpu", "cpu_ne"],
|
|
30
|
+
help="Default: cpu_ne for ANE bundles; cpu_gpu for ordinary exports",
|
|
31
|
+
)
|
|
32
|
+
predict.add_argument(
|
|
33
|
+
"--offline", action="store_true", help="Use local files or cached Hub snapshots only"
|
|
34
|
+
)
|
|
35
|
+
predict.add_argument("--revision", help="Pinned Hugging Face commit or revision")
|
|
36
|
+
args = parser.parse_args()
|
|
37
|
+
if args.command == "convert":
|
|
38
|
+
from .convert import convert
|
|
39
|
+
|
|
40
|
+
convert(
|
|
41
|
+
args.source,
|
|
42
|
+
args.output,
|
|
43
|
+
max_length=args.max_length,
|
|
44
|
+
flexible=not args.fixed,
|
|
45
|
+
batch_size=args.batch_size,
|
|
46
|
+
max_options=args.max_options,
|
|
47
|
+
precision=args.precision,
|
|
48
|
+
revision=args.revision,
|
|
49
|
+
attention=args.attention,
|
|
50
|
+
shape_mode=args.shape_mode,
|
|
51
|
+
)
|
|
52
|
+
else:
|
|
53
|
+
from pathlib import Path
|
|
54
|
+
|
|
55
|
+
from .agent import load
|
|
56
|
+
|
|
57
|
+
agent = load(
|
|
58
|
+
args.model_dir,
|
|
59
|
+
compute_units=args.compute_units,
|
|
60
|
+
local_files_only=args.offline,
|
|
61
|
+
revision=args.revision,
|
|
62
|
+
)
|
|
63
|
+
print(
|
|
64
|
+
json.dumps(
|
|
65
|
+
agent.predict(args.state, json.loads(Path(args.questions).read_text())),
|
|
66
|
+
ensure_ascii=False,
|
|
67
|
+
indent=2,
|
|
68
|
+
)
|
|
69
|
+
)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
if __name__ == "__main__":
|
|
73
|
+
main()
|
laya_coreml/common.py
ADDED
|
@@ -0,0 +1,119 @@
|
|
|
1
|
+
"""Laya prompt construction and calibration, adapted from upstream (see NOTICE)."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import math
|
|
5
|
+
from typing import Dict, List, Optional, Union
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
QTYPES = {"choice": 0, "score": 1, "noul": 2}
|
|
10
|
+
QTYPE_NAMES = {v: k for k, v in QTYPES.items()}
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def serialize_state(state: Union[str, dict, list]) -> str:
|
|
14
|
+
if isinstance(state, str):
|
|
15
|
+
return state
|
|
16
|
+
return json.dumps(state, ensure_ascii=False)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def render_criterion(value) -> str:
|
|
20
|
+
"""Render one criterion value as text.
|
|
21
|
+
|
|
22
|
+
Strings pass through; anything structured (dict, list, number) becomes compact JSON, so a
|
|
23
|
+
rubric reads as JSON rather than a Python repr. Without this a dict-valued criterion
|
|
24
|
+
crashed `noul` outright and leaked `{'desc': ...}` into `choice` and `score` prompts.
|
|
25
|
+
"""
|
|
26
|
+
if isinstance(value, str):
|
|
27
|
+
return value
|
|
28
|
+
return json.dumps(value, ensure_ascii=False, separators=(", ", ": "), default=str)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def render_options(q: Dict) -> List[str]:
|
|
32
|
+
"""Render option texts in label-index order. Noul is always [false, true]."""
|
|
33
|
+
t, crit = q["t"], q.get("crit")
|
|
34
|
+
if t == "choice":
|
|
35
|
+
# only None/"" mean "no description"; 0 and False are legitimate criterion values
|
|
36
|
+
return [
|
|
37
|
+
k if v is None or v == "" else "%s: %s" % (k, render_criterion(v))
|
|
38
|
+
for k, v in crit.items()
|
|
39
|
+
]
|
|
40
|
+
if t == "score":
|
|
41
|
+
return ["level %d: %s" % (i, render_criterion(c)) for i, c in enumerate(crit)]
|
|
42
|
+
crit = crit or {}
|
|
43
|
+
false_crit, true_crit = crit.get("false"), crit.get("true")
|
|
44
|
+
return [
|
|
45
|
+
"false: "
|
|
46
|
+
+ (
|
|
47
|
+
render_criterion(false_crit)
|
|
48
|
+
if false_crit not in (None, "")
|
|
49
|
+
else "no, the statement does not hold"
|
|
50
|
+
),
|
|
51
|
+
"true: "
|
|
52
|
+
+ (
|
|
53
|
+
render_criterion(true_crit)
|
|
54
|
+
if true_crit not in (None, "")
|
|
55
|
+
else "yes, the statement holds"
|
|
56
|
+
),
|
|
57
|
+
]
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def build_prefix(tok, q: Dict, head_max_len: int = 192, option_order=None):
|
|
61
|
+
"""Build the question-only prefix, before state tokens and final truncation."""
|
|
62
|
+
mask_tok = tok.mask_token
|
|
63
|
+
opts = render_options(q)
|
|
64
|
+
order = option_order if option_order is not None else list(range(len(opts)))
|
|
65
|
+
ins = str(q["ins"]).replace(mask_tok, " ")
|
|
66
|
+
head_ids = tok("%s question: %s" % (q["t"], ins), add_special_tokens=False)["input_ids"]
|
|
67
|
+
opt_ids = []
|
|
68
|
+
for i in order:
|
|
69
|
+
opt_ids.append(
|
|
70
|
+
[tok.mask_token_id]
|
|
71
|
+
+ tok(" " + opts[i].replace(mask_tok, " "), add_special_tokens=False)["input_ids"][:48]
|
|
72
|
+
)
|
|
73
|
+
opt_budget = head_max_len - sum(len(o) for o in opt_ids)
|
|
74
|
+
if opt_budget < 16:
|
|
75
|
+
per = max(4, (head_max_len - 16) // max(1, len(opt_ids)))
|
|
76
|
+
opt_ids = [o[:per] for o in opt_ids]
|
|
77
|
+
opt_budget = head_max_len - sum(len(o) for o in opt_ids)
|
|
78
|
+
head_ids = head_ids[: max(8, opt_budget)]
|
|
79
|
+
ids = [tok.cls_token_id] + head_ids + [tok.sep_token_id]
|
|
80
|
+
markers = []
|
|
81
|
+
for o in opt_ids:
|
|
82
|
+
markers.append(len(ids))
|
|
83
|
+
ids.extend(o)
|
|
84
|
+
ids.append(tok.sep_token_id)
|
|
85
|
+
return ids, markers
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def build_sequence(
|
|
89
|
+
tok,
|
|
90
|
+
state: Union[str, dict, list],
|
|
91
|
+
q: Dict,
|
|
92
|
+
max_len: int = 512,
|
|
93
|
+
head_max_len: int = 192,
|
|
94
|
+
option_order: Optional[List[int]] = None,
|
|
95
|
+
truncate_left: bool = False,
|
|
96
|
+
):
|
|
97
|
+
"""Format: [CLS] <type> instructions [SEP] [MASK] opt0 [MASK] opt1 ... [SEP] state [SEP]."""
|
|
98
|
+
ids, markers = build_prefix(tok, q, head_max_len, option_order)
|
|
99
|
+
room = max(0, max_len - len(ids) - 1)
|
|
100
|
+
st = tok(serialize_state(state).replace(tok.mask_token, " "), add_special_tokens=False)[
|
|
101
|
+
"input_ids"
|
|
102
|
+
]
|
|
103
|
+
st = st[-room:] if truncate_left else st[:room]
|
|
104
|
+
ids = ids + st + [tok.sep_token_id]
|
|
105
|
+
return ids[:max_len], [m for m in markers if m < max_len]
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def confidence_from_probs(p: np.ndarray, k: int) -> float:
|
|
109
|
+
"""Normalized Shannon entropy confidence: 1 - H(p) / log(k)."""
|
|
110
|
+
if k < 2:
|
|
111
|
+
return 1.0
|
|
112
|
+
p = p[:k]
|
|
113
|
+
ent = -(p * np.log(np.clip(p, 1e-12, 1.0))).sum()
|
|
114
|
+
return float(np.clip(1.0 - ent / math.log(k), 0.0, 1.0))
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
def temp_bucket(qtype: int, k: int) -> str:
|
|
118
|
+
size = "2" if k <= 2 else "3-5" if k <= 5 else "6-10" if k <= 10 else "11+"
|
|
119
|
+
return "%s:%s" % (QTYPE_NAMES[int(qtype)], size)
|