hotchpotch commited on
Commit
12f38b0
·
verified ·
1 Parent(s): 66bc467

Document inference interface, batching and performance tuning

Browse files
Files changed (1) hide show
  1. inference_v0.py +294 -18
inference_v0.py CHANGED
@@ -1,4 +1,134 @@
1
- """Self-contained Bekko v0 inference; generated by bekko_system_one.export_v0."""
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2
  from __future__ import annotations
3
 
4
  from dataclasses import dataclass
@@ -328,10 +458,18 @@ from transformers import AutoConfig, AutoModel, AutoTokenizer, PreTrainedTokeniz
328
  from transformers.models.modernbert.modeling_modernbert import apply_rotary_pos_emb
329
 
330
  def _text(value):
 
331
  return value if isinstance(value, str) else json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(',', ':'))
332
 
333
  def input_groups(case_input, *, prefix_layout='instruction_state'):
334
- """Render native inference input; labels and row metadata are never accepted."""
 
 
 
 
 
 
 
335
  if not isinstance(case_input, dict) or set(case_input) != {'state_json', 'decisions'}:
336
  raise ValueError('Pass the input object only: state_json and decisions')
337
  state = case_input['state_json']
@@ -379,10 +517,12 @@ class _SDPAPrefix(nn.Module):
379
  """Independent prefix self-attention and suffix-aligned document attention."""
380
 
381
  def __init__(self, backbone):
 
382
  super().__init__()
383
  self.backbone = backbone
384
 
385
  def qkv(self, layer, hidden, positions):
 
386
  attn = layer.attn
387
  q, k, v = attn.Wqkv(layer.attn_norm(hidden)).view(*hidden.shape[:2], 3, -1, attn.head_dim).unbind(2)
388
  cos, sin = self.backbone.rotary_emb(hidden, positions, layer.attention_type)
@@ -391,6 +531,7 @@ class _SDPAPrefix(nn.Module):
391
 
392
  @staticmethod
393
  def attend(q, k, v, qmask, kmask, window):
 
394
  qp = qmask.long().cumsum(1) - 1 + (kmask.sum(1) - qmask.sum(1))[:, None]
395
  kp = kmask.long().cumsum(1) - 1
396
  allowed = kmask[:, None, :].expand(-1, q.shape[1], -1)
@@ -401,10 +542,12 @@ class _SDPAPrefix(nn.Module):
401
 
402
  @staticmethod
403
  def update(layer, hidden, attended):
 
404
  hidden = hidden + layer.attn.out_drop(layer.attn.Wo(attended.flatten(2)))
405
  return hidden + layer.mlp(layer.mlp_norm(hidden))
406
 
407
  def forward(self, prefix_ids, prefix_mask, doc_ids, doc_mask, owners):
 
408
  prefix = self.backbone.embeddings(prefix_ids)
409
  hidden = self.backbone.embeddings(doc_ids)
410
  pp = (prefix_mask.long().cumsum(1) - 1).clamp_min(0)
@@ -428,6 +571,7 @@ class BekkoInference(InputModule):
428
  budget_policy = 'adaptive-v1'
429
 
430
  def __init__(self, backbone, tokenizer, *, tasks, query_length=None, document_length=None, context_length=None, query_truncation='balanced', task_tokens=None, choice_interaction=None, prefix_layout='instruction_state'):
 
431
  super().__init__()
432
  if backbone.config.model_type != 'modernbert':
433
  raise ValueError('Expected ModernBERT-compatible weights')
@@ -460,27 +604,34 @@ class BekkoInference(InputModule):
460
  self.settings = dict(tasks=list(tasks), query_length=query_length, document_length=document_length, context_length=capacity, query_truncation=query_truncation, task_tokens=self.task_tokens, choice_interaction=choice_interaction, prefix_layout=prefix_layout)
461
 
462
  def compile_inference(self, *, mode='default', dynamic=True, backend='inductor'):
463
- """Compile tensor execution lazily; first calls include compilation cost.
464
 
465
- Tokenization/rendering stay eager. No silent eager fallback is enabled.
466
- Changing shapes can trigger new compilations despite dynamic=True.
 
 
 
467
  """
468
  self.eval().requires_grad_(False)
469
  self._compiled_forward = torch.compile(self.encoder.forward, mode=mode, dynamic=dynamic, backend=backend)
470
  return self
471
 
472
  def disable_compile(self):
 
473
  self._compiled_forward = None
474
 
475
  def preprocess(self, inputs, prompt=None, **kwargs):
 
476
  raise ValueError('Typed decisions require candidate groups; use model.predict(input_object)')
477
 
478
  @staticmethod
479
  def _validate_limits(query_length, document_length, capacity):
 
480
  if not 3 <= query_length <= capacity - 2 or not 2 <= document_length <= capacity - 3:
481
  raise ValueError('Invalid query/document limits for this backbone')
482
 
483
  def _limits(self, query_length, document_length, context_length=None):
 
484
  capacity = self.context_length if context_length is None else context_length
485
  if not 5 <= capacity <= self.encoder.backbone.config.max_position_embeddings:
486
  raise ValueError('Context length exceeds the backbone positional capacity')
@@ -490,6 +641,7 @@ class BekkoInference(InputModule):
490
  return (q, d, capacity)
491
 
492
  def _documents(self, documents, tasks, limit):
 
493
  if len(tasks) != len(documents):
494
  raise ValueError('Tasks and candidates must align')
495
  if not documents:
@@ -498,6 +650,7 @@ class BekkoInference(InputModule):
498
  return [[self.task_token_ids[t], *d[:limit - 2], self.tokenizer.sep_token_id] if t in self.task_token_ids else [*d, self.tokenizer.sep_token_id] for t, d in zip(tasks, encoded, strict=True)]
499
 
500
  def _queries(self, queries, parts, limits):
 
501
  if not queries:
502
  return []
503
  if self.query_truncation == 'balanced':
@@ -508,7 +661,13 @@ class BekkoInference(InputModule):
508
  return [[self.tokenizer.cls_token_id, *q[:limit - 2], self.tokenizer.sep_token_id] for q, limit in zip(encoded, limits, strict=True)]
509
 
510
  def tokenize_branches(self, queries, documents, document_tasks=None, *, query_parts=None, query_length=None, document_length=None, context_length=None):
511
- """Allocate shared capacity conservatively when no decision ownership is supplied."""
 
 
 
 
 
 
512
  qlimit, dlimit, capacity = self._limits(query_length, document_length, context_length)
513
  dids = self._documents(documents, document_tasks or ['reranker'] * len(documents), dlimit)
514
  qids = self._queries(queries, query_parts, [qlimit] * len(queries))
@@ -546,6 +705,7 @@ class BekkoInference(InputModule):
546
  return [PreparedGroup(tuple(qmap[key]), g.task, qmap[key], [dmap[g.task, d][:limit - 1] + [dmap[g.task, d][-1]] if len(dmap[g.task, d]) > limit else dmap[g.task, d] for d in g.candidates], g.target, g.metadata) for g, key, limit in zip(groups, keys, document_limits, strict=True)]
547
 
548
  def collate_tokens(self, queries, documents, owners):
 
549
 
550
  def pad(rows):
551
  ids = torch.nn.utils.rnn.pad_sequence([torch.tensor(r, dtype=torch.long) for r in rows], batch_first=True, padding_value=self.tokenizer.pad_token_id)
@@ -555,6 +715,15 @@ class BekkoInference(InputModule):
555
  return dict(prefix_ids=pi, prefix_mask=pm, doc_ids=di, doc_mask=dm, owners=torch.tensor(owners, dtype=torch.long))
556
 
557
  def forward(self, features, **kwargs):
 
 
 
 
 
 
 
 
 
558
  device = next(self.parameters()).device
559
  with torch.autocast(device.type, dtype=torch.bfloat16, enabled=device.type == 'cuda'):
560
  execute = self._compiled_forward or self.encoder
@@ -583,11 +752,23 @@ class BekkoInference(InputModule):
583
 
584
  @torch.inference_mode()
585
  def predict_groups(self, groups, *, batch_size=128, token_budget=64000, query_length=None, document_length=None, context_length=None, show_progress_bar=True):
586
- """Score bounded windows, bucket by length, and restore original group order.
587
-
588
- batch_size bounds the number of decisions tokenized at once. The token
589
- budget estimates padded prefix-plus-document work; an oversized decision
590
- runs alone. Candidate groups are never split.
 
 
 
 
 
 
 
 
 
 
 
 
591
  """
592
  if batch_size < 1 or token_budget < 1:
593
  raise ValueError('batch_size and token_budget must be positive')
@@ -631,16 +812,42 @@ class BekkoInference(InputModule):
631
 
632
  @staticmethod
633
  def _interpret(group, probabilities):
 
634
  assert group.metadata is not None and group.metadata.candidate_ids is not None
635
  if group.metadata.kind == 'ranking':
636
  return {'probabilities': dict(zip(group.metadata.candidate_ids, probabilities.tolist(), strict=True)), 'order': [group.metadata.candidate_ids[i] for i in probabilities.argsort(descending=True, stable=True).tolist()]}
637
  return asdict(interpret_prediction(group, probabilities))
638
 
639
  def predict(self, inputs, *, batch_size=128, token_budget=64000, query_length=None, document_length=None, context_length=None, prefix_layout=None, show_progress_bar=True):
640
- """Accept one input dict or a list; return matching typed results in order.
641
-
642
- batch_size bounds both cases rendered and decisions tokenized per window.
643
- Length overrides apply only to this call. Progress counts completed cases.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
644
  """
645
  single = isinstance(inputs, dict)
646
  cases = [inputs] if single else list(inputs)
@@ -664,12 +871,15 @@ class BekkoInference(InputModule):
664
  return results[0] if single else results
665
 
666
  def get_sentence_embedding_dimension(self):
 
667
  return 1
668
 
669
  def get_config_dict(self):
 
670
  return self.settings
671
 
672
  def save(self, output_path, *args, **kwargs):
 
673
  path = Path(output_path)
674
  path.mkdir(parents=True, exist_ok=True)
675
  self.save_config(str(path))
@@ -681,6 +891,13 @@ class BekkoInference(InputModule):
681
 
682
  @classmethod
683
  def load(cls, model_name_or_path, subfolder='', token=None, cache_folder=None, revision=None, local_files_only=False, init_defaults=None, **kwargs):
 
 
 
 
 
 
 
684
  hub: dict[str, Any] = dict(subfolder=subfolder, token=token, cache_folder=cache_folder, revision=revision, local_files_only=local_files_only)
685
  settings = cls.load_config(model_name_or_path, **hub)
686
  config_path = cls.load_file_path(model_name_or_path, 'backbone_config.json', **hub)
@@ -704,20 +921,73 @@ class BekkoSentenceTransformer(SentenceTransformer):
704
  """
705
 
706
  def __init__(self, *args, **kwargs):
 
707
  super().__init__(*args, **kwargs)
708
  if len(self) != 1 or getattr(self[0], 'config_file_name', None) != 'inference_config.json':
709
  raise ValueError('BekkoSentenceTransformer requires an exported v0 checkpoint')
710
 
711
  def predict(self, inputs, *, batch_size=128, token_budget=64000, query_length=None, document_length=None, context_length=None, prefix_layout=None, show_progress_bar=True):
712
- """Return typed results for one native input dict or a list in input order."""
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
713
  return cast(Any, self[0]).predict(inputs, batch_size=batch_size, token_budget=token_budget, query_length=query_length, document_length=document_length, context_length=context_length, prefix_layout=prefix_layout, show_progress_bar=show_progress_bar)
714
 
715
  def predict_groups(self, groups, *, batch_size=128, token_budget=64000, query_length=None, document_length=None, context_length=None, show_progress_bar=True):
716
- """Return probability tensors for already-rendered decision groups."""
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
717
  return cast(Any, self[0]).predict_groups(groups, batch_size=batch_size, token_budget=token_budget, query_length=query_length, document_length=document_length, context_length=context_length, show_progress_bar=show_progress_bar)
718
 
719
  def compile_inference(self, *, mode='default', dynamic=True, backend='inductor'):
720
- """Compile tensor execution lazily and return this model."""
 
 
 
 
 
 
721
  cast(Any, self[0]).compile_inference(mode=mode, dynamic=dynamic, backend=backend)
722
  return self
723
 
@@ -727,6 +997,12 @@ class BekkoSentenceTransformer(SentenceTransformer):
727
  return self
728
 
729
  def main():
 
 
 
 
 
 
730
  parser = argparse.ArgumentParser(description=__doc__)
731
  parser.add_argument('--model', required=True)
732
  parser.add_argument('--input', required=True, type=Path, help='JSON file containing only native inference input')
 
1
+ """Standalone Bekko v0 typed-decision inference and remote-code interface.
2
+
3
+ Install the exported requirements.txt (and a CUDA-compatible PyTorch build for
4
+ GPU use). The exported file includes its helpers and needs no Bekko training
5
+ package, datasets, PEFT, W&B, or external FlashAttention extension.
6
+
7
+ Load once, reuse for many requests
8
+ ---------------------------------
9
+ From a downloaded export, import BekkoSentenceTransformer from inference_v0.
10
+ To load the class directly from a Hugging Face model repository::
11
+
12
+ from transformers.dynamic_module_utils import get_class_from_dynamic_module
13
+
14
+ repo = "YOUR_ORG/YOUR_MODEL"
15
+ revision = "FULL_COMMIT_HASH" # Pin the same revision for code and weights.
16
+ Model = get_class_from_dynamic_module(
17
+ "inference_v0.BekkoSentenceTransformer", repo,
18
+ revision=revision, token=True,
19
+ )
20
+ model = Model(repo, revision=revision, token=True,
21
+ trust_remote_code=True, device="cuda") # Or device="cpu".
22
+
23
+ Authenticate with Hugging Face before accessing a private model. token=True
24
+ uses your saved token. Plain SentenceTransformer(..., trust_remote_code=True)
25
+ also loads the export, but its typed entry point is model[0].predict(...).
26
+ BekkoSentenceTransformer exposes predict() directly. Raw-text encode() is not
27
+ a typed-decision API.
28
+
29
+ Input and output interface
30
+ --------------------------
31
+ Pass only a native input dict with exactly state_json and decisions; do not
32
+ pass a dataset row, targets, labels, or provenance. Fields ending in _json are
33
+ JSON-encoded strings, including when they contain a plain string. For example::
34
+
35
+ import json
36
+ request = {
37
+ "state_json": json.dumps({"message": "Please refund a duplicate charge."}),
38
+ "decisions": [{
39
+ "id": "department", "kind": "judgment", "type": "choice",
40
+ "instructions_json": json.dumps("Which department should respond?"),
41
+ "system_prompt": "",
42
+ "criteria": [
43
+ {"id": "billing", "description_json": json.dumps("Payments and refunds"),
44
+ "value": None},
45
+ {"id": "technical", "description_json": json.dumps("Technical failures"),
46
+ "value": None},
47
+ ],
48
+ "documents": [], "scoring": None,
49
+ }],
50
+ }
51
+ result = model.predict(request, show_progress_bar=False)
52
+ selected = result["department"]["selected_id"]
53
+
54
+ Decision IDs must be nonempty and unique within each request; they may repeat
55
+ across requests. Candidate IDs must be nonempty and distinct within a decision.
56
+ Judgments use kind="judgment", criteria, empty documents, and scoring=None:
57
+ * choice: returns selected_id and probabilities keyed by candidate ID.
58
+ * noul: use exactly true/false or yes/no IDs and authored descriptions of both
59
+ meanings; returns probability_yes and probabilities.
60
+ * score: supply numeric value for each criterion (at least two distinct values);
61
+ returns score (expected value on that scale), normalized_score in [0, 1],
62
+ probabilities, and values. Values are never inferred from descriptions.
63
+ Relative ranking uses kind="ranking", type=None, scoring="relative", empty
64
+ criteria, and documents containing id and content_json. It returns probabilities
65
+ and order (document IDs sorted by descending probability).
66
+
67
+ Choose the batching interface
68
+ -----------------------------
69
+ Use predict(request) for one request: it returns a dict keyed by decision ID.
70
+ Use predict(requests) for a list of requests: it returns a list of result dicts
71
+ in input order. This is the usual throughput interface, including mixed tasks::
72
+
73
+ results = model.predict(
74
+ [request, request], batch_size=128, token_budget=64000,
75
+ show_progress_bar=False,
76
+ )
77
+
78
+ Prefer this list call to a Python loop of single-request calls. batch_size limits
79
+ both the requests rendered per window and the decisions tokenized per window;
80
+ it is not a fixed candidate count or GPU microbatch size. token_budget estimates
81
+ candidate_count * (max_query_tokens + max_document_tokens) for each microbatch.
82
+ An oversized decision runs alone, so this is not a strict memory limit. Candidate
83
+ groups are never split. Length bucketing reduces padding; original request and
84
+ candidate order is restored. Tokenization deduplicates text within each window,
85
+ and encoder prefixes are shared within each microbatch, not cached across calls.
86
+
87
+ Use predict_groups(groups) only if you already have rendered Group objects,
88
+ for example from input_groups(request). It returns one CPU FP32 probability
89
+ tensor per decision, in group/candidate order, without typed interpretation.
90
+ It is not required for ordinary list batching. Empty request lists return [];
91
+ a request with no decisions returns {}. Inputs are materialized, not streamed;
92
+ chunk very large datasets into lists in the caller. Progress defaults to stderr;
93
+ predict counts requests, predict_groups counts decisions.
94
+
95
+ Fast execution and memory tuning
96
+ --------------------------------
97
+ Reuse a loaded model, batch requests, and select device="cuda" when available.
98
+ CUDA encoder execution uses BF16 autocast and PyTorch SDPA; CPU uses FP32.
99
+ Heads and per-decision softmax use FP32. Small CPU/GPU differences are expected.
100
+ Tune batch_size for the tokenization/sorting window and token_budget for GPU
101
+ work per microbatch; reduce the latter when memory is tight. A single oversized
102
+ decision still runs alone and may require shorter inputs or fewer candidates.
103
+ For repeated workloads, optionally enable compilation before warmup::
104
+
105
+ model.compile_inference() # Lazy encoder-only torch.compile, default Inductor.
106
+ model.predict([request, request], show_progress_bar=False) # Warmup.
107
+ results = model.predict([request, request], show_progress_bar=False)
108
+ model.disable_compile() # Return to eager encoder execution.
109
+
110
+ Compilation leaves rendering, tokenization, heads, packing, and output processing
111
+ eager. First calls include compilation cost; new shapes can recompile despite
112
+ dynamic=True. Benchmark representative warmed batches, synchronizing CUDA when
113
+ timing; no universal speedup is guaranteed and compile errors are not silently
114
+ converted to eager execution.
115
+
116
+ Context limits are separate from the microbatch token_budget. Fresh v0 exports
117
+ use adaptive-v1: reserve half the context for query and candidate (candidate gets
118
+ the odd token), then lend unused capacity subject to branch caps. The default
119
+ context is the backbone positional capacity (7,999 for v0 17M); candidates default
120
+ to min(3800, context - 3) tokens. Special tokens count. All candidates in a decision
121
+ share one truncated query. Queries follow saved balanced/right truncation and
122
+ candidates truncate on the right. Per-call context_length, query_length, and
123
+ document_length override limits without modifying the checkpoint. context_length
124
+ must not exceed positional capacity. prefix_layout overrides instruction_state
125
+ or state_instruction rendering; normally keep the exported default.
126
+
127
+ CLI: python inference_v0.py --model PATH_OR_HUB_ID --input requests.json
128
+ Add --device cuda, --compile, --batch-size 128, --token-budget 64000, or
129
+ --no-show-progress-bar as needed. Input JSON can be one request or an array;
130
+ results are JSON on stdout and progress is on stderr.
131
+ """
132
  from __future__ import annotations
133
 
134
  from dataclasses import dataclass
 
458
  from transformers.models.modernbert.modeling_modernbert import apply_rotary_pos_emb
459
 
460
  def _text(value):
461
+ """Return strings unchanged; serialize other decoded JSON values deterministically."""
462
  return value if isinstance(value, str) else json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(',', ':'))
463
 
464
  def input_groups(case_input, *, prefix_layout='instruction_state'):
465
+ """Render one native request into ordered Group objects for predict_groups().
466
+
467
+ Accept exactly state_json and decisions, with JSON strings for state,
468
+ instructions and candidate descriptions/content. See the module input schema.
469
+ This performs no tokenization or inference. Noul descriptions are embedded
470
+ into the query; numeric Score values remain output metadata. Labels and row
471
+ metadata are not accepted. prefix_layout controls instruction/state order.
472
+ """
473
  if not isinstance(case_input, dict) or set(case_input) != {'state_json', 'decisions'}:
474
  raise ValueError('Pass the input object only: state_json and decisions')
475
  state = case_input['state_json']
 
517
  """Independent prefix self-attention and suffix-aligned document attention."""
518
 
519
  def __init__(self, backbone):
520
+ """Wrap a ModernBERT backbone for independent shared-prefix SDPA execution."""
521
  super().__init__()
522
  self.backbone = backbone
523
 
524
  def qkv(self, layer, hidden, positions):
525
+ """Project one layer and apply rotary positions; return query/key/value tensors."""
526
  attn = layer.attn
527
  q, k, v = attn.Wqkv(layer.attn_norm(hidden)).view(*hidden.shape[:2], 3, -1, attn.head_dim).unbind(2)
528
  cos, sin = self.backbone.rotary_emb(hidden, positions, layer.attention_type)
 
531
 
532
  @staticmethod
533
  def attend(q, k, v, qmask, kmask, window):
534
+ """Apply masked SDPA, optionally windowed, and zero padded query positions."""
535
  qp = qmask.long().cumsum(1) - 1 + (kmask.sum(1) - qmask.sum(1))[:, None]
536
  kp = kmask.long().cumsum(1) - 1
537
  allowed = kmask[:, None, :].expand(-1, q.shape[1], -1)
 
542
 
543
  @staticmethod
544
  def update(layer, hidden, attended):
545
+ """Apply attention output and MLP residual updates for one encoder layer."""
546
  hidden = hidden + layer.attn.out_drop(layer.attn.Wo(attended.flatten(2)))
547
  return hidden + layer.mlp(layer.mlp_norm(hidden))
548
 
549
  def forward(self, prefix_ids, prefix_mask, doc_ids, doc_mask, owners):
550
+ """Encode unique prefixes once and attend each candidate to its owner prefix."""
551
  prefix = self.backbone.embeddings(prefix_ids)
552
  hidden = self.backbone.embeddings(doc_ids)
553
  pp = (prefix_mask.long().cumsum(1) - 1).clamp_min(0)
 
571
  budget_policy = 'adaptive-v1'
572
 
573
  def __init__(self, backbone, tokenizer, *, tasks, query_length=None, document_length=None, context_length=None, query_truncation='balanced', task_tokens=None, choice_interaction=None, prefix_layout='instruction_state'):
574
+ """Build a runtime from a ModernBERT backbone, tokenizer, heads and saved limits."""
575
  super().__init__()
576
  if backbone.config.model_type != 'modernbert':
577
  raise ValueError('Expected ModernBERT-compatible weights')
 
604
  self.settings = dict(tasks=list(tasks), query_length=query_length, document_length=document_length, context_length=capacity, query_truncation=query_truncation, task_tokens=self.task_tokens, choice_interaction=choice_interaction, prefix_layout=prefix_layout)
605
 
606
  def compile_inference(self, *, mode='default', dynamic=True, backend='inductor'):
607
+ """Enable lazy compilation of encoder.forward and return this runtime.
608
 
609
+ mode, dynamic and backend are forwarded to torch.compile; defaults are
610
+ "default", True and "inductor". Heads, tokenization, rendering and output
611
+ interpretation stay eager. First calls pay compilation cost and changing
612
+ shapes may recompile. Warm representative batches before timing. Compile
613
+ failures propagate; call disable_compile() to explicitly use eager mode.
614
  """
615
  self.eval().requires_grad_(False)
616
  self._compiled_forward = torch.compile(self.encoder.forward, mode=mode, dynamic=dynamic, backend=backend)
617
  return self
618
 
619
  def disable_compile(self):
620
+ """Clear the compiled encoder callable and restore eager execution; return None."""
621
  self._compiled_forward = None
622
 
623
  def preprocess(self, inputs, prompt=None, **kwargs):
624
+ """Reject raw-text encode(); typed decisions require predict() or predict_groups()."""
625
  raise ValueError('Typed decisions require candidate groups; use model.predict(input_object)')
626
 
627
  @staticmethod
628
  def _validate_limits(query_length, document_length, capacity):
629
+ """Validate branch caps including special tokens against shared positional capacity."""
630
  if not 3 <= query_length <= capacity - 2 or not 2 <= document_length <= capacity - 3:
631
  raise ValueError('Invalid query/document limits for this backbone')
632
 
633
  def _limits(self, query_length, document_length, context_length=None):
634
+ """Resolve per-call caps, clamping saved defaults to a smaller requested context."""
635
  capacity = self.context_length if context_length is None else context_length
636
  if not 5 <= capacity <= self.encoder.backbone.config.max_position_embeddings:
637
  raise ValueError('Context length exceeds the backbone positional capacity')
 
641
  return (q, d, capacity)
642
 
643
  def _documents(self, documents, tasks, limit):
644
+ """Tokenize candidates with task markers and final SEP, right-truncated to limit."""
645
  if len(tasks) != len(documents):
646
  raise ValueError('Tasks and candidates must align')
647
  if not documents:
 
650
  return [[self.task_token_ids[t], *d[:limit - 2], self.tokenizer.sep_token_id] if t in self.task_token_ids else [*d, self.tokenizer.sep_token_id] for t, d in zip(tasks, encoded, strict=True)]
651
 
652
  def _queries(self, queries, parts, limits):
653
+ """Tokenize queries under per-query caps with balanced or right truncation."""
654
  if not queries:
655
  return []
656
  if self.query_truncation == 'balanced':
 
661
  return [[self.tokenizer.cls_token_id, *q[:limit - 2], self.tokenizer.sep_token_id] for q, limit in zip(encoded, limits, strict=True)]
662
 
663
  def tokenize_branches(self, queries, documents, document_tasks=None, *, query_parts=None, query_length=None, document_length=None, context_length=None):
664
+ """Tokenize separate branch lists and return (query_ids, document_ids).
665
+
666
+ This low-level helper has no decision ownership, so allocation uses the
667
+ longest branches conservatively. Use prepare_groups() for per-decision
668
+ adaptive allocation or predict() for the full native-input interface.
669
+ Balanced truncation requires query_parts matching each rendered query.
670
+ """
671
  qlimit, dlimit, capacity = self._limits(query_length, document_length, context_length)
672
  dids = self._documents(documents, document_tasks or ['reranker'] * len(documents), dlimit)
673
  qids = self._queries(queries, query_parts, [qlimit] * len(queries))
 
705
  return [PreparedGroup(tuple(qmap[key]), g.task, qmap[key], [dmap[g.task, d][:limit - 1] + [dmap[g.task, d][-1]] if len(dmap[g.task, d]) > limit else dmap[g.task, d] for d in g.candidates], g.target, g.metadata) for g, key, limit in zip(groups, keys, document_limits, strict=True)]
706
 
707
  def collate_tokens(self, queries, documents, owners):
708
+ """Pad token lists on CPU and map each document to a unique prefix via owners."""
709
 
710
  def pad(rows):
711
  ids = torch.nn.utils.rnn.pad_sequence([torch.tensor(r, dtype=torch.long) for r in rows], batch_first=True, padding_value=self.tokenizer.pad_token_id)
 
715
  return dict(prefix_ids=pi, prefix_mask=pm, doc_ids=di, doc_mask=dm, owners=torch.tensor(owners, dtype=torch.long))
716
 
717
  def forward(self, features, **kwargs):
718
+ """Score collated tensor features and return the same dict with raw logits.
719
+
720
+ Requires prefix_ids, prefix_mask, doc_ids, doc_mask and owners; multiple
721
+ heads require head_indices. Optional Choice interaction also uses its
722
+ group indices/mask. Adds scores and sentence_embedding, both (N, 1).
723
+ These are logits, not typed predictions or normalized probabilities.
724
+ Prefer predict()/predict_groups(), which also manage batching and
725
+ inference mode. CUDA encoder autocast is BF16; task heads use FP32.
726
+ """
727
  device = next(self.parameters()).device
728
  with torch.autocast(device.type, dtype=torch.bfloat16, enabled=device.type == 'cuda'):
729
  execute = self._compiled_forward or self.encoder
 
752
 
753
  @torch.inference_mode()
754
  def predict_groups(self, groups, *, batch_size=128, token_budget=64000, query_length=None, document_length=None, context_length=None, show_progress_bar=True):
755
+ """Predict probability tensors for already-rendered decision groups.
756
+
757
+ Use predict(list_of_requests) for ordinary native-input batching. This
758
+ lower-level API accepts an iterable of Group objects (materialized),
759
+ such as input_groups(request), and skips typed result interpretation.
760
+ With balanced query truncation, each group needs matching QueryParts.
761
+
762
+ Returns a list of CPU FP32 tensors, one per group, each shaped
763
+ (number_of_candidates,) and normalized within that group. Group order
764
+ and candidate order match the input, regardless of internal bucketing.
765
+ Supported tasks must have heads in this checkpoint; empty input gives [].
766
+
767
+ batch_size (128) bounds decisions tokenized/sorted per window.
768
+ token_budget (64,000) bounds estimated padded microbatch work; a decision
769
+ exceeding it runs alone, with all candidates together. Length overrides
770
+ are per call; see predict() and the module's adaptive context description.
771
+ show_progress_bar counts completed decisions on stderr, not requests.
772
  """
773
  if batch_size < 1 or token_budget < 1:
774
  raise ValueError('batch_size and token_budget must be positive')
 
812
 
813
  @staticmethod
814
  def _interpret(group, probabilities):
815
+ """Convert one group distribution into typed fields or stable relative ranking."""
816
  assert group.metadata is not None and group.metadata.candidate_ids is not None
817
  if group.metadata.kind == 'ranking':
818
  return {'probabilities': dict(zip(group.metadata.candidate_ids, probabilities.tolist(), strict=True)), 'order': [group.metadata.candidate_ids[i] for i in probabilities.argsort(descending=True, stable=True).tolist()]}
819
  return asdict(interpret_prediction(group, probabilities))
820
 
821
  def predict(self, inputs, *, batch_size=128, token_budget=64000, query_length=None, document_length=None, context_length=None, prefix_layout=None, show_progress_bar=True):
822
+ """Predict typed decisions for one request or a batch of native requests.
823
+
824
+ Args:
825
+ inputs: Dict with exactly state_json and decisions, or an iterable
826
+ of those dicts (materialized as a list). See the module example.
827
+ Use a list for ordinary batched inference, including mixed tasks.
828
+ batch_size: Positive limit on requests rendered per window and on
829
+ decisions tokenized per window (default 128), not candidate count.
830
+ token_budget: Positive padded-work estimate per microbatch (64,000).
831
+ Complete decisions stay together; oversized ones run alone.
832
+ Lower this to reduce microbatch work; it is not a memory cap.
833
+ query_length: Optional query token cap, including special tokens.
834
+ document_length: Optional candidate token cap, including special tokens.
835
+ context_length: Optional shared query/candidate positional limit.
836
+ All length overrides apply only to this call; see adaptive-v1
837
+ allocation in the module docstring.
838
+ prefix_layout: Optional instruction_state or state_instruction;
839
+ None uses the exported rendering order.
840
+ show_progress_bar: Show completed requests on stderr (default True).
841
+
842
+ Returns:
843
+ One dict keyed by decision ID for a dict input, otherwise a list of
844
+ those dicts in input order. Each value contains the typed fields
845
+ documented above (Choice, Noul, Score, or relative ranking).
846
+ No decisions yields {}; an empty request list yields [].
847
+
848
+ For already-rendered Group objects and raw probability tensors, use
849
+ predict_groups(). For repeated GPU workloads, compile_inference() is
850
+ optional; warm representative batches before measuring throughput.
851
  """
852
  single = isinstance(inputs, dict)
853
  cases = [inputs] if single else list(inputs)
 
871
  return results[0] if single else results
872
 
873
  def get_sentence_embedding_dimension(self):
874
+ """Return the ST-compatible scalar logit dimension; this is not a text embedding."""
875
  return 1
876
 
877
  def get_config_dict(self):
878
+ """Return serializable runtime settings used by the ST module configuration."""
879
  return self.settings
880
 
881
  def save(self, output_path, *args, **kwargs):
882
+ """Save module config, backbone config, safetensors and tokenizer to a directory."""
883
  path = Path(output_path)
884
  path.mkdir(parents=True, exist_ok=True)
885
  self.save_config(str(path))
 
891
 
892
  @classmethod
893
  def load(cls, model_name_or_path, subfolder='', token=None, cache_folder=None, revision=None, local_files_only=False, init_defaults=None, **kwargs):
894
+ """Load this ST module from a local directory or Hugging Face repository.
895
+
896
+ Called by SentenceTransformer with subfolder, authentication, revision,
897
+ cache and offline settings. Restores FP32 weights with strict matching,
898
+ SDPA attention, eval mode and gradients disabled. Prefer constructing
899
+ BekkoSentenceTransformer for the public typed interface and device setup.
900
+ """
901
  hub: dict[str, Any] = dict(subfolder=subfolder, token=token, cache_folder=cache_folder, revision=revision, local_files_only=local_files_only)
902
  settings = cls.load_config(model_name_or_path, **hub)
903
  config_path = cls.load_file_path(model_name_or_path, 'backbone_config.json', **hub)
 
921
  """
922
 
923
  def __init__(self, *args, **kwargs):
924
+ """Load an exported v0 model with standard ST arguments, then validate its module."""
925
  super().__init__(*args, **kwargs)
926
  if len(self) != 1 or getattr(self[0], 'config_file_name', None) != 'inference_config.json':
927
  raise ValueError('BekkoSentenceTransformer requires an exported v0 checkpoint')
928
 
929
  def predict(self, inputs, *, batch_size=128, token_budget=64000, query_length=None, document_length=None, context_length=None, prefix_layout=None, show_progress_bar=True):
930
+ """Predict typed decisions for one request or a batch of native requests.
931
+
932
+ Args:
933
+ inputs: Dict with exactly state_json and decisions, or an iterable
934
+ of those dicts (materialized as a list). See the module example.
935
+ Use a list for ordinary batched inference, including mixed tasks.
936
+ batch_size: Positive limit on requests rendered per window and on
937
+ decisions tokenized per window (default 128), not candidate count.
938
+ token_budget: Positive padded-work estimate per microbatch (64,000).
939
+ Complete decisions stay together; oversized ones run alone.
940
+ Lower this to reduce microbatch work; it is not a memory cap.
941
+ query_length: Optional query token cap, including special tokens.
942
+ document_length: Optional candidate token cap, including special tokens.
943
+ context_length: Optional shared query/candidate positional limit.
944
+ All length overrides apply only to this call; see adaptive-v1
945
+ allocation in the module docstring.
946
+ prefix_layout: Optional instruction_state or state_instruction;
947
+ None uses the exported rendering order.
948
+ show_progress_bar: Show completed requests on stderr (default True).
949
+
950
+ Returns:
951
+ One dict keyed by decision ID for a dict input, otherwise a list of
952
+ those dicts in input order. Each value contains the typed fields
953
+ documented above (Choice, Noul, Score, or relative ranking).
954
+ No decisions yields {}; an empty request list yields [].
955
+
956
+ For already-rendered Group objects and raw probability tensors, use
957
+ predict_groups(). For repeated GPU workloads, compile_inference() is
958
+ optional; warm representative batches before measuring throughput.
959
+ """
960
  return cast(Any, self[0]).predict(inputs, batch_size=batch_size, token_budget=token_budget, query_length=query_length, document_length=document_length, context_length=context_length, prefix_layout=prefix_layout, show_progress_bar=show_progress_bar)
961
 
962
  def predict_groups(self, groups, *, batch_size=128, token_budget=64000, query_length=None, document_length=None, context_length=None, show_progress_bar=True):
963
+ """Predict probability tensors for already-rendered decision groups.
964
+
965
+ Use predict(list_of_requests) for ordinary native-input batching. This
966
+ lower-level API accepts an iterable of Group objects (materialized),
967
+ such as input_groups(request), and skips typed result interpretation.
968
+ With balanced query truncation, each group needs matching QueryParts.
969
+
970
+ Returns a list of CPU FP32 tensors, one per group, each shaped
971
+ (number_of_candidates,) and normalized within that group. Group order
972
+ and candidate order match the input, regardless of internal bucketing.
973
+ Supported tasks must have heads in this checkpoint; empty input gives [].
974
+
975
+ batch_size (128) bounds decisions tokenized/sorted per window.
976
+ token_budget (64,000) bounds estimated padded microbatch work; a decision
977
+ exceeding it runs alone, with all candidates together. Length overrides
978
+ are per call; see predict() and the module's adaptive context description.
979
+ show_progress_bar counts completed decisions on stderr, not requests.
980
+ """
981
  return cast(Any, self[0]).predict_groups(groups, batch_size=batch_size, token_budget=token_budget, query_length=query_length, document_length=document_length, context_length=context_length, show_progress_bar=show_progress_bar)
982
 
983
  def compile_inference(self, *, mode='default', dynamic=True, backend='inductor'):
984
+ """Enable lazy encoder compilation and return self for optional chaining.
985
+
986
+ For repeated inference, call once before representative warmup batches.
987
+ mode="default", dynamic=True and backend="inductor" pass to torch.compile.
988
+ Shapes may recompile; first-call latency includes compilation. Rendering,
989
+ tokenization, heads and typed outputs stay eager. See the module guide.
990
+ """
991
  cast(Any, self[0]).compile_inference(mode=mode, dynamic=dynamic, backend=backend)
992
  return self
993
 
 
997
  return self
998
 
999
  def main():
1000
+ """Run standalone inference from a JSON request or request array.
1001
+
1002
+ --model accepts a local export or Hub ID; private Hub access uses saved
1003
+ authentication. --compile enables lazy encoder compilation. Writes results
1004
+ to stdout and optional request progress to stderr. See --help for controls.
1005
+ """
1006
  parser = argparse.ArgumentParser(description=__doc__)
1007
  parser.add_argument('--model', required=True)
1008
  parser.add_argument('--input', required=True, type=Path, help='JSON file containing only native inference input')