kiracromo35/BITNET_B
0
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 