svd-code / sdg /inference.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
6.27 kB
"""
vLLM wrapper for generation and validation inference.
Pattern follows filter_with_uq.py β€” single LLM instance shared for both
generation (n=num_generations) and validation (n=num_validation_votes).
Inference results are cached per-prompt in MongoDB so that interrupted runs
can resume without re-doing expensive vLLM calls, and multiple shards can
share a single cache.
"""
from __future__ import annotations
import hashlib
import json
from pymongo import MongoClient
from vllm import LLM, SamplingParams
from sdg.config import SDGConfig
class MongoCache:
"""Dict-like interface over a MongoDB collection for inference caching."""
def __init__(self, uri: str, db_name: str):
self.client = MongoClient(uri)
self.db = self.client[db_name]
self.collection = self.db["inference_cache"]
self.collection.create_index("key", unique=True)
# Fail fast if MongoDB is unreachable
self.client.admin.command("ping")
def get(self, key: str, default=None):
doc = self.collection.find_one({"key": key})
if doc is not None:
return doc["value"]
return default
def __setitem__(self, key: str, value):
self.collection.update_one(
{"key": key},
{"$set": {"key": key, "value": value}},
upsert=True,
)
class VLLMEngine:
"""Thin wrapper around vllm.LLM with two pre-built SamplingParams."""
def __init__(self, config: SDGConfig):
print(f"Loading model: {config.model_name}")
self.model_name = config.model_name
self.llm = LLM(
model=config.model_name,
tensor_parallel_size=config.tensor_parallel_size,
max_model_len=config.max_model_len,
gpu_memory_utilization=config.gpu_memory_utilization,
trust_remote_code=True,
)
self.gen_params = SamplingParams(
temperature=config.gen_temperature,
top_p=config.gen_top_p,
max_tokens=config.gen_max_tokens,
n=config.num_generations,
)
if config.num_validation_votes > 0:
self.val_params = SamplingParams(
temperature=config.val_temperature,
top_p=config.val_top_p,
max_tokens=config.val_max_tokens,
n=config.num_validation_votes,
)
else:
self.val_params = None
# ── inference cache ──────────────────────────────────────────────
db_name = config.mongo_db_name
print(f"Connecting to MongoDB: {config.mongo_uri} db={db_name}")
self.cache = MongoCache(config.mongo_uri, db_name)
print("MongoDB cache connected successfully")
print("Model loaded successfully")
# ── helpers ──────────────────────────────────────────────────────────
def _cache_key(self, prompt: str, params: SamplingParams) -> str:
raw = json.dumps(
[self.model_name, params.temperature, params.top_p,
params.max_tokens, params.n, prompt],
ensure_ascii=False,
)
return hashlib.sha256(raw.encode()).hexdigest()
def apply_chat_template(self, prompt: str) -> str:
tokenizer = self.llm.get_tokenizer()
messages = [{"role": "user", "content": prompt}]
return tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
def _generate_cached(
self, prompts: list[str], params: SamplingParams,
) -> list[list[str]]:
"""Generate with per-prompt caching. Returns list[list[str]]."""
templated = [self.apply_chat_template(p) for p in prompts]
# Split into cached hits vs misses
results: list[list[str] | None] = [None] * len(prompts)
miss_indices: list[int] = []
miss_templated: list[str] = []
cache_hits = 0
for i, (prompt, tmpl) in enumerate(zip(prompts, templated)):
key = self._cache_key(tmpl, params)
cached = self.cache.get(key)
if cached is not None:
results[i] = cached
cache_hits += 1
else:
miss_indices.append(i)
miss_templated.append(tmpl)
# Run vLLM only for misses
if miss_templated:
outputs = self.llm.generate(miss_templated, params)
assert len(outputs) == len(miss_indices), (
f"vLLM returned {len(outputs)} outputs for {len(miss_indices)} prompts"
)
for idx, tmpl, out in zip(miss_indices, miss_templated, outputs):
texts = [o.text for o in out.outputs]
key = self._cache_key(tmpl, params)
self.cache[key] = texts
results[idx] = texts
if cache_hits:
print(f" Cache: {cache_hits} hits, {len(miss_indices)} misses")
return results # type: ignore[return-value]
# ── public API ───────────────────────────────────────────────────────
def generate_multi(self, prompts: list[str]) -> list[list[str]]:
"""Generate *num_generations* samples per prompt (candidate generation)."""
return self._generate_cached(prompts, self.gen_params)
def generate_with_votes(self, prompts: list[str]) -> list[list[str]]:
"""Generate *num_validation_votes* samples per prompt (UQ voting)."""
return self._generate_cached(prompts, self.val_params)
def generate_single(self, prompts: list[str]) -> list[str]:
"""Generate one sample per prompt (e.g. inferred-question generation)."""
single_params = SamplingParams(
temperature=self.val_params.temperature,
top_p=self.val_params.top_p,
max_tokens=self.val_params.max_tokens,
n=1,
)
results = self._generate_cached(prompts, single_params)
return [r[0] for r in results]