"""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()