Xunzhuo commited on
Commit
2973ad4
·
verified ·
1 Parent(s): ffe291b

Decision 2.0 package DEV2.0-2B (release) from 0ca621e0ecbb20e5ae1802b9f71f8dea801fa693-src_training_decision2

Browse files
Files changed (3) hide show
  1. MODEL_MANIFEST.json +13 -13
  2. README.md +1 -0
  3. decision2/qwen.py +105 -20
MODEL_MANIFEST.json CHANGED
@@ -2,7 +2,7 @@
2
  "base": null,
3
  "builder": {
4
  "module_sha256": "bd498737bb797532fe74039088c45cfe2a75585d111034ee7d9bd04db838587c",
5
- "source_commit": "539d2769fe4d13ae888dcef4cda8695cb4bafc9e"
6
  },
7
  "calibration": null,
8
  "card": {
@@ -11,14 +11,14 @@
11
  "assets/jevarena-v3-rank.svg": "26a36aa0569d9176357be3438320145c07cedbcc280d65864828c435b77ea17b",
12
  "assets/jevbench-public231-rank.svg": "7ed4bf852ef67c95c91ad8a85ba7faba53bdde0115f6a30d492dad2a7fb7ed4e"
13
  },
14
- "readme_sha256": "3d679fe27250706a8f872910527bc6bd96fd660ce51c34bb1bd6502996acc61d"
15
  },
16
  "files_sha256": {
17
  "ATTRIBUTIONS.md": "69c8e5237dafeeaa1ec77c5a6636068bcb51c3054c07ce6688c7de691f0594e3",
18
  "LICENSE": "bbedc3fda3305820b977265f01b8619d87570a6739de3a5582c3464840f1e57a",
19
  "LICENSES/Qwen3.5-2B-LICENSE.txt": "bbedc3fda3305820b977265f01b8619d87570a6739de3a5582c3464840f1e57a",
20
  "NOTICE": "f82ed40c15223a2a5e93dbdd2931b48a86c2aab2c181838f3f6d7f2491a503ee",
21
- "README.md": "3d679fe27250706a8f872910527bc6bd96fd660ce51c34bb1bd6502996acc61d",
22
  "assets/DEV2.0-2B-owl-banner.png": "94008fafd629e5b8412e9c7b92d9dec07ff18a30e2cb973d4a2e43a8f927515a",
23
  "assets/jevarena-v3-model-task.svg": "5616473321481ad6afe45343adc863dfe4e97537b761833302c12bcdd75e102e",
24
  "assets/jevarena-v3-rank.svg": "26a36aa0569d9176357be3438320145c07cedbcc280d65864828c435b77ea17b",
@@ -40,7 +40,7 @@
40
  "decision2/_vendor/dev2model/lora.py": "7049e35e2a7bf50dd5888902cd0b7030d9bbd448d7d7aba8875ceb89aebff231",
41
  "decision2/_vendor/dev2model/source.py": "ef7b30171c4befb26163ea4d3d6e5add9dc602228a476d6f0ead9d40b450c7e3",
42
  "decision2/api.py": "3157c0509e8fb6c56038486f19b42ca09f44edcfce2ef32bc9ed5f4768829b0f",
43
- "decision2/qwen.py": "59eb1605ece2066db3061530c2b687a667719202ea0c61ef88f748a3293ac49d",
44
  "decision_config.json": "0b6c3429ee06032739d7386659c332ed7bb62d0b96ccddfcad4fad999e22cd4f",
45
  "decision_head.safetensors": "33b6541bb6636677eb91a4d8e06acd4db81152b11097088840796f11d49c707a",
46
  "evaluation/EVALUATION.md": "ecc747078712611e97ed56724e165a758bc3526eee5feb4a1f47c61b1dfedf32",
@@ -135,9 +135,9 @@
135
  "AutoModel": "modeling_decision2.Decision2Model"
136
  },
137
  "automap_source": {
138
- "commit": "8e808244072cef8735b8f3fdcfb8487f36605c5c",
139
- "content_manifest_sha256": "565576439331ef038c4b5c4794b867563bc6ad56964d19b0de948ed4f931f538",
140
- "tree": "31010c6ec5e465b2c8d8567e994c9f1dfb571177"
141
  },
142
  "custom_pipelines": {
143
  "decision": {
@@ -173,7 +173,7 @@
173
  },
174
  "repo_id": "llm-semantic-router/DEV2.0-2B",
175
  "runtime": {
176
- "equivalence": "decision2/qwen.py loads this full checkpoint with the vendored training/model sources whose SHA-256 equal the scored adapter sources (checked at build time) and applies the per-item batching, BF16-backbone / FP32-head execution, raw probabilities (temperature 1; no calibration file) and answer normalization of v2.dec.infer_dec, which for this non-residual checkpoint wraps the same DecisionModel without extra readouts. The scored checkpoint (073bd1f2) stored every tensor in FP32; this package (v2.release.bf16_copy, receipt a5229ef1) stores its 186 Linear projection matrices in BF16 exactly as BF16 autocast rounds them and every other tensor bit for bit in FP32. From this revision the runtime holds the backbone's BF16-exact Linear weights in BF16, the values BF16 autocast multiplies with, instead of FP32 copies cast before every matmul; it was checked with 0 answer changes on every scored prompt and on mlx-diag (2,275). From this revision the package also ships 🤗 Transformers remote code (configuration_decision2.py, modeling_decision2.py, pipeline_decision2.py; config.json gains model_type, auto_map and custom_pipelines): AutoModel with trust_remote_code loads the package through this runtime, which now also refuses Transformers' remote-code prompt while loading the tokenizer; it was checked with 0 answer changes against the native runtime on every scored prompt and on mlx-diag (2,275). Checked on one GPU of the scoring node against the T = 1 predictions derived exactly from the sealed CAL698 predictions of every scored prompt (typed-final 1,600, css15 6,547, public231 231) and of the mlx-diag diagnostic (2,275) by release.sh --parity, with the scored run's persisted Triton autotune cache.",
177
  "requirements": {
178
  "causal-conv1d": "1.7.0 (GPU convolution kernels)",
179
  "flash-linear-attention": "0.5.2 (GPU gated-delta kernels)",
@@ -185,9 +185,9 @@
185
  "triton": "3.7.1 (persisted autotune cache)"
186
  },
187
  "runtime_source": {
188
- "commit": "8e808244072cef8735b8f3fdcfb8487f36605c5c",
189
- "content_manifest_sha256": "565576439331ef038c4b5c4794b867563bc6ad56964d19b0de948ed4f931f538",
190
- "tree": "31010c6ec5e465b2c8d8567e994c9f1dfb571177"
191
  },
192
  "scored_runtime_check": {
193
  "checked": {
@@ -267,9 +267,9 @@
267
  },
268
  "decision2/qwen.py": {
269
  "rewritten": false,
270
- "sha256": "59eb1605ece2066db3061530c2b687a667719202ea0c61ef88f748a3293ac49d",
271
  "source": "v2/release/runtime/qwen.py",
272
- "source_sha256": "59eb1605ece2066db3061530c2b687a667719202ea0c61ef88f748a3293ac49d"
273
  }
274
  },
275
  "schema": "dev2-package-manifest/1",
 
2
  "base": null,
3
  "builder": {
4
  "module_sha256": "bd498737bb797532fe74039088c45cfe2a75585d111034ee7d9bd04db838587c",
5
+ "source_commit": "0ca621e0ecbb20e5ae1802b9f71f8dea801fa693"
6
  },
7
  "calibration": null,
8
  "card": {
 
11
  "assets/jevarena-v3-rank.svg": "26a36aa0569d9176357be3438320145c07cedbcc280d65864828c435b77ea17b",
12
  "assets/jevbench-public231-rank.svg": "7ed4bf852ef67c95c91ad8a85ba7faba53bdde0115f6a30d492dad2a7fb7ed4e"
13
  },
14
+ "readme_sha256": "acbd4d50b6452339339e19613f2fc34681fad549709d7f2cefe84a80835bfe32"
15
  },
16
  "files_sha256": {
17
  "ATTRIBUTIONS.md": "69c8e5237dafeeaa1ec77c5a6636068bcb51c3054c07ce6688c7de691f0594e3",
18
  "LICENSE": "bbedc3fda3305820b977265f01b8619d87570a6739de3a5582c3464840f1e57a",
19
  "LICENSES/Qwen3.5-2B-LICENSE.txt": "bbedc3fda3305820b977265f01b8619d87570a6739de3a5582c3464840f1e57a",
20
  "NOTICE": "f82ed40c15223a2a5e93dbdd2931b48a86c2aab2c181838f3f6d7f2491a503ee",
21
+ "README.md": "acbd4d50b6452339339e19613f2fc34681fad549709d7f2cefe84a80835bfe32",
22
  "assets/DEV2.0-2B-owl-banner.png": "94008fafd629e5b8412e9c7b92d9dec07ff18a30e2cb973d4a2e43a8f927515a",
23
  "assets/jevarena-v3-model-task.svg": "5616473321481ad6afe45343adc863dfe4e97537b761833302c12bcdd75e102e",
24
  "assets/jevarena-v3-rank.svg": "26a36aa0569d9176357be3438320145c07cedbcc280d65864828c435b77ea17b",
 
40
  "decision2/_vendor/dev2model/lora.py": "7049e35e2a7bf50dd5888902cd0b7030d9bbd448d7d7aba8875ceb89aebff231",
41
  "decision2/_vendor/dev2model/source.py": "ef7b30171c4befb26163ea4d3d6e5add9dc602228a476d6f0ead9d40b450c7e3",
42
  "decision2/api.py": "3157c0509e8fb6c56038486f19b42ca09f44edcfce2ef32bc9ed5f4768829b0f",
43
+ "decision2/qwen.py": "97ba565f95d5a5ac9af6632e7397c05cfe9ab47fde68b56df19aabfa2d8fb187",
44
  "decision_config.json": "0b6c3429ee06032739d7386659c332ed7bb62d0b96ccddfcad4fad999e22cd4f",
45
  "decision_head.safetensors": "33b6541bb6636677eb91a4d8e06acd4db81152b11097088840796f11d49c707a",
46
  "evaluation/EVALUATION.md": "ecc747078712611e97ed56724e165a758bc3526eee5feb4a1f47c61b1dfedf32",
 
135
  "AutoModel": "modeling_decision2.Decision2Model"
136
  },
137
  "automap_source": {
138
+ "commit": "99432d1a7da5adbc70212df78ae7ebf7e50b41e4",
139
+ "content_manifest_sha256": "3bede2ca1fc0b805c167bf5ae6b2bd824a3b64d1b6ae98478d04add1cc1444ff",
140
+ "tree": "fbf37474b1ce2efd228b7f50803848c6f60ebbd2"
141
  },
142
  "custom_pipelines": {
143
  "decision": {
 
173
  },
174
  "repo_id": "llm-semantic-router/DEV2.0-2B",
175
  "runtime": {
176
+ "equivalence": "decision2/qwen.py loads this full checkpoint with the vendored training/model sources whose SHA-256 equal the scored adapter sources (checked at build time) and applies the per-item batching, BF16-backbone / FP32-head execution, raw probabilities (temperature 1; no calibration file) and answer normalization of v2.dec.infer_dec, which for this non-residual checkpoint wraps the same DecisionModel without extra readouts. The scored checkpoint (073bd1f2) stored every tensor in FP32; this package (v2.release.bf16_copy, receipt a5229ef1) stores its 186 Linear projection matrices in BF16 exactly as BF16 autocast rounds them and every other tensor bit for bit in FP32. From this revision the runtime holds the backbone's BF16-exact Linear weights in BF16, the values BF16 autocast multiplies with, instead of FP32 copies cast before every matmul; it was checked with 0 answer changes on every scored prompt and on mlx-diag (2,275). From this revision the package also ships 🤗 Transformers remote code (configuration_decision2.py, modeling_decision2.py, pipeline_decision2.py; config.json gains model_type, auto_map and custom_pipelines): AutoModel with trust_remote_code loads the package through this runtime, which now also refuses Transformers' remote-code prompt while loading the tokenizer; it was checked with 0 answer changes against the native runtime on every scored prompt and on mlx-diag (2,275). From this revision the runtime runs a request whose padded question batch would put more than 2**30 elements in a gated-delta q / k / v tensor (the FLA kernels' 32-bit offsets) as several GPU-sized batches, longest questions first; requests within that budget, every scored prompt among them, keep the single-batch path, and it was checked with 0 answer changes on every scored prompt and on mlx-diag (2,275). Checked on one GPU of the scoring node against the T = 1 predictions derived exactly from the sealed CAL698 predictions of every scored prompt (typed-final 1,600, css15 6,547, public231 231) and of the mlx-diag diagnostic (2,275) by release.sh --parity, with the scored run's persisted Triton autotune cache.",
177
  "requirements": {
178
  "causal-conv1d": "1.7.0 (GPU convolution kernels)",
179
  "flash-linear-attention": "0.5.2 (GPU gated-delta kernels)",
 
185
  "triton": "3.7.1 (persisted autotune cache)"
186
  },
187
  "runtime_source": {
188
+ "commit": "99432d1a7da5adbc70212df78ae7ebf7e50b41e4",
189
+ "content_manifest_sha256": "3bede2ca1fc0b805c167bf5ae6b2bd824a3b64d1b6ae98478d04add1cc1444ff",
190
+ "tree": "fbf37474b1ce2efd228b7f50803848c6f60ebbd2"
191
  },
192
  "scored_runtime_check": {
193
  "checked": {
 
267
  },
268
  "decision2/qwen.py": {
269
  "rewritten": false,
270
+ "sha256": "97ba565f95d5a5ac9af6632e7397c05cfe9ab47fde68b56df19aabfa2d8fb187",
271
  "source": "v2/release/runtime/qwen.py",
272
+ "source_sha256": "97ba565f95d5a5ac9af6632e7397c05cfe9ab47fde68b56df19aabfa2d8fb187"
273
  }
274
  },
275
  "schema": "dev2-package-manifest/1",
README.md CHANGED
@@ -164,6 +164,7 @@ print(json.dumps(result["answers"], indent=2))
164
  - **Weights:** the uniform average of three full fine-tunes of Decision 1.0 Sol with the same recipe and start; the seeds differ only in data order, and each seed's checkpoint was chosen on a held-out selection set (updates 366, 741 and 638).
165
  - **Comparator licences:** Decider 2B's model card declares Apache-2.0 and This-That 1.2's declares MIT (it is adapted from Decider 2B); neither repository has a LICENSE file. Bosun v3.1 1.7B ships Apache-2.0 LICENSE and NOTICE files.
166
  - **Runtime update:** BF16-resident weights; answers unchanged; latency p50 24.1 → 23.5 ms, p95 24.8 → 23.7 ms; peak GPU memory 7.1 → 4.6 GiB (400 single requests, 346 input tokens on average, one AMD MI325X GPU).
 
167
 
168
  ### Training
169
 
 
164
  - **Weights:** the uniform average of three full fine-tunes of Decision 1.0 Sol with the same recipe and start; the seeds differ only in data order, and each seed's checkpoint was chosen on a held-out selection set (updates 366, 741 and 638).
165
  - **Comparator licences:** Decider 2B's model card declares Apache-2.0 and This-That 1.2's declares MIT (it is adapted from Decider 2B); neither repository has a LICENSE file. Bosun v3.1 1.7B ships Apache-2.0 LICENSE and NOTICE files.
166
  - **Runtime update:** BF16-resident weights; answers unchanged; latency p50 24.1 → 23.5 ms, p95 24.8 → 23.7 ms; peak GPU memory 7.1 → 4.6 GiB (400 single requests, 346 input tokens on average, one AMD MI325X GPU).
167
+ - **Runtime update:** very long multi-question requests are processed in GPU-sized batches (previously they could fail on extremely long inputs); answers unchanged.
168
 
169
  ### Training
170
 
decision2/qwen.py CHANGED
@@ -7,10 +7,15 @@ records an import-only rewrite. A base-bound adapter's source files are pinned
7
  by repository revision and SHA-256; a supplied or downloaded copy is verified
8
  before use. GPU inference runs the backbone under BF16 autocast with its
9
  BF16-exact Linear weights held in BF16 and every other tensor, the head
10
- included, in FP32; CPU uses FP32.
 
 
 
11
  A package whose manifest names a weight ``storage`` codec (bf16z) is first restored
12
  to exact safetensors files in a cache directory; the restored checkpoint must
13
- reproduce the scored identity.
 
 
14
  """
15
 
16
  from __future__ import annotations
@@ -98,6 +103,53 @@ def keep_linear_bf16(module: Any, torch: Any) -> dict[str, int]:
98
  return counts
99
 
100
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
101
  class QwenDecision:
102
  def __init__(
103
  self,
@@ -109,6 +161,8 @@ class QwenDecision:
109
  torch: Any,
110
  score_bias: dict[int, list[float]] | None = None,
111
  residency: dict[str, int] | None = None,
 
 
112
  ):
113
  self.model = model
114
  self.tokenizer = tokenizer
@@ -118,6 +172,8 @@ class QwenDecision:
118
  self.torch = torch
119
  self.score_bias = score_bias
120
  self.residency = residency
 
 
121
 
122
  @classmethod
123
  def load(
@@ -150,7 +206,12 @@ class QwenDecision:
150
  (model_root / "decision_config.json").read_text(encoding="utf-8")
151
  )
152
  residual = metadata.get("dec_residual") is not None
153
- if residual:
 
 
 
 
 
154
  from ._vendor.dev2model.dec_model import dec_fingerprint
155
 
156
  identity = dec_fingerprint(model_root, source)
@@ -176,7 +237,16 @@ class QwenDecision:
176
  not torch.cuda.is_available() or not torch.cuda.is_bf16_supported()
177
  ):
178
  raise RuntimeError("A CUDA/ROCm BF16 GPU is required for GPU inference")
179
- if residual:
 
 
 
 
 
 
 
 
 
180
  from ._vendor.dev2model.dec_model import load_dec_checkpoint
181
 
182
  model, tokenizer = load_dec_checkpoint(model_root, source)
@@ -191,6 +261,11 @@ class QwenDecision:
191
  model = model.to(target).eval()
192
  if tokenizer.pad_token_id is None and tokenizer.eos_token_id is None:
193
  raise ValueError("Tokenizer needs a pad or EOS token")
 
 
 
 
 
194
  return cls(
195
  model,
196
  tokenizer,
@@ -200,6 +275,8 @@ class QwenDecision:
200
  torch,
201
  score_bias,
202
  residency,
 
 
203
  )
204
 
205
  def parameter_count(self) -> int:
@@ -229,7 +306,7 @@ class QwenDecision:
229
  for qid, question in questions.items():
230
  try:
231
  row = question_to_row(item, qid, question)
232
- encoded = encode(row, self.tokenizer, self.cap)
233
  except ValueError as exc:
234
  reason = (
235
  "max_length_exceeded"
@@ -250,21 +327,29 @@ class QwenDecision:
250
  pad_id = self.tokenizer.pad_token_id
251
  if pad_id is None:
252
  pad_id = self.tokenizer.eos_token_id
253
- batch = {
254
- key: value.to(self.device) if self.torch.is_tensor(value) else value
255
- for key, value in collate(
256
- [encoded for _, _, encoded in jobs], pad_id
257
- ).items()
258
- }
259
- autocast = (
260
- self.torch.autocast(device_type="cuda", dtype=self.torch.bfloat16)
261
- if self.device.type == "cuda"
262
- else nullcontext()
263
- )
264
- with self.torch.inference_mode(), autocast:
265
- logits = self.model(**batch)
266
- if len(logits) != len(jobs):
267
- raise RuntimeError("Model returned the wrong number of question answers")
 
 
 
 
 
 
 
 
268
  if self.score_bias is not None:
269
  from ._vendor.dev2model.score_bias import apply as apply_score_bias
270
  for (qid, row, encoded), values in zip(jobs, logits):
 
7
  by repository revision and SHA-256; a supplied or downloaded copy is verified
8
  before use. GPU inference runs the backbone under BF16 autocast with its
9
  BF16-exact Linear weights held in BF16 and every other tensor, the head
10
+ included, in FP32; CPU uses FP32. A request's questions run as one padded
11
+ batch unless, on a GPU, that batch would put more than 2**30 elements in a
12
+ gated-delta q / k / v tensor (``forward_token_budget``); then they run as
13
+ several batches.
14
  A package whose manifest names a weight ``storage`` codec (bf16z) is first restored
15
  to exact safetensors files in a cache directory; the restored checkpoint must
16
+ reproduce the scored identity. A checkpoint whose ``decision_config.json`` declares
17
+ ``readout: label_token`` is read through its tied LM head's label-token logits at
18
+ the answer cue (vendored ``label_token.py``) instead of a candidate head.
19
  """
20
 
21
  from __future__ import annotations
 
103
  return counts
104
 
105
 
106
+ INT32_MAX = 2**31 - 1
107
+
108
+
109
+ def forward_token_budget(config: Any) -> int | None:
110
+ """Most padded tokens one GPU forward may hold; None when nothing limits it.
111
+
112
+ The FLA gated-delta kernels of Qwen3.5-family backbones compute some
113
+ element offsets of their [batch, tokens, value heads, head dim] q / k / v
114
+ tensors in 32-bit integers (the L2-norm forward and the fused KKT-solve
115
+ kernel among them). A forward whose tensors exceed 2**31 - 1 elements reads
116
+ and writes the wrong memory for the later batch rows: wrong answers,
117
+ non-finite logits or a GPU memory fault. On DEV2.0-27B, forwards with
118
+ tensors between about 2**30 and 2**31 elements also hung or crashed the
119
+ process, so the budget keeps these tensors within 2**30 elements.
120
+ """
121
+ config = getattr(config, "text_config", None) or config
122
+ heads = getattr(config, "linear_num_value_heads", None)
123
+ if not heads:
124
+ return None
125
+ width = heads * max(config.linear_key_head_dim, config.linear_value_head_dim)
126
+ return (INT32_MAX // 2) // width
127
+
128
+
129
+ def micro_batches(lengths: list[int], budget: int | None) -> list[list[int]]:
130
+ """Indices of each forward: one batch when its padded size fits the budget.
131
+
132
+ Otherwise the questions go longest first into batches whose padded size
133
+ (``collate`` pads to the longest, rounded up to 8) stays within the budget,
134
+ each batch listing its indices in request order.
135
+ """
136
+
137
+ def padded(length: int) -> int:
138
+ return -(-length // 8) * 8
139
+
140
+ if budget is None or padded(max(lengths)) * len(lengths) <= budget:
141
+ return [list(range(len(lengths)))]
142
+ order = sorted(range(len(lengths)), key=lambda i: (-lengths[i], i))
143
+ if padded(lengths[order[0]]) > budget:
144
+ raise ValueError("A single question exceeds the forward token budget")
145
+ groups: list[list[int]] = []
146
+ while order:
147
+ rows = budget // padded(lengths[order[0]])
148
+ groups.append(sorted(order[:rows]))
149
+ order = order[rows:]
150
+ return groups
151
+
152
+
153
  class QwenDecision:
154
  def __init__(
155
  self,
 
161
  torch: Any,
162
  score_bias: dict[int, list[float]] | None = None,
163
  residency: dict[str, int] | None = None,
164
+ encode_fn: Any = None,
165
+ batch_tokens: int | None = None,
166
  ):
167
  self.model = model
168
  self.tokenizer = tokenizer
 
172
  self.torch = torch
173
  self.score_bias = score_bias
174
  self.residency = residency
175
+ self.encode_fn = encode_fn
176
+ self.batch_tokens = batch_tokens
177
 
178
  @classmethod
179
  def load(
 
206
  (model_root / "decision_config.json").read_text(encoding="utf-8")
207
  )
208
  residual = metadata.get("dec_residual") is not None
209
+ label = metadata.get("readout") == "label_token"
210
+ if label:
211
+ from ._vendor.dev2model.label_token import label_fingerprint
212
+
213
+ identity = label_fingerprint(model_root, source)
214
+ elif residual:
215
  from ._vendor.dev2model.dec_model import dec_fingerprint
216
 
217
  identity = dec_fingerprint(model_root, source)
 
237
  not torch.cuda.is_available() or not torch.cuda.is_bf16_supported()
238
  ):
239
  raise RuntimeError("A CUDA/ROCm BF16 GPU is required for GPU inference")
240
+ encode_fn = None
241
+ if label:
242
+ from ._vendor.dev2model.label_token import (
243
+ encode_label,
244
+ load_label_checkpoint,
245
+ )
246
+
247
+ model, tokenizer = load_label_checkpoint(model_root, source)
248
+ encode_fn = encode_label
249
+ elif residual:
250
  from ._vendor.dev2model.dec_model import load_dec_checkpoint
251
 
252
  model, tokenizer = load_dec_checkpoint(model_root, source)
 
261
  model = model.to(target).eval()
262
  if tokenizer.pad_token_id is None and tokenizer.eos_token_id is None:
263
  raise ValueError("Tokenizer needs a pad or EOS token")
264
+ batch_tokens = (
265
+ forward_token_budget(model.backbone.config)
266
+ if target.type == "cuda"
267
+ else None
268
+ )
269
  return cls(
270
  model,
271
  tokenizer,
 
275
  torch,
276
  score_bias,
277
  residency,
278
+ encode_fn,
279
+ batch_tokens,
280
  )
281
 
282
  def parameter_count(self) -> int:
 
306
  for qid, question in questions.items():
307
  try:
308
  row = question_to_row(item, qid, question)
309
+ encoded = (self.encode_fn or encode)(row, self.tokenizer, self.cap)
310
  except ValueError as exc:
311
  reason = (
312
  "max_length_exceeded"
 
327
  pad_id = self.tokenizer.pad_token_id
328
  if pad_id is None:
329
  pad_id = self.tokenizer.eos_token_id
330
+ logits: list[Any] = [None] * len(jobs)
331
+ for group in micro_batches(
332
+ [len(encoded["ids"]) for _, _, encoded in jobs], self.batch_tokens
333
+ ):
334
+ batch = {
335
+ key: value.to(self.device) if self.torch.is_tensor(value) else value
336
+ for key, value in collate(
337
+ [jobs[index][2] for index in group], pad_id
338
+ ).items()
339
+ }
340
+ autocast = (
341
+ self.torch.autocast(device_type="cuda", dtype=self.torch.bfloat16)
342
+ if self.device.type == "cuda"
343
+ else nullcontext()
344
+ )
345
+ with self.torch.inference_mode(), autocast:
346
+ output = self.model(**batch)
347
+ if len(output) != len(group):
348
+ raise RuntimeError(
349
+ "Model returned the wrong number of question answers"
350
+ )
351
+ for index, values in zip(group, output):
352
+ logits[index] = values
353
  if self.score_bias is not None:
354
  from ._vendor.dev2model.score_bias import apply as apply_score_bias
355
  for (qid, row, encoded), values in zip(jobs, logits):