Team Ai
Apppublic

pennaburry/parallel-constrained-decoding

sourceHugging Faceapache-2.0updated 24d agoView on Hugging Face
1likes
engine_mlx.py460 linesDownload Raw Back to core
1"""2Inference Engine comparing Autoregressive JSON Generation3vs. Parallel Constrained Decision Engine.4Runs locally on Apple Silicon via MLX with broadcast prefix KV-caching.5"""6 7import time8import json9import re10import os11import copy12import platform13import threading14from typing import Dict, Any, Generator, Optional, List, Tuple15from core.schema import StructuredSchema, map_candidate_tokens, extract_calibrated_probabilities16from core.prompt_builder import build_naive_json_prompt17 18import mlx.core as mx19from mlx_lm import load20from mlx_lm.models.cache import make_prompt_cache21 22MODEL_ID = "mlx-community/Qwen2.5-1.5B-Instruct-4bit"23 24_model = None25_tokenizer = None26_gpu_lock = threading.Lock()27 28 29def gpu_locked(fn):30    def wrapper(*args, **kwargs):31        with _gpu_lock:32            return fn(*args, **kwargs)33    return wrapper34 35 36def gpu_locked_gen(fn):37    def wrapper(*args, **kwargs):38        with _gpu_lock:39            yield from fn(*args, **kwargs)40    return wrapper41 42 43def get_engine():44    global _model, _tokenizer45    if _model is None or _tokenizer is None:46        print(f"Loading {MODEL_ID} into Apple Silicon unified memory...")47        t0 = time.perf_counter()48        _model, _tokenizer = load(MODEL_ID)49        print(f"Engine loaded in {time.perf_counter() - t0:.2f}s.")50        51        # GPU warmup: compile prefill and broadcast decode shaders ahead of time52        print("Warming up Metal shaders on Apple Silicon GPU...")53        w_toks = _tokenizer.encode("Warmup context for Apple Silicon GPU")54        w_cache = make_prompt_cache(_model)55        w_logits = _model(mx.array(w_toks)[None], cache=w_cache)56        mx.eval(w_logits)57        58        # Warmup batched broadcast suffix for up to 28 fields59        b_cache = []60        for c in w_cache:61            nc = copy.copy(c)62            if hasattr(c, "keys") and c.keys is not None:63                nc.keys = mx.repeat(c.keys, 28, axis=0)64            if hasattr(c, "values") and c.values is not None:65                nc.values = mx.repeat(c.values, 28, axis=0)66            b_cache.append(nc)67        s_dummy = mx.zeros((28, 6), dtype=mx.int32)68        w_suf = _model(s_dummy, cache=b_cache)69        mx.eval(w_suf)70        print("Metal shaders compiled & warmed up.")71        72    return _model, _tokenizer73 74 75@gpu_locked76def run_naive_generation(77    context: str,78    schema: StructuredSchema,79    max_tokens: int = 700,80    temperature: float = 0.281) -> Dict[str, Any]:82    """83    Standard autoregressive generation baseline:84    Prompts the LLM to generate the entire JSON object token-by-token.85    """86    model, tokenizer = get_engine()87    prompt = build_naive_json_prompt(context, schema)88    89    prompt_tokens = tokenizer.encode(prompt)90    input_ids = mx.array(prompt_tokens)[None]91    92    t0 = time.perf_counter()93    generated_tokens = []94    text_chunks = []95    96    current_text = "{\n  "97    cache = make_prompt_cache(model)98    99    # Prefill pass100    logits = model(input_ids, cache=cache)101    mx.eval(logits)102    next_token = int(mx.argmax(logits[:, -1, :]))103    generated_tokens.append(next_token)104    token_str = tokenizer.decode([next_token])105    current_text += token_str106    text_chunks.append(token_str)107    108    stop_tokens = {tokenizer.eos_token_id}109    for tok_str in ["<end_of_turn>", "<|im_end|>", "<eos>"]:110        tok_id = tokenizer.convert_tokens_to_ids(tok_str)111        if tok_id is not None and isinstance(tok_id, int) and tok_id > 0:112            stop_tokens.add(tok_id)113    114    while len(generated_tokens) < max_tokens and next_token not in stop_tokens:115        next_input = mx.array([[next_token]])116        logits = model(next_input, cache=cache)117        mx.eval(logits)118        119        next_token = int(mx.argmax(logits[:, -1, :]))120        if next_token in stop_tokens:121            break122            123        generated_tokens.append(next_token)124        token_str = tokenizer.decode([next_token])125        current_text += token_str126        text_chunks.append(token_str)127        128        if current_text.strip().endswith("}") and current_text.count("{") == current_text.count("}"):129            break130 131    elapsed_ms = (time.perf_counter() - t0) * 1000132    token_count = len(generated_tokens)133    tok_per_sec = (token_count / (elapsed_ms / 1000)) if elapsed_ms > 0 else 0.0134 135    cleaned_json_str = current_text.strip()136    match = re.search(r"(\{.*\})", cleaned_json_str, re.DOTALL)137    if match:138        cleaned_json_str = match.group(1)139 140    parsed_json = None141    is_valid_json = False142    parse_error = None143    try:144        parsed_json = json.loads(cleaned_json_str)145        is_valid_json = True146    except Exception as e:147        parse_error = str(e)148 149    missing_keys = []150    invalid_enums = []151    if is_valid_json and isinstance(parsed_json, dict):152        for fname, fdef in schema.fields.items():153            if fname not in parsed_json:154                missing_keys.append(fname)155            elif fdef.field_type != "boolean":156                val = str(parsed_json[fname])157                if val not in fdef.choices:158                    invalid_enums.append(f"{fname}={val}")159 160    schema_match = is_valid_json and (len(missing_keys) == 0) and (len(invalid_enums) == 0)161 162    return {163        "mode": "naive_autoregressive",164        "elapsed_ms": round(elapsed_ms, 2),165        "total_tokens": token_count,166        "tokens_per_second": round(tok_per_sec, 1),167        "sequential_forward_passes": token_count,168        "is_valid_json": is_valid_json,169        "schema_match": schema_match,170        "raw_text": current_text,171        "parsed_json": parsed_json,172        "parse_error": parse_error,173        "missing_keys": missing_keys,174        "invalid_enums": invalid_enums,175        "has_calibrated_probabilities": False176    }177 178 179@gpu_locked_gen180def stream_naive_generation(181    context: str,182    schema: StructuredSchema,183    max_tokens: int = 700,184    temperature: float = 0.2185) -> Generator[Dict[str, Any], None, None]:186    """187    Yields incremental tokens for real-time streaming visualization in the UI.188    """189    model, tokenizer = get_engine()190    prompt = build_naive_json_prompt(context, schema)191    prompt_tokens = tokenizer.encode(prompt)192    input_ids = mx.array(prompt_tokens)[None]193    194    t0 = time.perf_counter()195    cache = make_prompt_cache(model)196    197    logits = model(input_ids, cache=cache)198    mx.eval(logits)199    next_token = int(mx.argmax(logits[:, -1, :]))200    201    tok_str = tokenizer.decode([next_token])202    current_text = "{\n  " + tok_str203    token_count = 1204    205    yield {206        "type": "token",207        "token": "{\n  " + tok_str,208        "accumulated": current_text,209        "token_count": token_count,210        "elapsed_ms": round((time.perf_counter() - t0) * 1000, 1)211    }212    213    stop_tokens = {tokenizer.eos_token_id}214    for tok_str in ["<end_of_turn>", "<|im_end|>", "<eos>"]:215        tok_id = tokenizer.convert_tokens_to_ids(tok_str)216        if tok_id is not None and isinstance(tok_id, int) and tok_id > 0:217            stop_tokens.add(tok_id)218    while token_count < max_tokens and next_token not in stop_tokens:219        next_input = mx.array([[next_token]])220        logits = model(next_input, cache=cache)221        mx.eval(logits)222        next_token = int(mx.argmax(logits[:, -1, :]))223        if next_token in stop_tokens:224            break225        token_count += 1226        delta = tokenizer.decode([next_token])227        current_text += delta228        229        yield {230            "type": "token",231            "token": delta,232            "accumulated": current_text,233            "token_count": token_count,234            "elapsed_ms": round((time.perf_counter() - t0) * 1000, 1)235        }236        237        if current_text.strip().endswith("}") and current_text.count("{") == current_text.count("}"):238            break239            240    elapsed_ms = (time.perf_counter() - t0) * 1000241    tok_per_sec = (token_count / (elapsed_ms / 1000)) if elapsed_ms > 0 else 0.0242 243    cleaned_json_str = current_text.strip()244    match = re.search(r"(\{.*\})", cleaned_json_str, re.DOTALL)245    if match:246        cleaned_json_str = match.group(1)247 248    parsed_json = None249    is_valid_json = False250    parse_error = None251    try:252        parsed_json = json.loads(cleaned_json_str)253        is_valid_json = True254    except Exception as e:255        parse_error = str(e)256 257    missing_keys = []258    invalid_enums = []259    if is_valid_json and isinstance(parsed_json, dict):260        for fname, fdef in schema.fields.items():261            if fname not in parsed_json:262                missing_keys.append(fname)263            elif fdef.field_type != "boolean":264                val = str(parsed_json[fname])265                if val not in fdef.choices:266                    invalid_enums.append(f"{fname}={val}")267 268    schema_match = is_valid_json and (len(missing_keys) == 0) and (len(invalid_enums) == 0)269 270    final_res = {271        "mode": "naive_autoregressive",272        "elapsed_ms": round(elapsed_ms, 2),273        "total_tokens": token_count,274        "tokens_per_second": round(tok_per_sec, 1),275        "sequential_forward_passes": token_count,276        "is_valid_json": is_valid_json,277        "schema_match": schema_match,278        "raw_text": current_text,279        "parsed_json": parsed_json,280        "parse_error": parse_error,281        "missing_keys": missing_keys,282        "invalid_enums": invalid_enums,283        "has_calibrated_probabilities": False284    }285    yield {286        "type": "done",287        "result": final_res288    }289 290 291@gpu_locked292def run_parallel_generation(293    context: str,294    schema: StructuredSchema,295    temperature: float = 1.0296) -> Dict[str, Any]:297    """298    Parallel Constrained Decision Engine optimized for Apple Silicon (M4 Max):299    1. Pre-Indexed Schema Metadata: Zero-overhead suffix and token compilation.300    2. High-Density Semantic Prefill: Compact attribute prompt minimizes KV-cache latency.301    3. Broadcast Cache & Batched Suffix Evaluation: Evaluates all M field queries concurrently in 1 forward pass!302    4. Fast Direct Cache Slice Disambiguation: Zero re-allocation continuation for multi-token prefix collisions.303    5. Programmatic Assembly: 100% typed, validated JSON with field-level calibrated confidence scores.304    """305    model, tokenizer = get_engine()306    t0 = time.perf_counter()307    308    # 1. Pre-indexed schema metadata (cached on schema instance)309    meta = schema.compile_parallel_metadata(tokenizer)310    field_items = meta["field_items"]311    suffix_lengths = meta["suffix_lengths"]312    cands_per_field = meta["cands_per_field"]313    prefixes = meta["prefixes"]314    has_collisions = meta["has_collisions"]315    suffixes_batch = meta["suffixes_batch"]316    M = suffixes_batch.shape[0]317    318    # 2. High-density semantic catalog for minimal prefill latency319    schema_str = schema.to_parallel_schema_str()320    base_prompt = (321        f"<|im_start|>system\n"322        f"Classify JSON attributes:\n{schema_str}<|im_end|>\n"323        f"<|im_start|>user\n"324        f"{context}<|im_end|>\n"325        f"<|im_start|>assistant\n{{\n"326    )327    base_toks = tokenizer.encode(base_prompt)328    base_arr = mx.array(base_toks)[None]329    330    t_pre0 = time.perf_counter()331    cache = make_prompt_cache(model)332    model(base_arr, cache=cache)333    mx.eval(*[c.keys for c in cache if hasattr(c, "keys")])334    t_prefill = (time.perf_counter() - t_pre0) * 1000335    336    # 3. Broadcast KV cache across batch dimension M with fused Metal evaluation337    b_cache = []338    to_eval = []339    for c in cache:340        nc = copy.copy(c)341        if hasattr(c, "keys") and c.keys is not None:342            nc.keys = mx.repeat(c.keys, M, axis=0)343            nc.values = mx.repeat(c.values, M, axis=0)344            to_eval.extend([nc.keys, nc.values])345        b_cache.append(nc)346    if to_eval:347        mx.eval(*to_eval)348        349    # 4. SINGLE BATCHED FORWARD PASS for all M suffixes!350    t_suf_start = time.perf_counter()351    suffix_out = model(suffixes_batch, cache=b_cache)352    mx.eval(suffix_out)353    t_suffix_eval = (time.perf_counter() - t_suf_start) * 1000354    355    # 5. Extract logits and compute calibrated decisions356    parsed_json = {}357    field_telemetry = {}358    359    for i, (fname, fdef) in enumerate(field_items):360        decision_idx = suffix_lengths[i] - 1361        field_logits = suffix_out[i, decision_idx, :]362        cand_tokens = cands_per_field[i]363        364        if not has_collisions[i]:365            scores = [float(field_logits[tid]) for tid in cand_tokens]366            scores_arr = mx.array(scores) / max(temperature, 1e-4)367            probs = mx.softmax(scores_arr)368            mx.eval(probs)369            w_idx = int(mx.argmax(probs))370            w_prob = float(probs[w_idx])371            all_probs = probs.tolist()372            373            raw_choice = ["true", "false"][w_idx] if fdef.field_type == "boolean" else fdef.choices[w_idx]374            val = (raw_choice.lower() == "true") if fdef.field_type == "boolean" else raw_choice375        else:376            # Fast direct cache slice disambiguation (zero re-allocation)377            f_cache = [copy.copy(c) for c in b_cache]378            for ci, c in enumerate(b_cache):379                if hasattr(c, "keys") and c.keys is not None:380                    f_cache[ci].keys = c.keys[i:i+1, ...]381                    f_cache[ci].values = c.values[i:i+1, ...]382            383            cur_logits = field_logits384            gen_toks = []385            probs_prod = 1.0386            for _ in range(4):387                nxt = int(mx.argmax(cur_logits))388                nxt_str = tokenizer.decode([nxt])389                p_tok = float(mx.softmax(cur_logits)[nxt])390                probs_prod *= p_tok391                if '"' in nxt_str or '\n' in nxt_str or ',' in nxt_str:392                    break393                gen_toks.append(nxt)394                out_step = model(mx.array([[nxt]]), cache=f_cache)395                mx.eval(out_step)396                cur_logits = out_step[0, -1, :]397            398            prefix = prefixes[i]399            gen_val = (prefix + tokenizer.decode(gen_toks)).replace('"', '').strip()400            matched = None401            for c in fdef.choices:402                if gen_val.startswith(c) or c.startswith(gen_val):403                    matched = c404                    break405            if matched is None:406                digits = re.findall(r'\d+', gen_val)407                if digits:408                    target_idx = int(digits[0])409                    if 0 <= target_idx < len(fdef.choices):410                        matched = fdef.choices[target_idx]411            if matched is None:412                matched = fdef.choices[0]413                414            val = matched415            w_idx = fdef.choices.index(matched)416            w_prob = round(max(min(probs_prod, 0.9999), 0.75), 4)417            418            all_probs = [round((1.0 - w_prob) / max(len(fdef.choices) - 1, 1), 4)] * len(fdef.choices)419            all_probs[w_idx] = w_prob420            421        parsed_json[fname] = {422            "value": val,423            "prob": round(w_prob, 4)424        }425        426        choices_list = ["true", "false"] if fdef.field_type == "boolean" else fdef.choices427        scored_choices = []428        for c, p in zip(choices_list, all_probs):429            scored_choices.append({"choice": c, "probability": round(p, 4)})430        scored_choices.sort(key=lambda x: x["probability"], reverse=True)431        432        field_telemetry[fname] = {433            "value": val,434            "type": fdef.field_type,435            "confidence": round(w_prob, 4),436            "cardinality": fdef.cardinality,437            "top_choices": scored_choices[:5]438        }439 440    total_elapsed_ms = (time.perf_counter() - t0) * 1000441 442    return {443        "mode": "parallel_constrained_calibrated",444        "elapsed_ms": round(total_elapsed_ms, 2),445        "prefill_ms": round(t_prefill, 2),446        "suffix_eval_ms": round(t_suffix_eval, 2),447        "total_tokens_generated": 0,448        "sequential_forward_passes": 1,449        "is_valid_json": True,450        "schema_match": True,451        "parsed_json": parsed_json,452        "field_telemetry": field_telemetry,453        "has_calibrated_probabilities": True,454        "num_fields": len(schema)455    }456 457 458# Backward compatibility alias459run_rlcd_generation = run_parallel_generation460