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.
@@ -0,0 +1,6 @@
1
+ """Core ML runtime for Laya. PyTorch is only needed when converting weights."""
2
+
3
+ from .agent import Agent, load
4
+
5
+ __all__ = ["Agent", "load"]
6
+ __version__ = "0.1.0"
@@ -0,0 +1,3 @@
1
+ from .cli import main
2
+
3
+ main()
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
@@ -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)