Team Ai
Apppublic

valeriow/parallel-constrained-decoding

sourceHugging Faceapache-2.0updated 25d agoView on Hugging Face
0likes
schema.py224 linesDownload Raw Back to core
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