pactbench / pact /gemma_server.py
BBoran's picture
Publish current portable PACTBench release
f1fc3a0 verified
Raw
History Blame Contribute Delete
7.58 kB
"""Minimal OpenAI-compatible server for the PACT L2 decision LLM (default: google/gemma-4-E2B-it).
Run in an env whose transformers supports the model. gemma-4-* needs transformers>=5 -> use the
dedicated venv `<gemma4-venv>`. Loads the
model ONCE, resident, and serves POST /v1/chat/completions matching what
pact/eval_l2_decision.py::call_openai_compatible expects. Runs as a SEPARATE process from the
closed loop (GameMaster 5B in its own env); they talk over HTTP so VRAM/deps don't collide.
CUDA_VISIBLE_DEVICES=6 GEMMA_PORT=8000 \
<gemma4-venv>/bin/python pact/gemma_server.py
Then point the loop / eval at it:
export GEMMA_OPENAI_BASE_URL=http://127.0.0.1:8000/v1
export GEMMA_MODEL=gemma-4-E2B-it
"""
import base64, json, os, sys, time
from io import BytesIO
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
import torch
from PIL import Image
from transformers import AutoConfig, AutoProcessor
MODEL_ID = os.environ.get("GEMMA_MODEL_PATH", "google/gemma-4-E2B-it")
MODEL_REVISION = MODEL_ID.rstrip("/").split("/")[-1] if "/snapshots/" in MODEL_ID else None
PORT = int(os.environ.get("GEMMA_PORT", "8000"))
LOCAL_ONLY = os.environ.get("HF_HUB_OFFLINE") == "1" or os.environ.get("TRANSFORMERS_OFFLINE") == "1"
print(f"[gemma_server] loading {MODEL_ID} ...", flush=True)
_proc = AutoProcessor.from_pretrained(MODEL_ID, local_files_only=LOCAL_ONLY)
_mtype = AutoConfig.from_pretrained(MODEL_ID, local_files_only=LOCAL_ONLY).model_type
if _mtype == "gemma4":
from transformers import Gemma4ForConditionalGeneration as _CLS
elif _mtype == "gemma3n":
from transformers import Gemma3nForConditionalGeneration as _CLS
elif _mtype == "gemma3":
from transformers import Gemma3ForConditionalGeneration as _CLS
else:
raise ValueError(
f"gemma_server expected model_type in {{'gemma4','gemma3n','gemma3'}}, got {_mtype!r} from {MODEL_ID}")
_model = _CLS.from_pretrained(
MODEL_ID, dtype=torch.bfloat16, device_map="cuda:0", local_files_only=LOCAL_ONLY).eval()
print(f"[gemma_server] loaded ({_mtype}).", flush=True)
def _to_text_content(c):
if isinstance(c, str):
return c
if isinstance(c, list): # OpenAI multimodal list -> concat text parts
return "".join(p.get("text", "") for p in c if isinstance(p, dict) and p.get("type") == "text")
return str(c)
def _image_from_url(url):
if not isinstance(url, str):
raise ValueError("image_url.url must be a string")
if url.startswith("data:image/"):
_, b64 = url.split(",", 1)
return Image.open(BytesIO(base64.b64decode(b64))).convert("RGB")
if url.startswith("file://"):
return Image.open(url[7:]).convert("RGB")
if os.path.exists(url):
return Image.open(url).convert("RGB")
raise ValueError("only data:image, file://, or local image paths are supported")
def _to_gemma_content(c):
if isinstance(c, str):
return [{"type": "text", "text": c}]
if not isinstance(c, list):
return [{"type": "text", "text": str(c)}]
parts = []
for p in c:
if not isinstance(p, dict):
parts.append({"type": "text", "text": str(p)})
continue
typ = p.get("type")
if typ == "text":
parts.append({"type": "text", "text": str(p.get("text", ""))})
elif typ == "image_url":
image_url = p.get("image_url", {})
url = image_url.get("url") if isinstance(image_url, dict) else image_url
parts.append({"type": "image", "image": _image_from_url(url)})
elif typ == "image":
image = p.get("image")
if isinstance(image, Image.Image):
parts.append({"type": "image", "image": image.convert("RGB")})
else:
parts.append({"type": "image", "image": _image_from_url(image)})
return parts
def generate(messages, temperature, max_new_tokens):
# Gemma chat template has NO system role -> fold system text into the first user turn.
sys_txt = "\n".join(_to_text_content(m["content"]) for m in messages if m.get("role") == "system")
turns = []
for m in messages:
if m.get("role") == "system":
continue
parts = _to_gemma_content(m["content"])
if m["role"] == "user" and sys_txt:
if parts and parts[0].get("type") == "text":
parts[0]["text"] = sys_txt + "\n\n" + parts[0].get("text", "")
else:
parts.insert(0, {"type": "text", "text": sys_txt})
sys_txt = ""
turns.append({"role": m["role"], "content": parts})
inputs = _proc.apply_chat_template(
turns, add_generation_prompt=True, tokenize=True,
return_dict=True, return_tensors="pt").to(_model.device)
in_len = inputs["input_ids"].shape[-1]
with torch.inference_mode():
out = _model.generate(
**inputs, max_new_tokens=max_new_tokens, do_sample=temperature > 0,
temperature=temperature if temperature > 0 else None,
pad_token_id=_proc.tokenizer.eos_token_id)
return _proc.tokenizer.decode(out[0][in_len:], skip_special_tokens=True).strip()
class H(BaseHTTPRequestHandler):
def log_message(self, *a): # quiet
pass
def _send(self, code, obj):
b = json.dumps(obj).encode()
try:
self.send_response(code)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(b)))
self.end_headers()
self.wfile.write(b)
except (BrokenPipeError, ConnectionResetError):
# The caller may time out while generation is still finishing.
# There is no client left to receive a secondary 500 response.
return
def do_GET(self):
if self.path.rstrip("/") in ("/health", "/v1/models"):
self._send(200, {
"status": "ok", "model": MODEL_ID, "model_id": MODEL_ID,
"model_revision": MODEL_REVISION, "model_type": _mtype,
})
else:
self._send(404, {"error": "not found"})
def do_POST(self):
if not self.path.rstrip("/").endswith("/chat/completions"):
self._send(404, {"error": "not found"})
return
try:
body = json.loads(self.rfile.read(int(self.headers.get("Content-Length", 0))))
requested = str(body.get("model", ""))
requested_key = "".join(c for c in requested.split("/")[-1].lower() if c.isalnum())
actual_key = "".join(c for c in MODEL_ID.lower() if c.isalnum())
if not requested_key or requested_key not in actual_key:
raise ValueError(f"requested model {requested!r}, but this server hosts {MODEL_ID!r}")
text = generate(body["messages"],
float(body.get("temperature", 0.0)),
int(body.get("max_tokens", 256)))
self._send(200, {"id": "cmpl", "object": "chat.completion", "created": int(time.time()),
"model": MODEL_ID,
"choices": [{"index": 0, "finish_reason": "stop",
"message": {"role": "assistant", "content": text}}]})
except Exception as e:
self._send(500, {"error": f"{type(e).__name__}: {e}"})
if __name__ == "__main__":
print(f"[gemma_server] serving on http://127.0.0.1:{PORT}/v1 (POST /v1/chat/completions)", flush=True)
ThreadingHTTPServer(("127.0.0.1", PORT), H).serve_forever()