"""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 ``. 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 \ /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()