File size: 4,951 Bytes
4dd405a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
"""Constrained-decoding inference wrapper using lm-format-enforcer.

Forces the LLM output to match the V1 JSON schema:
  - segments: list of {intent: enum, slots: dict, text: str}
  - intent: one of 51 enum values (incl. "unknown")
  - abstain_reason: null | str

This eliminates the "invented intent" failure mode at inference time.

Two backends:
  - mlx-lm path: integrates lm-format-enforcer's TokenEnforcer with mlx_lm.generate
  - transformers path: standard JsonSchemaParser + LogitsProcessor for HF models

Usage:
    from infer_constrained import constrained_infer_mlx
    out = constrained_infer_mlx(model, tok, system, user, schema=V1_SCHEMA)
"""
from __future__ import annotations
import json
from pathlib import Path

ROOT = Path(__file__).resolve().parents[3]
INTENTS_50 = sorted(json.loads((ROOT / "poc/deberta_intent/checkpoints-base/label_mapping.json").read_text())["intent2id"].keys())
INTENTS_51 = INTENTS_50 + ["unknown"]
SLOTS_LOWER = ["altimeter_setting", "altitude", "approach_type", "call_sign", "clock_position",
               "direction", "distance", "facility", "fix", "frequency", "heading", "pattern_leg",
               "route", "runway", "speed", "taxiway", "time", "transponder_code", "turn_direction",
               "sequence"]


def build_schema() -> dict:
    """JSON schema enforcing V1 contract: segments + abstain_reason, enum-constrained."""
    return {
        "type": "object",
        "properties": {
            "segments": {
                "type": "array",
                "minItems": 1,
                "maxItems": 6,
                "items": {
                    "type": "object",
                    "properties": {
                        "intent": {"type": "string", "enum": INTENTS_51},
                        "slots": {
                            "type": "object",
                            "additionalProperties": False,
                            "properties": {k: {"type": "string"} for k in SLOTS_LOWER},
                        },
                        "text": {"type": "string"},
                    },
                    "required": ["intent", "slots", "text"],
                    "additionalProperties": False,
                },
            },
            "abstain_reason": {"type": ["string", "null"]},
        },
        "required": ["segments", "abstain_reason"],
        "additionalProperties": False,
    }


V1_SCHEMA = build_schema()


def constrained_infer_transformers(model, tokenizer, prompt: str, max_tokens: int = 512) -> str:
    """Constrained generation with HF transformers + lm-format-enforcer.

    Use this on the Modal-trained Qwen3-32B + adapter (transformers stack).
    """
    from lmformatenforcer import JsonSchemaParser
    from lmformatenforcer.integrations.transformers import build_transformers_prefix_allowed_tokens_fn
    parser = JsonSchemaParser(V1_SCHEMA)
    prefix_fn = build_transformers_prefix_allowed_tokens_fn(tokenizer, parser)
    inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
    output_ids = model.generate(
        **inputs, max_new_tokens=max_tokens, do_sample=False,
        prefix_allowed_tokens_fn=prefix_fn, pad_token_id=tokenizer.eos_token_id,
    )
    return tokenizer.decode(output_ids[0][inputs.input_ids.shape[1]:], skip_special_tokens=True)


def constrained_infer_mlx(model, tokenizer, system: str, user: str, max_tokens: int = 512) -> str:
    """Best-effort constrained generation for MLX.

    lm-format-enforcer doesn't have a first-class MLX integration; we use
    its TokenEnforcerTokenizerData + a custom logits-processor wrapper.
    Falls back to unconstrained if the integration fails.
    """
    try:
        from lmformatenforcer import JsonSchemaParser, TokenEnforcer
        from lmformatenforcer.integrations.mlx import build_mlx_logits_processor
    except ImportError:
        # mlx integration not available — fall back
        from mlx_lm import generate
        msgs = [{"role": "system", "content": system}, {"role": "user", "content": user}]
        prompt = tokenizer.apply_chat_template(msgs, add_generation_prompt=True, tokenize=False)
        return generate(model, tokenizer, prompt=prompt, max_tokens=max_tokens, verbose=False)

    parser = JsonSchemaParser(V1_SCHEMA)
    msgs = [{"role": "system", "content": system}, {"role": "user", "content": user}]
    prompt = tokenizer.apply_chat_template(msgs, add_generation_prompt=True, tokenize=False)
    # Custom processor — implementation depends on which mlx-lm version exposes logits_processors
    from mlx_lm import generate
    return generate(model, tokenizer, prompt=prompt, max_tokens=max_tokens,
                    logits_processors=[build_mlx_logits_processor(tokenizer, parser)],
                    verbose=False)


if __name__ == "__main__":
    schema = V1_SCHEMA
    print(f"V1 schema enums: {len(INTENTS_51)} intents, {len(SLOTS_LOWER)} slot keys")
    print(json.dumps(schema, indent=2)[:500])