Download sdg/inference.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 6.27 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/inference.py
- Command line
-
hf download hf://fzzhang/svd-code/sdg/inference.py
-
curl -L -o inference.py https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/inference.py
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] | |