harshatheg's picture
Fix tuple index unpacking in PyTorch engine
292b8ea verified
Raw
History Blame Contribute Delete
11 kB
"""
PyTorch / CUDA / CPU Inference Engine for Parallel Constrained Decoding.
Optimized for Linux containers, Hugging Face Spaces (ZeroGPU & CUDA), and cloud environments.
"""
import os
import time
import json
import copy
import re
import threading
from typing import Dict, Any, Generator, Optional, List, Tuple
import torch
import torch.nn.functional as F
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
from transformers.cache_utils import DynamicCache
from core.schema import StructuredSchema
from core.prompt_builder import build_naive_json_prompt
MODEL_ID = os.environ.get("MODEL_ID", "Qwen/Qwen2.5-1.5B-Instruct")
_torch_model = None
_torch_tokenizer = None
_torch_device = None
_gpu_lock = threading.Lock()
# Support Hugging Face Spaces ZeroGPU if available
try:
import spaces
gpu_decorator = spaces.GPU(duration=60)
except Exception:
def gpu_decorator(fn=None, **kwargs):
if fn is not None:
return fn
return lambda f: f
def get_torch_engine():
global _torch_model, _torch_tokenizer, _torch_device
if _torch_model is None or _torch_tokenizer is None:
_torch_device = "cuda" if torch.cuda.is_available() else "cpu"
if _torch_device == "cuda":
dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
else:
dtype = torch.float32
print(f"Loading {MODEL_ID} on {_torch_device} ({dtype})...")
t0 = time.perf_counter()
_torch_tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
load_kwargs = {
"torch_dtype": dtype,
"low_cpu_mem_usage": True
}
if _torch_device == "cuda":
load_kwargs["device_map"] = "auto"
_torch_model = AutoModelForCausalLM.from_pretrained(MODEL_ID, **load_kwargs)
if _torch_device == "cpu":
_torch_model = _torch_model.to("cpu")
_torch_model.eval()
print(f"Engine loaded on {_torch_device} in {time.perf_counter() - t0:.2f}s.")
return _torch_model, _torch_tokenizer, _torch_device
@gpu_decorator
def run_parallel_generation_torch(
context: str,
schema: StructuredSchema,
temperature: float = 1.0
) -> Dict[str, Any]:
"""
Parallel Constrained Decision Engine running on PyTorch (CUDA / CPU).
Evaluates all schema fields concurrently against a broadcast prefix KV-cache.
"""
model, tokenizer, device = get_torch_engine()
t0 = time.perf_counter()
# 1. Compile schema metadata
meta = schema.compile_parallel_metadata(tokenizer)
field_items = meta["field_items"]
suffix_lengths = meta["suffix_lengths"]
cands_per_field = meta["cands_per_field"]
prefixes = meta["prefixes"]
has_collisions = meta["has_collisions"]
suffixes_batch = meta["suffixes_batch"]
M = len(field_items)
# 2. High-density semantic catalog prefill
schema_str = schema.to_parallel_schema_str()
base_prompt = (
f"<|im_start|>system\n"
f"Classify JSON attributes:\n{schema_str}<|im_end|>\n"
f"<|im_start|>user\n"
f"{context}<|im_end|>\n"
f"<|im_start|>assistant\n{{\n"
)
base_toks = tokenizer.encode(base_prompt, return_tensors="pt").to(device)
t_pre0 = time.perf_counter()
with torch.no_grad():
base_out = model(base_toks, use_cache=True)
base_cache = base_out.past_key_values
t_prefill = (time.perf_counter() - t_pre0) * 1000
# 3. Parallel Suffix Evaluation
t_suf0 = time.perf_counter()
pad_id = tokenizer.pad_token_id or tokenizer.eos_token_id or 0
suffix_arr = torch.tensor(suffixes_batch, dtype=torch.long, device=device)
suffix_mask = (suffix_arr != pad_id).long()
# Broadcast KV cache to batch size M
with torch.no_grad():
batched_cache = copy.deepcopy(base_cache)
if hasattr(batched_cache, "batch_repeat_interleave"):
batched_cache.batch_repeat_interleave(M)
elif isinstance(batched_cache, tuple):
batched_cache = tuple(
tuple(t.repeat(M, 1, 1, 1) for t in layer)
for layer in batched_cache
)
prefix_len = base_toks.shape[1]
prefix_mask = torch.ones((M, prefix_len), dtype=torch.long, device=device)
full_mask = torch.cat([prefix_mask, suffix_mask], dim=1)
out = model(suffix_arr, past_key_values=batched_cache, attention_mask=full_mask)
suffix_out = out.logits
t_suffix_eval = (time.perf_counter() - t_suf0) * 1000
# 4. Slicing, Disambiguation & Softmax
parsed_json = {}
field_telemetry = {}
for i, (fname, fdef) in enumerate(field_items):
decision_idx = suffix_lengths[i] - 1
field_logits = suffix_out[i, decision_idx, :]
cand_tokens = cands_per_field[i]
scores = [float(field_logits[tid].item()) for tid in cand_tokens]
scores_t = torch.tensor(scores, dtype=torch.float32) / max(temperature, 1e-4)
probs = F.softmax(scores_t, dim=-1).tolist()
w_idx = int(torch.argmax(scores_t).item())
w_prob = float(probs[w_idx])
all_probs = probs
if fdef.field_type == "boolean":
val = (w_idx == 0)
else:
val = fdef.choices[w_idx]
parsed_json[fname] = {
"value": val,
"prob": round(w_prob, 4)
}
choices_list = ["true", "false"] if fdef.field_type == "boolean" else fdef.choices
scored_choices = []
for c, p in zip(choices_list, all_probs):
scored_choices.append({"choice": c, "probability": round(p, 4)})
scored_choices.sort(key=lambda x: x["probability"], reverse=True)
field_telemetry[fname] = {
"value": val,
"type": fdef.field_type,
"confidence": round(w_prob, 4),
"cardinality": fdef.cardinality,
"top_choices": scored_choices[:5]
}
total_elapsed_ms = (time.perf_counter() - t0) * 1000
return {
"mode": "parallel_constrained_calibrated",
"elapsed_ms": round(total_elapsed_ms, 2),
"prefill_ms": round(t_prefill, 2),
"suffix_eval_ms": round(t_suffix_eval, 2),
"total_tokens_generated": 0,
"sequential_forward_passes": 1,
"is_valid_json": True,
"schema_match": True,
"parsed_json": parsed_json,
"field_telemetry": field_telemetry,
"has_calibrated_probabilities": True,
"num_fields": len(schema),
"device": device
}
@gpu_decorator
def run_naive_generation_torch(
context: str,
schema: StructuredSchema,
temperature: float = 0.2,
max_new_tokens: int = 512
) -> Dict[str, Any]:
"""
Standard autoregressive baseline using PyTorch.
"""
model, tokenizer, device = get_torch_engine()
t0 = time.perf_counter()
prompt = build_naive_json_prompt(context, schema)
input_ids = tokenizer.encode(prompt, return_tensors="pt").to(device)
prompt_tokens = input_ids.shape[1]
with torch.no_grad():
output_ids = model.generate(
input_ids,
max_new_tokens=max_new_tokens,
do_sample=(temperature > 0.0),
temperature=max(temperature, 1e-4),
pad_token_id=tokenizer.pad_token_id or tokenizer.eos_token_id
)
elapsed_ms = (time.perf_counter() - t0) * 1000
gen_tokens = output_ids.shape[1] - prompt_tokens
tok_per_sec = (gen_tokens / (elapsed_ms / 1000.0)) if elapsed_ms > 0 else 0.0
raw_text = tokenizer.decode(output_ids[0][prompt_tokens:], skip_special_tokens=True)
# Parse JSON
parsed_json = None
is_valid = False
try:
first_brace = raw_text.find("{")
last_brace = raw_text.rfind("}")
if first_brace != -1 and last_brace != -1:
cleaned = raw_text[first_brace:last_brace + 1]
parsed_json = json.loads(cleaned)
is_valid = True
except Exception:
pass
schema_match = False
if is_valid and isinstance(parsed_json, dict):
expected_keys = set(schema.get_field_names())
schema_match = (set(parsed_json.keys()) == expected_keys)
return {
"mode": "autoregressive_naive",
"elapsed_ms": round(elapsed_ms, 2),
"total_tokens": gen_tokens,
"tokens_per_second": round(tok_per_sec, 1),
"sequential_forward_passes": gen_tokens,
"is_valid_json": is_valid,
"schema_match": schema_match,
"raw_text": raw_text,
"parsed_json": parsed_json,
"device": device
}
def stream_naive_generation_torch(
context: str,
schema: StructuredSchema,
temperature: float = 0.2,
max_new_tokens: int = 512
) -> Generator[Dict[str, Any], None, None]:
"""
Generator streaming individual tokens for side-by-side comparison visualizer.
"""
model, tokenizer, device = get_torch_engine()
t0 = time.perf_counter()
prompt = build_naive_json_prompt(context, schema)
input_ids = tokenizer.encode(prompt, return_tensors="pt").to(device)
prompt_tokens = input_ids.shape[1]
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
gen_kwargs = {
"input_ids": input_ids,
"max_new_tokens": max_new_tokens,
"do_sample": (temperature > 0.0),
"temperature": max(temperature, 1e-4),
"pad_token_id": tokenizer.pad_token_id or tokenizer.eos_token_id,
"streamer": streamer
}
thread = threading.Thread(target=model.generate, kwargs=gen_kwargs)
thread.start()
full_text = ""
tok_count = 0
for token_str in streamer:
tok_count += 1
full_text += token_str
yield {
"type": "token",
"token": token_str,
"token_count": tok_count
}
thread.join()
elapsed_ms = (time.perf_counter() - t0) * 1000
parsed_json = None
is_valid = False
try:
first_brace = full_text.find("{")
last_brace = full_text.rfind("}")
if first_brace != -1 and last_brace != -1:
cleaned = full_text[first_brace:last_brace + 1]
parsed_json = json.loads(cleaned)
is_valid = True
except Exception:
pass
schema_match = False
if is_valid and isinstance(parsed_json, dict):
expected_keys = set(schema.get_field_names())
schema_match = (set(parsed_json.keys()) == expected_keys)
result = {
"mode": "autoregressive_naive",
"elapsed_ms": round(elapsed_ms, 2),
"total_tokens": tok_count,
"tokens_per_second": round((tok_count / (elapsed_ms / 1000.0)) if elapsed_ms > 0 else 0.0, 1),
"sequential_forward_passes": tok_count,
"is_valid_json": is_valid,
"schema_match": schema_match,
"raw_text": full_text,
"parsed_json": parsed_json,
"device": device
}
yield {
"type": "done",
"result": result
}