Skip to content

Commit 92571ba

Browse files
Create vllm_engine.py
1 parent c860c13 commit 92571ba

1 file changed

Lines changed: 92 additions & 0 deletions

File tree

prototype/vllm_engine.py

Lines changed: 92 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,92 @@
1+
# prototype/vllm_engine.py (NEW)
2+
from typing import List, Optional
3+
import torch
4+
from transformers import AutoTokenizer
5+
6+
class VllmInferenceEngine:
7+
"""Production inference engine using vLLM for optimal GPU utilization."""
8+
9+
def __init__(
10+
self,
11+
model_id: str,
12+
gpu_memory_fraction: float = 0.8,
13+
max_num_seqs: int = 64,
14+
tensor_parallel_size: int = 1,
15+
enforce_eager: bool = False
16+
):
17+
self.model_id = model_id
18+
self.gpu_memory_fraction = gpu_memory_fraction
19+
self.max_num_seqs = max_num_seqs
20+
self.tensor_parallel_size = tensor_parallel_size
21+
22+
import vllm
23+
24+
# Initialize vLLM engine with optimal settings
25+
self.engine = vllm.LLM(
26+
model=model_id,
27+
gpu_memory_fraction=self.gpu_memory_fraction,
28+
max_num_seqs=self.max_num_seqs,
29+
tensor_parallel_size=self.tensor_parallel_size,
30+
enforce_eager=enforce_eager, # For debugging/control
31+
enable_prefix_caching=True, # Performance optimization
32+
)
33+
34+
self.tokenizer = AutoTokenizer.from_pretrained(model_id)
35+
36+
def generate(
37+
self,
38+
prompts: List[str],
39+
max_tokens: int = 256,
40+
temperature: float = 1.0,
41+
top_p: float = 0.9,
42+
stop: Optional[List[str]] = None
43+
) -> List[str]:
44+
"""Batched generation with vLLM."""
45+
46+
# Tokenize prompts
47+
input_ids = self.tokenizer(
48+
prompts,
49+
return_tensors="pt",
50+
padding=True
51+
).to(self.engine.device)
52+
53+
# Generate using vLLM's optimized path
54+
outputs = self.engine.generate(
55+
**input_ids,
56+
max_tokens=max_tokens,
57+
temperature=temperature,
58+
top_p=top_p,
59+
stop=stop
60+
)
61+
62+
# Decode and return
63+
decoded = self.tokenizer.batch_decode(outputs)
64+
return decoded
65+
66+
def tokenized_generate(
67+
self,
68+
input_ids: torch.Tensor,
69+
attention_mask: Optional[torch.Tensor] = None,
70+
max_tokens: int = 256
71+
) -> List[str]:
72+
"""Generate from pre-tokenized input (for pipeline parallelism)."""
73+
74+
outputs = self.engine.generate(
75+
input_ids=input_ids,
76+
attention_mask=attention_mask,
77+
max_tokens=max_tokens
78+
)
79+
80+
decoded = self.tokenizer.batch_decode(outputs)
81+
return decoded
82+
83+
def get_model_profile(self) -> dict:
84+
"""Get vLLM engine profile for monitoring."""
85+
profile = {
86+
"num_active_requests": len(self.engine.get_cache_config()),
87+
"cache_hit_rate": self.engine.get_cache_config().cache_hit_rate if hasattr(self.engine, 'get_cache_config') else None,
88+
"gpu_memory_utilization": self.engine.get_gpu_memory_utilization(),
89+
"num_tokens_seen": self.engine.get_num_prompt_tokens_seen() + self.engine.get_num_generation_tokens_seen(),
90+
}
91+
92+
return profile

0 commit comments

Comments
 (0)