""" 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]