LevArtesa's picture
Stage_V_3 serve: upload training/serve.py
bc3b40f verified
Raw
History Blame Contribute Delete
8.89 kB
"""Minimal HTTP inference server for Qwen3-4B (non-thinking mode).
Serves a single POST /generate endpoint that takes a prompt and returns a
completion. Designed to run on the same HF Space as training/evaluation.
Auto-pauses the Space after IDLE_SECONDS of no requests to avoid billing
surprises.
Environment variables:
HF_TOKEN: HuggingFace token for auto-pausing the Space on idle.
IDLE_SECONDS: Seconds of inactivity before auto-pause (default 900 = 15 min).
Endpoints:
GET / → "ok"
POST /generate → body {"prompt_text": str, "system_prompt"?: str,
"max_new_tokens"?: int, "temperature"?: float,
"top_p"?: float, "enable_thinking"?: bool}
→ {"completion": str, "gen_time_s": float}
POST /shutdown → pauses the Space immediately
"""
import json
import logging
import os
import sys
import threading
import time
from http.server import BaseHTTPRequestHandler, HTTPServer
import torch
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
BASE_MODEL = os.environ.get("MODEL_REPO", "Qwen/Qwen3-4B")
ADAPTER_REPO = os.environ.get("ADAPTER_REPO", "LevArtesa/grpo-humanizer-de-lora")
# Subfolder of a specific checkpoint, or "" for the top-level adapter
ADAPTER_SUBFOLDER = os.environ.get("ADAPTER_SUBFOLDER", "")
LOAD_ADAPTER = os.environ.get("LOAD_ADAPTER", "1").lower() in ("1", "true", "yes")
IDLE_SECONDS = int(os.environ.get("IDLE_SECONDS", "900"))
logger = logging.getLogger(__name__)
# Track last activity timestamp for idle-based auto-pause
_last_request_ts = time.time()
_lock = threading.Lock()
def _setup_logging() -> None:
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
handlers=[logging.StreamHandler(sys.stdout)],
)
def _load_model():
logger.info("Loading %s (4-bit)", BASE_MODEL)
bnb = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
)
tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(
BASE_MODEL,
quantization_config=bnb,
device_map="auto",
)
adapter_loaded = None
if LOAD_ADAPTER:
from huggingface_hub import snapshot_download
hf_token = os.environ.get("HF_TOKEN")
try:
if ADAPTER_SUBFOLDER:
logger.info("Downloading adapter %s/%s", ADAPTER_REPO, ADAPTER_SUBFOLDER)
local_dir = snapshot_download(
repo_id=ADAPTER_REPO,
allow_patterns=[f"{ADAPTER_SUBFOLDER}/*"],
token=hf_token,
)
adapter_path = str(__import__("pathlib").Path(local_dir) / ADAPTER_SUBFOLDER)
else:
logger.info("Downloading top-level adapter from %s", ADAPTER_REPO)
adapter_path = snapshot_download(
repo_id=ADAPTER_REPO,
allow_patterns=["adapter_*", "*.json", "*.txt", "tokenizer*", "chat_template*"],
token=hf_token,
)
logger.info("Applying LoRA adapter from %s", adapter_path)
model = PeftModel.from_pretrained(model, adapter_path)
adapter_loaded = adapter_path
except Exception as exc: # noqa: BLE001
logger.exception("Failed to load adapter; falling back to base model")
model.eval()
logger.info(
"Model loaded. Adapter=%s",
adapter_loaded if adapter_loaded else "none (base only)",
)
return model, tokenizer
def _pause_space() -> None:
space_id = os.environ.get("SPACE_ID")
hf_token = os.environ.get("HF_TOKEN")
if not space_id or not hf_token:
logger.warning("No SPACE_ID/HF_TOKEN — cannot auto-pause")
return
try:
from huggingface_hub import HfApi
logger.info("Pausing Space %s", space_id)
HfApi(token=hf_token).pause_space(space_id)
except Exception as exc: # noqa: BLE001
logger.error("Failed to pause: %s", exc)
def _idle_watchdog():
while True:
time.sleep(60)
with _lock:
idle = time.time() - _last_request_ts
if idle > IDLE_SECONDS:
logger.warning("Idle %.0fs > %d — auto-pausing Space", idle, IDLE_SECONDS)
_pause_space()
os._exit(0)
def _generate(model, tokenizer, body: dict) -> dict:
prompt_text = body.get("prompt_text") or ""
system_prompt = body.get("system_prompt")
max_new_tokens = int(body.get("max_new_tokens", 1024))
temperature = float(body.get("temperature", 0.9))
top_p = float(body.get("top_p", 0.95))
enable_thinking = bool(body.get("enable_thinking", False))
messages: list[dict] = []
if system_prompt:
messages.append({"role": "system", "content": system_prompt})
messages.append({"role": "user", "content": prompt_text})
# v4/v5/v51 patched: enable_thinking is Qwen-specific
try:
inputs = tokenizer.apply_chat_template(
messages,
add_generation_prompt=True,
return_tensors="pt",
enable_thinking=enable_thinking,
).to(model.device)
except Exception: # noqa: BLE001
inputs = tokenizer.apply_chat_template(
messages,
add_generation_prompt=True,
return_tensors="pt",
).to(model.device)
t0 = time.time()
with torch.no_grad():
output = model.generate(
inputs,
max_new_tokens=max_new_tokens,
do_sample=True,
temperature=temperature,
top_p=top_p,
pad_token_id=tokenizer.pad_token_id,
)
gen_time = time.time() - t0
new_tokens = output[0][inputs.shape[1]:]
completion = tokenizer.decode(new_tokens, skip_special_tokens=True).strip()
return {"completion": completion, "gen_time_s": round(gen_time, 2)}
def _make_handler(model, tokenizer):
class Handler(BaseHTTPRequestHandler):
def log_message(self, format, *args): # noqa: A002
pass
def do_GET(self):
global _last_request_ts
with _lock:
_last_request_ts = time.time()
if self.path in ("/", "/health"):
self.send_response(200)
self.send_header("Content-Type", "text/plain")
self.end_headers()
self.wfile.write(b"ok\n")
else:
self.send_response(404)
self.end_headers()
def do_POST(self):
global _last_request_ts
with _lock:
_last_request_ts = time.time()
if self.path == "/shutdown":
self.send_response(200)
self.end_headers()
self.wfile.write(b'{"status":"pausing"}')
threading.Thread(target=lambda: (time.sleep(1), _pause_space(), os._exit(0)), daemon=True).start()
return
if self.path != "/generate":
self.send_response(404)
self.end_headers()
return
length = int(self.headers.get("Content-Length", 0))
raw = self.rfile.read(length) if length else b"{}"
try:
body = json.loads(raw)
result = _generate(model, tokenizer, body)
payload = json.dumps(result, ensure_ascii=False).encode("utf-8")
self.send_response(200)
self.send_header("Content-Type", "application/json; charset=utf-8")
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)
except Exception as exc: # noqa: BLE001
logger.exception("generate error")
err = json.dumps({"error": str(exc)}).encode("utf-8")
self.send_response(500)
self.send_header("Content-Type", "application/json")
self.end_headers()
self.wfile.write(err)
return Handler
def main() -> None:
_setup_logging()
logger.info("=== Inference server — starting ===")
model, tokenizer = _load_model()
threading.Thread(target=_idle_watchdog, daemon=True).start()
logger.info("Idle watchdog armed: %d seconds", IDLE_SECONDS)
handler = _make_handler(model, tokenizer)
server = HTTPServer(("0.0.0.0", 7860), handler)
logger.info("Listening on :7860")
server.serve_forever()
if __name__ == "__main__":
main()