from __future__ import annotations import base64 import json import hashlib import os import re import tempfile import threading import time import traceback import types import urllib.parse from datetime import datetime, timezone from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from typing import Any import torch from peft import PeftModel, get_peft_model_state_dict from peft.tuners.lora.layer import LoraLayer from safetensors import safe_open from transformers import AutoModelForMultimodalLM, AutoProcessor from transformers.video_utils import load_video MODEL_DIR = Path(os.environ["GEMMA4_MODEL_DIR"]) ADAPTER_DIR = Path(os.environ["GEMMA4_ADAPTER_DIR"]) TEMPLATE_PATH = Path(os.environ["GEMMA4_TEMPLATE_PATH"]) STATE_PATH = Path(os.environ["GEMMA4_STATE_PATH"]) HOST = os.environ.get("GEMMA4_HOST", "127.0.0.1") PORT = int(os.environ.get("GEMMA4_PORT", "8091")) ALIAS = os.environ.get( "GEMMA4_ALIAS", "anru-human10mb-selected-native-multimodal-128k-mix045" ) LORA_SCALE = float(os.environ.get("GEMMA4_LORA_SCALE", "0.45")) MAX_CONTEXT = int(os.environ.get("GEMMA4_MAX_CONTEXT", "131072")) MAX_IMAGE_TOKENS = int(os.environ.get("GEMMA4_MAX_IMAGE_TOKENS", "1120")) MAX_GPU_MEMORY = os.environ.get("GEMMA4_MAX_GPU_MEMORY", "14GiB") MAX_CPU_MEMORY = os.environ.get("GEMMA4_MAX_CPU_MEMORY", "48GiB") OFFLOAD_DIR = Path(os.environ["GEMMA4_OFFLOAD_DIR"]) ATTN_IMPLEMENTATION = os.environ.get("GEMMA4_ATTN_IMPLEMENTATION", "sdpa") SYSTEM_PROMPT_PATH = Path( os.environ.get("GEMMA4_SYSTEM_PROMPT_PATH", str(Path(__file__).with_name("system_prompt.md"))) ) for required in (MODEL_DIR, ADAPTER_DIR, TEMPLATE_PATH, SYSTEM_PROMPT_PATH): if not required.exists(): raise FileNotFoundError(required) OFFLOAD_DIR.mkdir(parents=True, exist_ok=True) STATE_PATH.parent.mkdir(parents=True, exist_ok=True) DEFAULT_SYSTEM_PROMPT = SYSTEM_PROMPT_PATH.read_text(encoding="utf-8-sig").strip() if not DEFAULT_SYSTEM_PROMPT: raise ValueError(f"system prompt is empty: {SYSTEM_PROMPT_PATH}") SYSTEM_PROMPT_BYTES = SYSTEM_PROMPT_PATH.read_bytes() SYSTEM_PROMPT_SHA256 = hashlib.sha256(SYSTEM_PROMPT_BYTES).hexdigest().upper() SYSTEM_PROMPT_LINE_COUNT = len(DEFAULT_SYSTEM_PROMPT.splitlines()) started_at = time.perf_counter() print(f"[{datetime.now(timezone.utc).isoformat()}] loading processor", flush=True) processor = AutoProcessor.from_pretrained(MODEL_DIR, local_files_only=True) processor.chat_template = TEMPLATE_PATH.read_text(encoding="utf-8") def _fetch_videos_pyav(self, video_url_or_urls, sample_indices_fn=None): if isinstance(video_url_or_urls, list): return list( zip(*[_fetch_videos_pyav(self, item, sample_indices_fn=sample_indices_fn) for item in video_url_or_urls]) ) return load_video(video_url_or_urls, backend="pyav", sample_indices_fn=sample_indices_fn) if getattr(processor, "video_processor", None) is not None: processor.video_processor.fetch_videos = types.MethodType(_fetch_videos_pyav, processor.video_processor) print(f"[{datetime.now(timezone.utc).isoformat()}] loading Gemma4 multimodal weights", flush=True) model = AutoModelForMultimodalLM.from_pretrained( MODEL_DIR, dtype=torch.bfloat16, device_map="auto", max_memory={0: MAX_GPU_MEMORY, "cpu": MAX_CPU_MEMORY}, offload_folder=str(OFFLOAD_DIR), offload_buffers=True, low_cpu_mem_usage=True, local_files_only=True, attn_implementation=ATTN_IMPLEMENTATION, ) print(f"[{datetime.now(timezone.utc).isoformat()}] attaching LoRA adapter", flush=True) model = PeftModel.from_pretrained( model, ADAPTER_DIR, is_trainable=False, local_files_only=True, key_mapping={r"^model\.": "model.language_model."}, device_map="auto", max_memory={0: MAX_GPU_MEMORY, "cpu": MAX_CPU_MEMORY}, offload_folder=str(OFFLOAD_DIR), ) lora_modules = 0 for module in model.modules(): if isinstance(module, LoraLayer) and "default" in module.active_adapters: module.scale_layer(LORA_SCALE) lora_modules += 1 if not lora_modules: raise RuntimeError("LoRA adapter attached zero modules") adapter_state = get_peft_model_state_dict(model, adapter_name="default") adapter_tensor_count = len(adapter_state) adapter_file = ADAPTER_DIR / "adapter_model.safetensors" with safe_open(adapter_file, framework="pt", device="cpu") as adapter_reader: source_adapter_tensor_count = len(adapter_reader.keys()) adapter_nonzero_tensors = sum( int(torch.count_nonzero(adapter_reader.get_tensor(key)).item() > 0) for key in adapter_reader.keys() ) if adapter_tensor_count == 0 or adapter_nonzero_tensors == 0: raise RuntimeError("LoRA adapter state is empty") if adapter_tensor_count != source_adapter_tensor_count: raise RuntimeError( f"LoRA tensor count mismatch: runtime={adapter_tensor_count}, source={source_adapter_tensor_count}" ) model.eval() generation_lock = threading.Lock() loaded_seconds = round(time.perf_counter() - started_at, 3) def _json_bytes(value: Any) -> bytes: return json.dumps(value, ensure_ascii=False).encode("utf-8") def _data_url_to_file(url: str, temp_paths: list[Path]) -> str: match = re.fullmatch(r"data:([^;,]+)?(?:;charset=[^;,]+)?;base64,(.+)", url, re.DOTALL) if not match: return url mime = match.group(1) or "application/octet-stream" suffixes = { "image/jpeg": ".jpg", "image/png": ".png", "image/webp": ".webp", "audio/wav": ".wav", "audio/mpeg": ".mp3", "video/mp4": ".mp4", } suffix = suffixes.get(mime, ".bin") fd, name = tempfile.mkstemp(prefix="gemma4-input-", suffix=suffix) os.close(fd) path = Path(name) path.write_bytes(base64.b64decode(match.group(2), validate=True)) temp_paths.append(path) return str(path) def _normalise_url(value: Any, temp_paths: list[Path]) -> str: if isinstance(value, dict): value = value.get("url") if not isinstance(value, str) or not value: raise ValueError("multimodal content is missing its URL/data") if value.startswith("data:"): return _data_url_to_file(value, temp_paths) if value.startswith("file://"): return urllib.parse.unquote(urllib.parse.urlparse(value).path.lstrip("/") if os.name == "nt" else urllib.parse.urlparse(value).path) return value def _normalise_messages( messages: Any, temp_paths: list[Path], use_default_system_prompt: bool ) -> list[dict[str, Any]]: if not isinstance(messages, list) or not messages: raise ValueError("messages must be a non-empty list") result: list[dict[str, Any]] = [] for message in messages: role = message.get("role") content = message.get("content", "") if role not in {"system", "user", "assistant"}: raise ValueError(f"unsupported role: {role!r}") if isinstance(content, str): result.append({"role": role, "content": content}) continue if not isinstance(content, list): raise ValueError("message content must be text or a content-part list") parts: list[dict[str, Any]] = [] for part in content: kind = part.get("type") if kind in {"text", "input_text"}: parts.append({"type": "text", "text": str(part.get("text", ""))}) elif kind in {"image", "image_url", "input_image"}: source = part.get("url", part.get("image_url", part.get("image"))) parts.append({"type": "image", "url": _normalise_url(source, temp_paths)}) elif kind in {"audio", "audio_url", "input_audio"}: source = part.get("url", part.get("audio_url", part.get("audio"))) if source is None and isinstance(part.get("input_audio"), dict): audio = part["input_audio"] fmt = audio.get("format", "wav") source = f"data:audio/{fmt};base64,{audio.get('data', '')}" parts.append({"type": "audio", "url": _normalise_url(source, temp_paths)}) elif kind in {"video", "video_url", "input_video"}: source = part.get("url", part.get("video_url", part.get("video"))) parts.append({"type": "video", "url": _normalise_url(source, temp_paths)}) else: raise ValueError(f"unsupported content part: {kind!r}") result.append({"role": role, "content": parts}) if use_default_system_prompt: if result and result[0]["role"] == "system": existing = result[0]["content"] if isinstance(existing, str): merged = DEFAULT_SYSTEM_PROMPT + "\n\n---\n\n" + existing else: merged = [{"type": "text", "text": DEFAULT_SYSTEM_PROMPT}, *existing] result[0] = {"role": "system", "content": merged} else: result.insert(0, {"role": "system", "content": DEFAULT_SYSTEM_PROMPT}) return result def _complete(payload: dict[str, Any]) -> dict[str, Any]: temp_paths: list[Path] = [] began = time.perf_counter() try: use_default_system_prompt = payload.get("use_default_system_prompt", True) if not isinstance(use_default_system_prompt, bool): raise ValueError("use_default_system_prompt must be a JSON boolean") messages = _normalise_messages( payload.get("messages"), temp_paths, use_default_system_prompt ) enable_thinking = bool(payload.get("enable_thinking", False)) inputs = processor.apply_chat_template( messages, tokenize=True, return_dict=True, return_tensors="pt", add_generation_prompt=True, enable_thinking=enable_thinking, processor_kwargs={"images_kwargs": {"max_soft_tokens": MAX_IMAGE_TOKENS}}, ) prompt_tokens = int(inputs["input_ids"].shape[-1]) max_new_tokens = max(1, min(int(payload.get("max_tokens", payload.get("max_completion_tokens", 256))), 4096)) if prompt_tokens + max_new_tokens > MAX_CONTEXT: raise ValueError( f"requested {prompt_tokens + max_new_tokens} tokens exceeds {MAX_CONTEXT}-token service context" ) temperature = float(payload.get("temperature", 0.0)) top_p = float(payload.get("top_p", 0.95)) inputs = inputs.to(model.device) generate_args: dict[str, Any] = { "max_new_tokens": max_new_tokens, "do_sample": temperature > 0.0, "use_cache": True, } if temperature > 0.0: generate_args.update(temperature=temperature, top_p=top_p) with generation_lock, torch.inference_mode(): generated = model.generate(**inputs, **generate_args) completion_tokens = int(generated.shape[-1] - prompt_tokens) text = processor.decode(generated[0][prompt_tokens:], skip_special_tokens=True) for stop in payload.get("stop", []) if isinstance(payload.get("stop", []), list) else [payload.get("stop")]: if stop and stop in text: text = text.split(stop, 1)[0] elapsed = round(time.perf_counter() - began, 3) return { "id": f"chatcmpl-native-{int(time.time() * 1000)}", "object": "chat.completion", "created": int(time.time()), "model": ALIAS, "choices": [ { "index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop" if completion_tokens < max_new_tokens else "length", } ], "usage": { "prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, "total_tokens": prompt_tokens + completion_tokens, }, "timing": {"elapsed_seconds": elapsed}, "system_prompt": { "default_injected": use_default_system_prompt, "sha256": SYSTEM_PROMPT_SHA256 if use_default_system_prompt else None, }, } finally: for path in temp_paths: try: path.unlink() except OSError: pass health = { "status": "ok", "backend": "transformers-native", "model": ALIAS, "modalities": ["text", "image", "audio", "video"], "context_tokens": MAX_CONTEXT, "attention_implementation": ATTN_IMPLEMENTATION, "default_system_prompt": True, "system_prompt_file": SYSTEM_PROMPT_PATH.name, "system_prompt_sha256": SYSTEM_PROMPT_SHA256, "system_prompt_bytes": len(SYSTEM_PROMPT_BYTES), "system_prompt_chars": len(DEFAULT_SYSTEM_PROMPT), "system_prompt_lines": SYSTEM_PROMPT_LINE_COUNT, "image_max_soft_tokens": MAX_IMAGE_TOKENS, "lora_scale": LORA_SCALE, "lora_modules": lora_modules, "adapter_tensor_count": adapter_tensor_count, "adapter_nonzero_tensors": adapter_nonzero_tensors, "load_seconds": loaded_seconds, "cuda_available": torch.cuda.is_available(), "cuda_allocated_bytes": torch.cuda.memory_allocated() if torch.cuda.is_available() else 0, "device_map": getattr(model, "hf_device_map", None), } STATE_PATH.write_text(json.dumps(health, ensure_ascii=False, indent=2), encoding="utf-8") print(json.dumps(health, ensure_ascii=False), flush=True) class Handler(BaseHTTPRequestHandler): server_version = "Gemma4Native/1.0" def _send(self, status: int, value: Any) -> None: body = _json_bytes(value) self.send_response(status) self.send_header("Content-Type", "application/json; charset=utf-8") self.send_header("Content-Length", str(len(body))) self.end_headers() self.wfile.write(body) def do_GET(self) -> None: # noqa: N802 path = urllib.parse.urlparse(self.path).path if path == "/health": self._send(200, health) elif path == "/v1/models": self._send(200, {"object": "list", "data": [{"id": ALIAS, "object": "model", "owned_by": "local"}]}) elif path == "/props": self._send( 200, { "model_alias": ALIAS, "total_slots": 1, "default_generation_settings": {"n_ctx": MAX_CONTEXT}, "modalities": health["modalities"], "backend": health["backend"], }, ) elif path == "/lora-adapters": self._send(200, [{"id": 0, "scale": LORA_SCALE, "path": str(ADAPTER_DIR)}]) else: self._send(404, {"error": {"message": "not found", "type": "invalid_request_error"}}) def do_POST(self) -> None: # noqa: N802 path = urllib.parse.urlparse(self.path).path if path != "/v1/chat/completions": self._send(404, {"error": {"message": "not found", "type": "invalid_request_error"}}) return try: length = int(self.headers.get("Content-Length", "0")) payload = json.loads(self.rfile.read(length)) if payload.get("stream"): raise ValueError("streaming is not enabled on this local backend") self._send(200, _complete(payload)) except ValueError as exc: self._send(400, {"error": {"message": str(exc), "type": "invalid_request_error"}}) except Exception as exc: traceback.print_exc() self._send(500, {"error": {"message": repr(exc), "type": "server_error"}}) def log_message(self, fmt: str, *args: Any) -> None: print(f"[{datetime.now(timezone.utc).isoformat()}] {self.client_address[0]} {fmt % args}", flush=True) print(f"[{datetime.now(timezone.utc).isoformat()}] listening on http://{HOST}:{PORT}", flush=True) ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()