Anrn-12B-R1 / serve.py
chirs345678's picture
Add files using upload-large-folder tool
4e6526a verified
Raw
History Blame Contribute Delete
16 kB
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()