TR-hash-tiny / app.py
Pacific-i64's picture
Test think prefill with pinned SFT epoch 3 revision
1631a4e verified
Raw
History Blame Contribute Delete
4.06 kB
"""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()