atc-parser-scripts / infer_constrained.py
rudeparis's picture
Upload infer_constrained.py with huggingface_hub
4dd405a verified
Raw
History Blame Contribute Delete
4.95 kB
"""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])