CoolFace
Apppublic

kiracromo35/BITNET_B

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes
server.py278 linesDownload Raw Back to root
1"""2OMNI-NEXUS · Worker de cómputo elástico3=======================================4Envuelve llama.cpp en una API compatible con OpenAI para que Odysseus lo5vea como un endpoint local más, sin modificar nada del proyecto.6 7Endpoints:8    GET  /health              → estado (lo usa el healthcheck de HF)9    GET  /v1/models           → catálogo (lo usa Odysseus al añadir el endpoint)10    POST /v1/chat/completions → inferencia, con y sin streaming11    GET  /warm                → despierta el Space sin gastar inferencia12 13Nota sobre /warm: los Spaces gratuitos se suspenden tras un rato inactivos y14tardan ~60-90 s en volver. El router del gateway llama a /warm para levantar15el siguiente worker mientras el actual trabaja, así el usuario no percibe el16arranque en frío.17"""18 19from __future__ import annotations20 21import json22import os23import time24import uuid25from typing import Any, AsyncIterator, Dict, List, Optional26 27import httpx28from fastapi import FastAPI, HTTPException29from fastapi.middleware.cors import CORSMiddleware30from fastapi.responses import StreamingResponse31from pydantic import BaseModel, Field32 33LLAMA_URL = f"http://127.0.0.1:{os.getenv('LLAMA_PORT', '8080')}"34WORKER_ID = os.getenv("WORKER_ID", "A")35N_CTX = int(os.getenv("N_CTX", "4096"))36TIMEOUT = float(os.getenv("LLAMA_TIMEOUT", "180.0"))37MODEL_NAME = os.getenv("MODEL_NAME", "omni-nexus/worker")38 39# Plantilla de chat. BitNet b1.58-2B-4T se entrenó con el formato de Llama 3,40# no con ChatML. Usar la plantilla equivocada degrada mucho la calidad sin dar41# ningún error visible — el modelo simplemente responde peor. Por eso es una42# variable explícita y no una suposición.43CHAT_TEMPLATE = os.getenv("CHAT_TEMPLATE", "llama3").lower()44 45app = FastAPI(title=f"Omni-Nexus Worker {WORKER_ID}", version="1.0.0")46 47# CORS abierto: el endpoint se consume desde el navegador de Odysseus.48app.add_middleware(49    CORSMiddleware,50    allow_origins=["*"],51    allow_methods=["*"],52    allow_headers=["*"],53)54 55_stats = {"requests": 0, "tokens": 0, "started_at": time.time()}56 57 58# ---------------------------------------------------------------------------59# Modelos de entrada (subconjunto del esquema OpenAI que Odysseus usa)60# ---------------------------------------------------------------------------61class Message(BaseModel):62    role: str63    content: Any = ""64 65 66class ChatRequest(BaseModel):67    model: Optional[str] = None68    messages: List[Message]69    temperature: float = 0.770    top_p: float = 0.9571    max_tokens: Optional[int] = Field(default=1024)72    stream: bool = False73    stop: Optional[List[str]] = None74 75 76# ---------------------------------------------------------------------------77def _flatten(content: Any) -> str:78    """El contenido puede venir como string o como lista de bloques."""79    if isinstance(content, str):80        return content81    if isinstance(content, list):82        parts = []83        for block in content:84            if isinstance(block, dict) and block.get("type") == "text":85                parts.append(block.get("text", ""))86            elif isinstance(block, str):87                parts.append(block)88        return "\n".join(parts)89    return str(content or "")90 91 92def _prompt_llama3(messages: List[Message]) -> str:93    """Formato Llama 3 — el correcto para BitNet b1.58-2B-4T."""94    out = ["<|begin_of_text|>"]95    for m in messages:96        role = m.role if m.role in ("system", "user", "assistant") else "user"97        out.append(98            f"<|start_header_id|>{role}<|end_header_id|>\n\n"99            f"{_flatten(m.content)}<|eot_id|>"100        )101    out.append("<|start_header_id|>assistant<|end_header_id|>\n\n")102    return "".join(out)103 104 105def _prompt_chatml(messages: List[Message]) -> str:106    """Formato ChatML — Qwen 2.5 y buena parte de los instruct actuales."""107    out = []108    for m in messages:109        role = m.role if m.role in ("system", "user", "assistant") else "user"110        out.append(f"<|im_start|>{role}\n{_flatten(m.content)}<|im_end|>")111    out.append("<|im_start|>assistant\n")112    return "\n".join(out)113 114 115def to_prompt(messages: List[Message]) -> str:116    if CHAT_TEMPLATE == "chatml":117        return _prompt_chatml(messages)118    return _prompt_llama3(messages)119 120 121# Se incluyen los stops de ambas plantillas: sobra un par de tokens en la122# lista y evita que un cambio de modelo deje basura al final de la respuesta.123STOP_TOKENS = [124    "<|eot_id|>", "<|start_header_id|>", "<|end_header_id|>",125    "<|im_end|>", "<|im_start|>", "</s>", "<|end_of_text|>",126]127 128 129# ---------------------------------------------------------------------------130@app.get("/health")131async def health():132    llama_ok = False133    try:134        async with httpx.AsyncClient(timeout=5.0) as c:135            r = await c.get(f"{LLAMA_URL}/health")136            llama_ok = r.status_code == 200137    except Exception:138        pass139    return {140        "status": "ok" if llama_ok else "starting",141        "worker": WORKER_ID,142        "llama": llama_ok,143        "ctx": N_CTX,144        "template": CHAT_TEMPLATE,145        "kind": os.getenv("WORKER_KIND", "chat"),146        "uptime_s": round(time.time() - _stats["started_at"]),147        "requests": _stats["requests"],148    }149 150 151@app.get("/warm")152async def warm():153    """Despierta el Space sin gastar inferencia. Responde apenas el proceso vive."""154    return {"worker": WORKER_ID, "awake": True, "ts": time.time()}155 156 157@app.get("/v1/models")158async def models():159    return {160        "object": "list",161        "data": [162            {163                "id": MODEL_NAME,164                "object": "model",165                "created": int(_stats["started_at"]),166                "owned_by": f"omni-nexus-worker-{WORKER_ID}",167                "context_length": N_CTX,168            }169        ],170    }171 172 173@app.post("/v1/chat/completions")174async def chat(req: ChatRequest):175    prompt = to_prompt(req.messages)176    payload = {177        "prompt": prompt,178        "temperature": req.temperature,179        "top_p": req.top_p,180        "n_predict": req.max_tokens or 1024,181        "stop": STOP_TOKENS + (req.stop or []),182        "stream": req.stream,183        "cache_prompt": True,  # reutiliza el KV cache entre turnos: gran ahorro184    }185    _stats["requests"] += 1186 187    if req.stream:188        return StreamingResponse(189            _stream(payload, req.model), media_type="text/event-stream"190        )191 192    try:193        async with httpx.AsyncClient(timeout=TIMEOUT) as c:194            r = await c.post(f"{LLAMA_URL}/completion", json=payload)195            r.raise_for_status()196            data = r.json()197    except httpx.TimeoutException:198        raise HTTPException(504, detail=f"Worker {WORKER_ID}: timeout de inferencia")199    except Exception as exc:200        raise HTTPException(502, detail=f"Worker {WORKER_ID}: {exc}")201 202    text = (data.get("content") or "").strip()203    for tok in STOP_TOKENS:204        text = text.replace(tok, "")205 206    n_out = data.get("tokens_predicted", 0)207    n_in = data.get("tokens_evaluated", 0)208    _stats["tokens"] += n_out209 210    return {211        "id": f"chatcmpl-{uuid.uuid4().hex[:12]}",212        "object": "chat.completion",213        "created": int(time.time()),214        "model": req.model or MODEL_NAME,215        "choices": [216            {217                "index": 0,218                "message": {"role": "assistant", "content": text.strip()},219                "finish_reason": "stop",220            }221        ],222        "usage": {223            "prompt_tokens": n_in,224            "completion_tokens": n_out,225            "total_tokens": n_in + n_out,226        },227        "_worker": WORKER_ID,228    }229 230 231async def _stream(payload: Dict[str, Any], model: Optional[str]) -> AsyncIterator[str]:232    cid = f"chatcmpl-{uuid.uuid4().hex[:12]}"233    created = int(time.time())234 235    def frame(delta: Dict[str, Any], finish: Optional[str] = None) -> str:236        return "data: " + json.dumps(237            {238                "id": cid,239                "object": "chat.completion.chunk",240                "created": created,241                "model": model or MODEL_NAME,242                "choices": [{"index": 0, "delta": delta, "finish_reason": finish}],243            }244        ) + "\n\n"245 246    yield frame({"role": "assistant", "content": ""})247 248    try:249        async with httpx.AsyncClient(timeout=TIMEOUT) as c:250            async with c.stream("POST", f"{LLAMA_URL}/completion", json=payload) as r:251                async for line in r.aiter_lines():252                    if not line.startswith("data: "):253                        continue254                    try:255                        chunk = json.loads(line[6:])256                    except json.JSONDecodeError:257                        continue258                    piece = chunk.get("content", "")259                    if piece:260                        _stats["tokens"] += 1261                        yield frame({"content": piece})262                    if chunk.get("stop"):263                        break264    except Exception as exc:265        yield frame({"content": f"\n[worker {WORKER_ID}: {exc}]"})266 267    yield frame({}, finish="stop")268    yield "data: [DONE]\n\n"269 270 271@app.get("/")272async def root():273    return {274        "name": f"Omni-Nexus Worker {WORKER_ID}",275        "usage": "Añádelo en Odysseus → Settings → Add Local Models con la URL /v1",276        "endpoints": ["/health", "/warm", "/v1/models", "/v1/chat/completions"],277    }278