#!/usr/bin/env python3 from __future__ import annotations import json import math import os import queue import subprocess import tempfile import threading import time import uuid from collections import deque from pathlib import Path from typing import Any, Generator, Iterable from fastapi import FastAPI, Header, HTTPException, Request from fastapi.responses import JSONResponse, StreamingResponse from pydantic import BaseModel, Field import uvicorn def load_dotenv(path: str | Path = ".env") -> None: p = Path(path) if not p.exists(): return for raw in p.read_text(encoding="utf-8").splitlines(): line = raw.strip() if not line or line.startswith("#") or "=" not in line: continue key, val = line.split("=", 1) key = key.strip() val = os.path.expandvars(val.strip().strip('"').strip("'")) os.environ.setdefault(key, val) load_dotenv() APP_ROOT = Path(__file__).resolve().parent DG_RUNTIME_PRESET = os.getenv("DG_RUNTIME_PRESET", "").strip().lower().replace("-", "_") MODEL_ID = os.getenv("DG_MODEL_ID", "diffusiongemma-26b-a4b-it-q4-local") API_KEY = os.getenv("DG_API_KEY", "").strip() VISUAL_SERVER_BIN = os.getenv("DG_VISUAL_SERVER_BIN", "").strip() MODEL_PATH = os.getenv("DG_MODEL", "").strip() DG_NGL_REQUESTED = int(os.getenv("DG_NGL", "18")) DG_ALLOW_FULL_NGL = int(os.getenv("DG_ALLOW_FULL_NGL", "0")) DG_MAXTOK = int(os.getenv("DG_MAXTOK", "2048")) DG_FA = int(os.getenv("DG_FA", "1")) DG_FUSED_MMQ_GLU = int(os.getenv("DG_FUSED_MMQ_GLU", "0")) DG_FUSED_MOE_DOWN_REDUCE = int(os.getenv("DG_FUSED_MOE_DOWN_REDUCE", "1")) DG_FUSED_MOE_DOWN_REDUCE_TEMP = int(os.getenv("DG_FUSED_MOE_DOWN_REDUCE_TEMP", "1")) DG_FUSED_MOE_DOWN_REDUCE_SHARED_FACTORS = int(os.getenv("DG_FUSED_MOE_DOWN_REDUCE_SHARED_FACTORS", "1")) DG_MMQ_MAX_X = os.getenv("DG_MMQ_MAX_X", "").strip() DG_MMQ_X = os.getenv("DG_MMQ_X", "").strip() DG_NO_SC_EMBT = int(os.getenv("DG_NO_SC_EMBT", "1")) DG_SC_CHUNK = int(os.getenv("DG_SC_CHUNK", "4096")) DG_SC_PRECONVERT_PROBS = int(os.getenv("DG_SC_PRECONVERT_PROBS", "1")) DG_SC_PREDEQUANT_EMB = int(os.getenv("DG_SC_PREDEQUANT_EMB", "1")) DG_DEVSAMPLE_LAZY_SAMPLE = int(os.getenv("DG_DEVSAMPLE_LAZY_SAMPLE", "1")) DG_DIFFUSION_MASK_CACHE = int(os.getenv("DG_DIFFUSION_MASK_CACHE", "1")) DG_OFFLOAD_INPUT_ENV = os.getenv("DG_OFFLOAD_INPUT") DG_LAZY_RESERVE = int(os.getenv("DG_LAZY_RESERVE", "1")) DG_VISUAL_FRAMES = int(os.getenv("DG_VISUAL_FRAMES", "0")) DG_SAFE_COOLDOWN_SECONDS = float(os.getenv("DG_SAFE_COOLDOWN_SECONDS", "30")) DG_N_BLOCKS = int(os.getenv("DG_N_BLOCKS", "2")) DG_MAX_BLOCKS = int(os.getenv("DG_MAX_BLOCKS", "8")) DG_REQUEST_TIMEOUT = int(os.getenv("DG_REQUEST_TIMEOUT", "900")) DG_STARTUP_TIMEOUT = int(os.getenv("DG_STARTUP_TIMEOUT", "900")) HOST = os.getenv("DG_HOST", "127.0.0.1") PORT = int(os.getenv("DG_PORT", "8080")) if DG_RUNTIME_PRESET in {"q4_safe", "quality", "quality_speed"}: if not MODEL_PATH: MODEL_PATH = str(APP_ROOT / "models/diffusiongemma/diffusiongemma-26B-A4B-it-Q4_K_M.gguf") if "DG_NGL" not in os.environ: DG_NGL_REQUESTED = 18 if "DG_ALLOW_FULL_NGL" not in os.environ: DG_ALLOW_FULL_NGL = 0 elif DG_RUNTIME_PRESET in {"iq3_fullgpu_fast", "fullgpu_iq3", "fullgpu_fast"}: if "DG_MODEL_ID" not in os.environ: MODEL_ID = "diffusiongemma-26b-a4b-it-iq3m-fullgpu" if not MODEL_PATH: MODEL_PATH = str(APP_ROOT / "models/diffusiongemma/diffusiongemma-26B-A4B-it-IQ3_M-from-Q4_K_M.gguf") if "DG_NGL" not in os.environ: DG_NGL_REQUESTED = 999 if "DG_ALLOW_FULL_NGL" not in os.environ: DG_ALLOW_FULL_NGL = 1 if "DG_MAXTOK" not in os.environ: DG_MAXTOK = 768 if "DG_FUSED_MMQ_GLU" not in os.environ: DG_FUSED_MMQ_GLU = 0 elif DG_RUNTIME_PRESET in {"q4_fullgpu_small", "fullgpu_q4_small"}: if "DG_MODEL_ID" not in os.environ: MODEL_ID = "diffusiongemma-26b-a4b-it-q4-fullgpu-small" if not MODEL_PATH: MODEL_PATH = str(APP_ROOT / "models/diffusiongemma/diffusiongemma-26B-A4B-it-Q4_K_M.gguf") if "DG_NGL" not in os.environ: DG_NGL_REQUESTED = 999 if "DG_ALLOW_FULL_NGL" not in os.environ: DG_ALLOW_FULL_NGL = 1 if "DG_MAXTOK" not in os.environ: DG_MAXTOK = 288 if "DG_SC_CHUNK" not in os.environ: DG_SC_CHUNK = 2048 if "DG_FUSED_MMQ_GLU" not in os.environ: DG_FUSED_MMQ_GLU = 0 DG_NGL = DG_NGL_REQUESTED if DG_ALLOW_FULL_NGL else min(DG_NGL_REQUESTED, 18) DG_OFFLOAD_INPUT = int(DG_OFFLOAD_INPUT_ENV if DG_OFFLOAD_INPUT_ENV is not None else ("1" if DG_ALLOW_FULL_NGL else "0")) class ChatMessage(BaseModel): role: str content: Any = "" class ChatCompletionRequest(BaseModel): model: str | None = None messages: list[ChatMessage] max_tokens: int | None = Field(default=None, ge=1) stream: bool = False seed: int | None = None temperature: float | None = None top_p: float | None = None n_blocks: int | None = Field(default=None, ge=1) diffusion_blocks: int | None = Field(default=None, ge=1) class DiffusionGemmaRunner: def __init__(self, bin_path: str, model_path: str, ngl: int, maxtok: int, fa: int): self.bin_path = str(Path(bin_path).expanduser()) self.model_path = str(Path(model_path).expanduser()) self.ngl = ngl self.maxtok = maxtok self.fa = fa self.proc: subprocess.Popen[str] | None = None self.stdout_q: queue.Queue[str] = queue.Queue() self.lock = threading.Lock() self.stderr_tail: deque[str] = deque(maxlen=120) self.started_at: float | None = None self.last_request_finished_at = 0.0 self._tmpdir = tempfile.TemporaryDirectory(prefix="dg-api-") def _require_paths(self) -> None: if not self.bin_path or not Path(self.bin_path).exists(): raise RuntimeError(f"DG_VISUAL_SERVER_BIN does not exist: {self.bin_path}") if not self.model_path or not Path(self.model_path).exists(): raise RuntimeError(f"DG_MODEL does not exist: {self.model_path}") def _read_stdout(self) -> None: assert self.proc is not None and self.proc.stdout is not None for line in self.proc.stdout: self.stdout_q.put(line.rstrip("\n")) def _read_stderr(self) -> None: assert self.proc is not None and self.proc.stderr is not None for line in self.proc.stderr: clean = line.rstrip("\n") self.stderr_tail.append(clean) print(f"[dg-runner] {clean}", flush=True) def start(self) -> None: self._require_paths() if self.proc and self.proc.poll() is None: return env = os.environ.copy() env.update({"NGL": str(self.ngl), "MAXTOK": str(self.maxtok), "FA": str(self.fa)}) if DG_FUSED_MMQ_GLU: env.setdefault("GGML_CUDA_MERGED_GATE_UP_GLU_NOSTREAM", "1") if DG_FUSED_MOE_DOWN_REDUCE: env.setdefault("GGML_CUDA_MOE_DOWN_REDUCE", "1") if DG_FUSED_MOE_DOWN_REDUCE_TEMP: env.setdefault("GGML_CUDA_MOE_DOWN_REDUCE_TEMP", "1") if DG_FUSED_MOE_DOWN_REDUCE_SHARED_FACTORS: env.setdefault("GGML_CUDA_MOE_DOWN_REDUCE_SHARED_FACTORS", "1") if DG_MMQ_MAX_X: env.setdefault("GGML_CUDA_MMQ_MAX_X", DG_MMQ_MAX_X) if DG_MMQ_X: env.setdefault("GGML_CUDA_MMQ_X", DG_MMQ_X) if DG_NO_SC_EMBT: env.setdefault("DG_SC_EMBT_TRANSPOSE_VIEW", "1") env.setdefault("GGML_CUDA_DG_SC_CHUNK", str(DG_SC_CHUNK)) if DG_SC_PRECONVERT_PROBS: env.setdefault("GGML_CUDA_DG_SC_PRECONVERT_PROBS", "1") if DG_SC_PREDEQUANT_EMB: env.setdefault("GGML_CUDA_DG_SC_PREDEQUANT_EMB", "1") if DG_DEVSAMPLE_LAZY_SAMPLE: env.setdefault("DG_DEVSAMPLE_LAZY_SAMPLE", "1") if DG_DIFFUSION_MASK_CACHE: env.setdefault("DG_DIFFUSION_MASK_CACHE", "1") if DG_OFFLOAD_INPUT: env.setdefault("DG_OFFLOAD_INPUT", "1") if DG_LAZY_RESERVE: env.setdefault("LLAMA_CUDA_LAZY_RESERVE", "1") env.setdefault("DG_VISUAL_FRAMES", str(DG_VISUAL_FRAMES)) cmd = [self.bin_path, self.model_path] print(f"Starting DiffusionGemma runner: {' '.join(cmd)}", flush=True) print( "Runtime env: " f"NGL={self.ngl}, MAXTOK={self.maxtok}, FA={self.fa}, " f"GGML_CUDA_MERGED_GATE_UP_GLU_NOSTREAM={env.get('GGML_CUDA_MERGED_GATE_UP_GLU_NOSTREAM', '0')}, " f"GGML_CUDA_MOE_DOWN_REDUCE={env.get('GGML_CUDA_MOE_DOWN_REDUCE', '0')}, " f"GGML_CUDA_MOE_DOWN_REDUCE_TEMP={env.get('GGML_CUDA_MOE_DOWN_REDUCE_TEMP', '0')}, " f"GGML_CUDA_MOE_DOWN_REDUCE_SHARED_FACTORS={env.get('GGML_CUDA_MOE_DOWN_REDUCE_SHARED_FACTORS', '0')}, " f"GGML_CUDA_MMQ_MAX_X={env.get('GGML_CUDA_MMQ_MAX_X', '0')}, " f"GGML_CUDA_MMQ_X={env.get('GGML_CUDA_MMQ_X', '0')}, " f"DG_SC_EMBT_TRANSPOSE_VIEW={env.get('DG_SC_EMBT_TRANSPOSE_VIEW', '0')}, " f"GGML_CUDA_DG_SC_CHUNK={env.get('GGML_CUDA_DG_SC_CHUNK', '0')}, " f"GGML_CUDA_DG_SC_PRECONVERT_PROBS={env.get('GGML_CUDA_DG_SC_PRECONVERT_PROBS', '0')}, " f"GGML_CUDA_DG_SC_PREDEQUANT_EMB={env.get('GGML_CUDA_DG_SC_PREDEQUANT_EMB', '0')}, " f"DG_DEVSAMPLE_LAZY_SAMPLE={env.get('DG_DEVSAMPLE_LAZY_SAMPLE', '0')}, " f"DG_DIFFUSION_MASK_CACHE={env.get('DG_DIFFUSION_MASK_CACHE', '0')}, " f"DG_OFFLOAD_INPUT={env.get('DG_OFFLOAD_INPUT', '0')}, " f"LLAMA_CUDA_LAZY_RESERVE={env.get('LLAMA_CUDA_LAZY_RESERVE', '0')}, " f"DG_VISUAL_FRAMES={env.get('DG_VISUAL_FRAMES', '0')}", flush=True, ) self.proc = subprocess.Popen( cmd, stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, encoding="utf-8", errors="replace", bufsize=1, env=env, ) threading.Thread(target=self._read_stdout, daemon=True).start() threading.Thread(target=self._read_stderr, daemon=True).start() self.started_at = time.time() deadline = time.monotonic() + DG_STARTUP_TIMEOUT while time.monotonic() < deadline: if self.proc.poll() is not None: tail = "\n".join(self.stderr_tail) raise RuntimeError(f"runner exited during startup with code {self.proc.returncode}\n{tail}") try: line = self.stdout_q.get(timeout=1) except queue.Empty: continue print(f"[dg-runner:stdout] {line}", flush=True) if line.startswith("READY"): return if line.startswith("ERR"): raise RuntimeError(f"runner startup error: {line}") raise TimeoutError("runner startup timed out; model may still be loading or OOM-looping") def stop(self) -> None: if not self.proc: return if self.proc.poll() is None: try: assert self.proc.stdin is not None self.proc.stdin.write("QUIT\n") self.proc.stdin.flush() self.proc.wait(timeout=10) except Exception: self.proc.terminate() try: self.proc.wait(timeout=10) except Exception: self.proc.kill() self.proc = None def _ensure_alive(self) -> None: if not self.proc or self.proc.poll() is not None: self.start() def _write_request_file(self, payload: dict[str, Any]) -> str: fd, path = tempfile.mkstemp(prefix="dg-request-", suffix=".json", dir=self._tmpdir.name, text=True) with os.fdopen(fd, "w", encoding="utf-8") as f: json.dump(payload, f, ensure_ascii=False) return path def _send_request_path(self, path: str) -> None: assert self.proc is not None and self.proc.stdin is not None self.proc.stdin.write(path + "\n") self.proc.stdin.flush() def _apply_safe_cooldown(self) -> None: if DG_SAFE_COOLDOWN_SECONDS <= 0 or self.last_request_finished_at <= 0: return wait = DG_SAFE_COOLDOWN_SECONDS - (time.monotonic() - self.last_request_finished_at) if wait > 0: print(f"[dg-safe] cooling down for {wait:.1f}s before next request", flush=True) time.sleep(wait) @staticmethod def _parse_committed(line: str) -> str | None: # Protocol after PR commit 7ea238c: C if not line.startswith("C "): return None parts = line.split(" ", 2) if len(parts) < 3: return "" try: return json.loads(parts[2]) except json.JSONDecodeError: return parts[2] def generate(self, messages: list[dict[str, Any]], n_blocks: int, seed: int, timeout: int) -> str: with self.lock: self._apply_safe_cooldown() self._ensure_alive() payload = {"seed": seed, "n_blocks": n_blocks, "messages": messages} path = self._write_request_file(payload) answer = "" try: self._send_request_path(path) deadline = time.monotonic() + timeout while True: if self.proc and self.proc.poll() is not None: tail = "\n".join(self.stderr_tail) raise RuntimeError(f"runner exited with code {self.proc.returncode}\n{tail}") remaining = deadline - time.monotonic() if remaining <= 0: raise TimeoutError("request timed out") try: line = self.stdout_q.get(timeout=min(1.0, remaining)) except queue.Empty: continue if line == "DONE": self.last_request_finished_at = time.monotonic() return answer if line.startswith("ERR"): raise RuntimeError(line) committed = self._parse_committed(line) if committed is not None: answer = committed finally: try: os.remove(path) except OSError: pass def stream_generate(self, messages: list[dict[str, Any]], n_blocks: int, seed: int, timeout: int) -> Generator[str, None, None]: with self.lock: self._apply_safe_cooldown() self._ensure_alive() payload = {"seed": seed, "n_blocks": n_blocks, "messages": messages} path = self._write_request_file(payload) previous = "" try: self._send_request_path(path) deadline = time.monotonic() + timeout while True: if self.proc and self.proc.poll() is not None: tail = "\n".join(self.stderr_tail) raise RuntimeError(f"runner exited with code {self.proc.returncode}\n{tail}") remaining = deadline - time.monotonic() if remaining <= 0: raise TimeoutError("request timed out") try: line = self.stdout_q.get(timeout=min(1.0, remaining)) except queue.Empty: continue if line == "DONE": self.last_request_finished_at = time.monotonic() return if line.startswith("ERR"): raise RuntimeError(line) committed = self._parse_committed(line) if committed is not None: if committed.startswith(previous): delta = committed[len(previous):] else: delta = committed previous = committed if delta: yield delta finally: try: os.remove(path) except OSError: pass runner = DiffusionGemmaRunner(VISUAL_SERVER_BIN, MODEL_PATH, DG_NGL, DG_MAXTOK, DG_FA) app = FastAPI(title="DiffusionGemma Local OpenAI-Compatible API") def check_auth(authorization: str | None) -> None: if not API_KEY: return expected = f"Bearer {API_KEY}" if authorization != expected: raise HTTPException(status_code=401, detail="Invalid or missing API key") def normalize_messages(messages: Iterable[ChatMessage]) -> list[dict[str, str]]: normalized: list[dict[str, str]] = [] for msg in messages: role = msg.role content = msg.content if isinstance(content, list): parts: list[str] = [] for item in content: if isinstance(item, dict) and item.get("type") == "text": parts.append(str(item.get("text", ""))) elif isinstance(item, str): parts.append(item) content = "\n".join(part for part in parts if part) elif content is None: content = "" else: content = str(content) normalized.append({"role": role, "content": content}) return normalized def blocks_from_request(req: ChatCompletionRequest) -> int: explicit = req.n_blocks or req.diffusion_blocks if explicit: return max(1, min(DG_MAX_BLOCKS, explicit)) if req.max_tokens: return max(1, min(DG_MAX_BLOCKS, math.ceil(req.max_tokens / 256))) return max(1, min(DG_MAX_BLOCKS, DG_N_BLOCKS)) def completion_id() -> str: return "chatcmpl-" + uuid.uuid4().hex def sse(data: dict[str, Any] | str) -> str: if data == "[DONE]": return "data: [DONE]\n\n" return "data: " + json.dumps(data, ensure_ascii=False) + "\n\n" @app.on_event("startup") def on_startup() -> None: runner.start() @app.on_event("shutdown") def on_shutdown() -> None: runner.stop() @app.get("/healthz") def healthz() -> dict[str, Any]: alive = runner.proc is not None and runner.proc.poll() is None return { "ok": alive, "model": MODEL_ID, "runtime_preset": DG_RUNTIME_PRESET, "model_path": MODEL_PATH, "ngl_requested": DG_NGL_REQUESTED, "allow_full_ngl": DG_ALLOW_FULL_NGL, "ngl": DG_NGL, "maxtok": DG_MAXTOK, "fa": DG_FA, "fused_mmq_glu": DG_FUSED_MMQ_GLU, "fused_moe_down_reduce": DG_FUSED_MOE_DOWN_REDUCE, "fused_moe_down_reduce_temp": DG_FUSED_MOE_DOWN_REDUCE_TEMP, "fused_moe_down_reduce_shared_factors": DG_FUSED_MOE_DOWN_REDUCE_SHARED_FACTORS, "mmq_max_x": DG_MMQ_MAX_X, "mmq_x": DG_MMQ_X, "no_sc_embt": DG_NO_SC_EMBT, "sc_chunk": DG_SC_CHUNK, "sc_preconvert_probs": DG_SC_PRECONVERT_PROBS, "sc_predequant_emb": DG_SC_PREDEQUANT_EMB, "devsample_lazy_sample": DG_DEVSAMPLE_LAZY_SAMPLE, "diffusion_mask_cache": DG_DIFFUSION_MASK_CACHE, "offload_input": DG_OFFLOAD_INPUT, "lazy_reserve": DG_LAZY_RESERVE, "safe_cooldown_seconds": DG_SAFE_COOLDOWN_SECONDS, "started_at": runner.started_at, "stderr_tail": list(runner.stderr_tail)[-20:], } @app.get("/v1/models") def models(authorization: str | None = Header(default=None)) -> dict[str, Any]: check_auth(authorization) return { "object": "list", "data": [{"id": MODEL_ID, "object": "model", "created": 0, "owned_by": "local"}], } @app.post("/v1/chat/completions") def chat_completions(req: ChatCompletionRequest, authorization: str | None = Header(default=None)) -> Any: check_auth(authorization) messages = normalize_messages(req.messages) n_blocks = blocks_from_request(req) seed = req.seed if req.seed is not None else int(time.time_ns() % (2**31 - 1)) request_timeout = DG_REQUEST_TIMEOUT cid = completion_id() created = int(time.time()) model = req.model or MODEL_ID if req.stream: def event_stream() -> Generator[str, None, None]: try: for delta in runner.stream_generate(messages, n_blocks, seed, request_timeout): yield sse({ "id": cid, "object": "chat.completion.chunk", "created": created, "model": model, "choices": [{"index": 0, "delta": {"content": delta}, "finish_reason": None}], }) yield sse({ "id": cid, "object": "chat.completion.chunk", "created": created, "model": model, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], }) yield sse("[DONE]") except Exception as e: yield sse({"error": {"message": str(e), "type": "runner_error"}}) yield sse("[DONE]") return StreamingResponse(event_stream(), media_type="text/event-stream") try: answer = runner.generate(messages, n_blocks, seed, request_timeout) except Exception as e: raise HTTPException(status_code=500, detail=str(e)) from e return JSONResponse({ "id": cid, "object": "chat.completion", "created": created, "model": model, "choices": [{ "index": 0, "message": {"role": "assistant", "content": answer}, "finish_reason": "stop", }], "usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}, }) @app.post("/v1/completions") async def completions(request: Request, authorization: str | None = Header(default=None)) -> Any: # Minimal legacy endpoint: wraps prompt as a user message and delegates to chat completions. check_auth(authorization) body = await request.json() prompt = body.get("prompt", "") if isinstance(prompt, list): prompt = "\n".join(map(str, prompt)) req = ChatCompletionRequest( model=body.get("model") or MODEL_ID, messages=[ChatMessage(role="user", content=str(prompt))], max_tokens=body.get("max_tokens"), stream=bool(body.get("stream", False)), seed=body.get("seed"), n_blocks=body.get("n_blocks") or body.get("diffusion_blocks"), ) return chat_completions(req, authorization=authorization) if __name__ == "__main__": uvicorn.run(app, host=HOST, port=PORT, log_level="info")