Spaces:
Running
Running
| """Thin Hugging Face Spaces launcher for the tr-hash-i64 API server.""" | |
| from __future__ import annotations | |
| import os | |
| import sys | |
| import threading | |
| from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer | |
| from pathlib import Path | |
| from huggingface_hub import snapshot_download | |
| def _available_cpus() -> int: | |
| try: | |
| return len(os.sched_getaffinity(0)) | |
| except AttributeError: | |
| return os.cpu_count() or 1 | |
| def _enabled(name: str, default: bool) -> bool: | |
| value = os.environ.get(name) | |
| if value is None: | |
| return default | |
| return value.strip().lower() in {"1", "true", "yes", "on"} | |
| cpu_count = _available_cpus() | |
| threads = max(1, min(int(os.environ.get("CPU_THREADS", cpu_count)), cpu_count)) | |
| os.environ.setdefault("OMP_NUM_THREADS", str(threads)) | |
| os.environ.setdefault("MKL_NUM_THREADS", str(threads)) | |
| os.environ.setdefault("OPENBLAS_NUM_THREADS", str(threads)) | |
| os.environ.setdefault("TR_HASH_I64_CPU_THREADS", str(threads)) | |
| os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") | |
| os.environ.setdefault("MALLOC_ARENA_MAX", "2") | |
| model_name = os.environ["MODEL_NAME"] | |
| model_repo = os.environ.get( | |
| "MODEL_REPO", "AETHORIA-AI/TR-HASH-MoE-200M-160B-SFT" | |
| ) | |
| model_revision = os.environ.get("MODEL_REVISION") or None | |
| model_dir_override = os.environ.get("MODEL_DIR") | |
| port = int(os.environ.get("PORT", "7860")) | |
| MODEL_FILES = [ | |
| "config.json", | |
| "model.safetensors", | |
| "chat_template.json", | |
| "chat_template.jinja", | |
| "model_config.yaml", | |
| "special_tokens_map.json", | |
| "tokenizer.json", | |
| "tokenizer_config.json", | |
| ] | |
| class _MaintenanceHandler(BaseHTTPRequestHandler): | |
| def do_GET(self) -> None: | |
| body = ( | |
| b'{"status":"maintenance",' | |
| b'"detail":"TR-HASH MoE 200M full-SFT checkpoint is loading."}' | |
| ) | |
| self.send_response(503) | |
| self.send_header("Content-Type", "application/json") | |
| self.send_header("Access-Control-Allow-Origin", "*") | |
| self.send_header("Retry-After", "30") | |
| self.send_header("Content-Length", str(len(body))) | |
| self.end_headers() | |
| self.wfile.write(body) | |
| def log_message(self, *_: object) -> None: | |
| return | |
| def _resolve_checkpoint() -> str: | |
| if model_dir_override: | |
| config_path = Path(model_dir_override) / "config.json" | |
| if not config_path.is_file(): | |
| raise FileNotFoundError(f"MODEL_DIR has no config.json: {config_path}") | |
| return model_dir_override | |
| server = ThreadingHTTPServer(("0.0.0.0", port), _MaintenanceHandler) | |
| thread = threading.Thread(target=server.serve_forever, daemon=True) | |
| thread.start() | |
| print( | |
| f"Resolving {model_repo} at revision {model_revision or 'main'}.", | |
| flush=True, | |
| ) | |
| try: | |
| checkpoint_dir = snapshot_download( | |
| repo_id=model_repo, | |
| revision=model_revision, | |
| allow_patterns=MODEL_FILES, | |
| ) | |
| finally: | |
| server.shutdown() | |
| server.server_close() | |
| thread.join(timeout=5) | |
| print(f"Checkpoint ready at {checkpoint_dir}; starting tr-hash-i64.", flush=True) | |
| return checkpoint_dir | |
| model_dir = _resolve_checkpoint() | |
| import torch | |
| use_cpu_int8 = not torch.cuda.is_available() and _enabled("CPU_INT8", True) | |
| argv = [ | |
| "tr-hash-i64", "serve", model_name, | |
| "--checkpoint", model_dir, | |
| "--host", "0.0.0.0", | |
| "--port", str(port), | |
| "--dtype", "float16" if torch.cuda.is_available() else "float32", | |
| "--quantization", "int8" if use_cpu_int8 else "none", | |
| "--max-batch-size", os.environ.get("MAX_BATCH_SIZE", "4"), | |
| "--max-kv-blocks", os.environ.get("MAX_KV_BLOCKS", "128"), | |
| "--chunk-size", os.environ.get("PREFILL_CHUNK_SIZE", "256"), | |
| "--rate-limit", os.environ.get("RATE_LIMIT", "60"), | |
| "--max-pending", os.environ.get("MAX_PENDING", "16"), | |
| "--context-compact-tokens", os.environ.get("CONTEXT_COMPACT_TOKENS", "1024"), | |
| ] | |
| api_key = os.environ.get("TR_HASH_I64_API_KEY") | |
| if api_key: | |
| argv.extend(["--api-key", api_key]) | |
| sys.argv = argv | |
| from tr_hash_i64.cli import main | |
| main() | |