valeriow/parallel-constrained-decoding
0
1"""2Schema definitions, validation, and sub-vocabulary token mapping for parallel constrained decisions.3Supports booleans and categorical enums with cardinality up to 255.4"""5 6from typing import Dict, Any, List, Tuple, Optional7import numpy as np8 9 10class FieldDefinition:11 def __init__(self, name: str, field_type: str, description: str, choices: Optional[List[str]] = None):12 self.name = name13 self.field_type = field_type.lower()14 self.description = description15 16 if self.field_type == "boolean":17 self.choices = ["true", "false"]18 elif self.field_type in ("enum", "choice", "selection"):19 if not choices or len(choices) == 0:20 raise ValueError(f"Field '{name}' of type enum must have choices defined.")21 if len(choices) > 255:22 raise ValueError(f"Field '{name}' exceeds maximum cardinality of 255 choices (got {len(choices)}).")23 self.choices = choices24 else:25 raise ValueError(f"Unsupported field type '{field_type}'. Supported types: 'boolean' and 'enum'.")26 27 self.cached_candidate_token_ids: Optional[List[List[int]]] = None28 29 @property30 def cardinality(self) -> int:31 return len(self.choices)32 33 def compile_candidate_tokens(self, tokenizer):34 """Pre-indexes and caches candidate token IDs so inference runs in microseconds."""35 if self.cached_candidate_token_ids is not None:36 return self.cached_candidate_token_ids37 38 candidate_tokens_per_choice = []39 if self.field_type == "boolean":40 true_variants = ['true', ' true', 'True', ' True', 'TRUE', 'yes', ' yes']41 true_ids = []42 for v in true_variants:43 toks = tokenizer.encode(v, add_special_tokens=False)44 if toks:45 true_ids.append(toks[0])46 candidate_tokens_per_choice.append(list(set(true_ids)))47 48 false_variants = ['false', ' false', 'False', ' False', 'FALSE', 'no', ' no']49 false_ids = []50 for v in false_variants:51 toks = tokenizer.encode(v, add_special_tokens=False)52 if toks:53 false_ids.append(toks[0])54 candidate_tokens_per_choice.append(list(set(false_ids)))55 else:56 for choice in self.choices:57 c_clean = str(choice).strip()58 variants = [' ' + c_clean, c_clean]59 ids = []60 for v in variants:61 toks = tokenizer.encode(v, add_special_tokens=False)62 if toks:63 ids.append(toks[0])64 candidate_tokens_per_choice.append(list(set(ids)))65 66 self.cached_candidate_token_ids = candidate_tokens_per_choice67 return self.cached_candidate_token_ids68 69 def to_dict(self) -> Dict[str, Any]:70 return {71 "name": self.name,72 "type": self.field_type,73 "description": self.description,74 "choices": self.choices,75 "cardinality": self.cardinality,76 }77 78 79class StructuredSchema:80 def __init__(self, schema_dict: Dict[str, Any], tokenizer=None):81 self.fields: Dict[str, FieldDefinition] = {}82 for field_name, spec in schema_dict.items():83 field_type = spec.get("type", "enum")84 description = spec.get("description", "")85 choices = spec.get("choices", None)86 fdef = FieldDefinition(87 name=field_name,88 field_type=field_type,89 description=description,90 choices=choices91 )92 if tokenizer is not None:93 fdef.compile_candidate_tokens(tokenizer)94 self.fields[field_name] = fdef95 96 def compile_all_tokens(self, tokenizer):97 for fdef in self.fields.values():98 fdef.compile_candidate_tokens(tokenizer)99 100 def get_field_names(self) -> List[str]:101 return list(self.fields.keys())102 103 def __getitem__(self, key: str) -> FieldDefinition:104 return self.fields[key]105 106 def __len__(self) -> int:107 return len(self.fields)108 109 def to_json_schema_prompt_str(self) -> str:110 """Returns a clean TypeScript/JSON schema representation for naive LLM prompting."""111 lines = ["{"]112 for name, field in self.fields.items():113 if field.field_type == "boolean":114 lines.append(f' "{name}": boolean, // {field.description}')115 else:116 choices_limit = 20 if len(field.choices) > 50 else len(field.choices)117 choices_str = " | ".join(f'"{c}"' for c in field.choices[:choices_limit])118 if len(field.choices) > choices_limit:119 choices_str += f" | ... ({len(field.choices)} total options)"120 lines.append(f' "{name}": {choices_str}, // {field.description}')121 lines.append("}")122 return "\n".join(lines)123 124 def to_parallel_schema_str(self) -> str:125 """Returns a high-density, compact description catalog for minimal prefill token latency."""126 lines = []127 for name, field in self.fields.items():128 desc = field.description.split('\n')[0].strip()129 lines.append(f' "{name}": {desc}')130 return "\n".join(lines)131 132 to_rlcd_schema_str = to_parallel_schema_str133 134 def compile_parallel_metadata(self, tokenizer):135 """Pre-indexes and caches compact suffixes, token candidate IDs, and common prefixes."""136 if hasattr(self, "_parallel_metadata") and self._parallel_metadata is not None:137 return self._parallel_metadata138 139 import os140 field_items = list(self.fields.items())141 suffix_tok_lists = []142 suffix_lengths = []143 cands_per_field = []144 prefixes = []145 has_collisions = []146 147 for fname, fdef in field_items:148 if fdef.field_type == "boolean":149 suffix = f' "{fname}": '150 cands = [151 tokenizer.encode("true", add_special_tokens=False)[0],152 tokenizer.encode("false", add_special_tokens=False)[0]153 ]154 prefix = ""155 else:156 prefix = os.path.commonprefix(fdef.choices)157 suffix = f' "{fname}": "{prefix}'158 cands = []159 for c in fdef.choices:160 rem = c[len(prefix):]161 c_toks = tokenizer.encode(rem, add_special_tokens=False)162 cands.append(c_toks[0] if c_toks else tokenizer.encode('"', add_special_tokens=False)[0])163 toks = tokenizer.encode(suffix, add_special_tokens=False)164 suffix_tok_lists.append(toks)165 suffix_lengths.append(len(toks))166 cands_per_field.append(cands)167 prefixes.append(prefix)168 has_collisions.append(len(set(cands)) < len(cands))169 170 max_s_len = max(suffix_lengths)171 pad_id = tokenizer.pad_token_id or 0172 padded = [s + [pad_id] * (max_s_len - len(s)) for s in suffix_tok_lists]173 try:174 import mlx.core as mx175 suffixes_batch = mx.array(padded, dtype=mx.int32)176 except Exception:177 suffixes_batch = np.array(padded, dtype=np.int32)178 179 self._parallel_metadata = {180 "field_items": field_items,181 "suffix_lengths": suffix_lengths,182 "cands_per_field": cands_per_field,183 "prefixes": prefixes,184 "has_collisions": has_collisions,185 "suffixes_batch": suffixes_batch186 }187 return self._parallel_metadata188 189 compile_rlcd_metadata = compile_parallel_metadata190 191 192def map_candidate_tokens(tokenizer, choices: List[str], is_boolean: bool = False) -> List[List[int]]:193 """Helper fallback when field definition is not pre-compiled."""194 f = FieldDefinition("tmp", "boolean" if is_boolean else "enum", "", choices if not is_boolean else None)195 return f.compile_candidate_tokens(tokenizer)196 197 198def extract_calibrated_probabilities(199 next_token_logits: np.ndarray,200 candidate_token_ids_list: List[List[int]],201 temperature: float = 1.0202) -> Tuple[int, float, List[float]]:203 """204 Takes the logits at the decision token position and computes exact205 calibrated probabilities across only the constrained candidate choices (K <= 255).206 """207 choice_scores = []208 for token_ids in candidate_token_ids_list:209 if not token_ids:210 choice_scores.append(-1e9)211 continue212 score = max(float(next_token_logits[tid]) for tid in token_ids)213 choice_scores.append(score)214 215 scores = np.array(choice_scores, dtype=np.float32) / max(temperature, 1e-4)216 shifted = scores - np.max(scores)217 exp_scores = np.exp(shifted)218 probs = exp_scores / (np.sum(exp_scores) + 1e-12)219 220 winner_idx = int(np.argmax(probs))221 winner_prob = float(probs[winner_idx])222 223 return winner_idx, winner_prob, probs.tolist()224 