Spaces:
Paused
Paused
| """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() | |