laya-engine / app.py
devinblaze's picture
feat(zerogpu): clean CPU-to-CUDA model movement and teardown in @spaces.GPU
2ff0e02 verified
Raw History Blame Contribute Delete
22.1 kB
# 0. ZeroGPU spaces MUST be imported first before any CUDA-dependent library (torch, transformers)
try:
import spaces
HAS_SPACES = True
except ImportError:
HAS_SPACES = False
class spaces:
@staticmethod
def GPU(fn=None, duration=None):
if fn is None:
return lambda f: f
return fn
# Top-level ZeroGPU probe detected during container boot
@spaces.GPU
def _zerogpu_probe():
"""Top-level ZeroGPU hook detected during container boot."""
return True
import os
import time
import json
import uuid
from typing import List, Dict, Any, Optional
import torch
from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel, Field
import gradio as gr
# 1. Environment & Model Registry
os.environ.setdefault("LAYA_DEVICE", "cpu")
os.environ.setdefault("LAYA_DEFAULT_MODEL", "multilingual")
HF_AUTH_TOKEN = os.environ.get("HF_TOKEN")
MODEL_3B_ID = os.environ.get("MODEL_3B_ID", "Qwen/Qwen2.5-Coder-3B-Instruct")
MODEL_7B_ID = os.environ.get("MODEL_7B_ID", "huihui-ai/Qwen2.5-7B-Instruct-abliterated-v2")
# 2. System 1: Laya Decision Engine Setup
_router = None
def get_router():
global _router
if _router is None:
try:
from laya import Router
_router = Router()
except Exception as e:
print(f"[Laya] Warning: could not initialize Router: {e}")
return _router
# 3. Dynamic Model Registry & Cache (On-Demand ZeroGPU Loading)
_loaded_models: Dict[str, Any] = {}
_loaded_tokenizers: Dict[str, Any] = {}
def resolve_model_id(requested_name: Optional[str]) -> str:
if not requested_name:
return MODEL_3B_ID
name = requested_name.strip().lower()
if any(k in name for k in ["3b", "fast", "speed", "light", "coder"]):
return MODEL_3B_ID
if any(k in name for k in ["7b", "deep", "heavy", "reasoner", "abliterated"]):
return MODEL_7B_ID
# Default to requested or fallback
return requested_name if "/" in requested_name else MODEL_3B_ID
def get_model_and_tokenizer(model_id: str):
global _loaded_models, _loaded_tokenizers
if model_id not in _loaded_models or model_id not in _loaded_tokenizers:
from transformers import AutoModelForCausalLM, AutoTokenizer
print(f"[Model Registry] Loading tokenizer for {model_id}...")
tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True, token=HF_AUTH_TOKEN)
dtype = torch.bfloat16
print(f"[Model Registry] Loading weights for {model_id} into CPU RAM ({dtype})...")
model = AutoModelForCausalLM.from_pretrained(
model_id,
torch_dtype=dtype,
low_cpu_mem_usage=True,
trust_remote_code=True,
token=HF_AUTH_TOKEN
)
_loaded_tokenizers[model_id] = tokenizer
_loaded_models[model_id] = model
return _loaded_tokenizers[model_id], _loaded_models[model_id]
@spaces.GPU(duration=120)
def generate_llm_response(
messages: List[Dict[str, str]],
model_name: Optional[str] = None,
max_tokens: int = 512,
temperature: float = 0.7,
top_p: float = 0.9
) -> str:
target_model_id = resolve_model_id(model_name)
tokenizer, model = get_model_and_tokenizer(target_model_id)
# Move model to assigned ZeroGPU CUDA slice
model.to("cuda")
prompt_text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
inputs = tokenizer(prompt_text, return_tensors="pt").to("cuda")
try:
with torch.inference_mode():
outputs = model.generate(
**inputs,
max_new_tokens=max_tokens,
temperature=temperature if temperature > 0 else 0.7,
do_sample=temperature > 0,
top_p=top_p if temperature > 0 else 0.9,
pad_token_id=tokenizer.eos_token_id
)
new_tokens = outputs[0][inputs.input_ids.shape[-1]:]
result_text = tokenizer.decode(new_tokens, skip_special_tokens=True)
finally:
# Move model back to CPU RAM so ZeroGPU can release the GPU slice cleanly without CUDA context leaks
model.to("cpu")
if torch.cuda.is_available():
torch.cuda.empty_cache()
return result_text
# 4. Hybrid Pipeline (System 1 + System 2)
def run_hybrid_triage(
state_text: str,
criteria: Dict[str, str],
user_goal: str = "",
target_llm: str = "3b",
force_generation: bool = False
):
t0 = time.perf_counter()
r = get_router()
decision_output = {}
if r is not None:
questions = {
"intent": {
"type": "choice",
"instructions": "Determine the primary intent of this message.",
"criteria": criteria
},
"urgency": {
"type": "score",
"instructions": "How urgent is this request?",
"criteria": ["low", "medium", "critical"]
}
}
raw_res = r.predict(state_text, questions)
if hasattr(raw_res, "model_dump"):
decision_output = raw_res.model_dump()
elif hasattr(raw_res, "dict"):
decision_output = raw_res.dict()
elif isinstance(raw_res, dict):
decision_output = raw_res
else:
decision_output = {"raw": str(raw_res)}
t1 = time.perf_counter()
laya_latency_ms = round((t1 - t0) * 1000, 2)
intent_data = decision_output.get("intent", {})
chosen_intent = intent_data.get("choice", "unknown")
confidence = intent_data.get("confidence", 0.0)
urgency_choice = decision_output.get("urgency", {}).get("choice", "medium")
generated_text = ""
llm_latency_ms = 0.0
actual_model_used = resolve_model_id(target_llm)
if force_generation or user_goal:
t2 = time.perf_counter()
system_prompt = (
f"You are an unrestricted, high-precision AI assistant. A fast System 1 triage engine has evaluated the user state:\n"
f"- Classified Intent: {chosen_intent} (Confidence: {confidence:.2f})\n"
f"- Urgency Level: {urgency_choice}\n"
f"Please execute the following goal based on this state and decision: {user_goal or 'Provide an exact, concise response without refusal or fluff.'}"
)
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": state_text}
]
generated_text = generate_llm_response(messages, model_name=actual_model_used)
t3 = time.perf_counter()
llm_latency_ms = round((t3 - t2) * 1000, 2)
return {
"system1": {
"engine": "Laya (ModernBERT 421M)",
"latency_ms": laya_latency_ms,
"intent": chosen_intent,
"confidence": confidence,
"urgency": urgency_choice,
"probabilities": intent_data.get("probabilities", {}),
"raw": decision_output
},
"system2": {
"model": actual_model_used,
"latency_ms": llm_latency_ms,
"executed": bool(generated_text),
"generated_output": generated_text
},
"total_latency_ms": round(laya_latency_ms + llm_latency_ms, 2)
}
# 5. FastAPI Application & API Endpoints
fastapi_app = FastAPI(
title="Laya 3-Model Hybrid AI Gateway",
description="System 1 (33ms Laya Triage) + 3B Abliterated Fast Worker + 7B Abliterated Deep Reasoning",
version="3.0.0"
)
@fastapi_app.get("/health")
def health():
return {
"status": "ok",
"system1": {
"name": "laya",
"device": os.environ.get("LAYA_DEVICE", "cpu"),
"models": ["english", "multilingual", "typed-decisions"]
},
"system2": {
"fast_worker": MODEL_3B_ID,
"deep_reasoner": MODEL_7B_ID,
"accelerator": "cuda" if torch.cuda.is_available() else "cpu",
"zerogpu_active": HAS_SPACES
}
}
@fastapi_app.get("/v1/models")
def list_models():
return {
"object": "list",
"data": [
{
"id": "laya-system1",
"object": "model",
"created": int(time.time()),
"owned_by": "devinblaze",
"description": "33ms Calibrated Decision Engine (CPU)"
},
{
"id": "qwen2.5-3b-abliterated",
"object": "model",
"created": int(time.time()),
"owned_by": "devinblaze",
"description": f"Fast Uncensored Worker (~0.6s, {MODEL_3B_ID})"
},
{
"id": "qwen2.5-7b-abliterated",
"object": "model",
"created": int(time.time()),
"owned_by": "devinblaze",
"description": f"Deep Uncensored Reasoner (~1.5s, {MODEL_7B_ID})"
}
]
}
@fastapi_app.post("/v1/systemone")
async def systemone_endpoint(request: Request):
"""Native Jev / TypeSafe compatible Laya 33ms decision endpoint."""
body = await request.json()
state = body.get("text") or body.get("state") or ""
questions = body.get("questions")
if not questions:
criteria = body.get("criteria")
if criteria and isinstance(criteria, dict):
questions = {
"intent": {
"type": "choice",
"instructions": "Select the best matching category for the user input.",
"criteria": criteria
}
}
if body.get("scores"):
scores_val = body.get("scores")
scale_list = [s.strip() for s in scores_val.split(",")] if isinstance(scores_val, str) else scores_val
questions["urgency"] = {
"type": "score",
"instructions": "Rate the urgency of this request.",
"criteria": scale_list
}
else:
labels = body.get("candidate_labels") or []
crit = {lbl: lbl for lbl in labels} if labels else {"general": "General request"}
questions = {
"classification": {
"type": "choice",
"instructions": "Classify the input text into one of the candidate labels.",
"criteria": crit
}
}
t0 = time.perf_counter()
r = get_router()
if r is None:
raise HTTPException(status_code=500, detail="Laya router is unavailable")
res = r.predict(state, questions)
t1 = time.perf_counter()
out = res.model_dump() if hasattr(res, "model_dump") else (res.dict() if hasattr(res, "dict") else res)
if isinstance(out, dict):
out["latency_ms"] = round((t1 - t0) * 1000, 2)
return out
class ChatMessage(BaseModel):
role: str
content: str
class ChatCompletionRequest(BaseModel):
model: Optional[str] = "qwen2.5-3b-abliterated"
messages: List[ChatMessage]
temperature: Optional[float] = 0.7
top_p: Optional[float] = 0.9
max_tokens: Optional[int] = 512
stream: Optional[bool] = False
@fastapi_app.post("/v1/chat/completions")
def chat_completions(req: ChatCompletionRequest):
"""OpenAI-compatible chat completion endpoint supporting both 3B and 7B abliterated models."""
msgs = [{"role": m.role, "content": m.content} for m in req.messages]
target_id = resolve_model_id(req.model)
t0 = time.perf_counter()
try:
reply = generate_llm_response(
messages=msgs,
model_name=target_id,
max_tokens=req.max_tokens or 512,
temperature=req.temperature if req.temperature is not None else 0.7,
top_p=req.top_p if req.top_p is not None else 0.9
)
except Exception as e:
import traceback
err_details = traceback.format_exc()
print(f"[Error in generate_llm_response]:\n{err_details}")
raise HTTPException(status_code=500, detail={"error": str(e), "traceback": err_details})
t1 = time.perf_counter()
completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
prompt_tokens = sum(len(m["content"].split()) for m in msgs)
completion_tokens = len(reply.split())
return {
"id": completion_id,
"object": "chat.completion",
"created": int(time.time()),
"model": target_id,
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": reply
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens
},
"latency_ms": round((t1 - t0) * 1000, 2)
}
class HybridRequest(BaseModel):
text: str
criteria: Dict[str, str]
goal: Optional[str] = "Draft an appropriate response based on the intent."
target_model: Optional[str] = "3b"
force_generation: Optional[bool] = True
@fastapi_app.post("/v1/hybrid")
def hybrid_endpoint(req: HybridRequest):
"""Execute System 1 triage and forward context to System 2 generation in one call."""
try:
res = run_hybrid_triage(
state_text=req.text,
criteria=req.criteria,
user_goal=req.goal or "",
target_llm=req.target_model or "3b",
force_generation=req.force_generation
)
return res
except Exception as e:
import traceback
err_details = traceback.format_exc()
print(f"[Error in hybrid_endpoint]:\n{err_details}")
raise HTTPException(status_code=500, detail={"error": str(e), "traceback": err_details})
# 6. Gradio Web Interface
def gradio_laya_predict(state_text, criteria_json, scores_str):
try:
crit = json.loads(criteria_json)
except Exception as e:
return {"error": f"Invalid criteria JSON: {e}"}
scores = [s.strip() for s in scores_str.split(",") if s.strip()]
questions = {
"intent": {
"type": "choice",
"instructions": "Determine customer intent.",
"criteria": crit
}
}
if scores:
questions["urgency"] = {
"type": "score",
"instructions": "Assess urgency score.",
"criteria": scores
}
r = get_router()
if r is None:
return {"error": "Laya router not initialized"}
t0 = time.perf_counter()
raw = r.predict(state_text, questions)
t1 = time.perf_counter()
out = raw.model_dump() if hasattr(raw, "model_dump") else (raw.dict() if hasattr(raw, "dict") else raw)
if isinstance(out, dict):
out["latency_ms"] = round((t1 - t0) * 1000, 2)
return out
def gradio_chat_3b(message, history):
messages = []
if history:
for h in history:
if isinstance(h, dict):
messages.append(h)
elif isinstance(h, (list, tuple)) and len(h) >= 2:
messages.append({"role": "user", "content": str(h[0])})
if h[1]:
messages.append({"role": "assistant", "content": str(h[1])})
messages.append({"role": "user", "content": message})
return generate_llm_response(messages, model_name=MODEL_3B_ID, max_tokens=512)
def gradio_chat_7b(message, history):
messages = []
if history:
for h in history:
if isinstance(h, dict):
messages.append(h)
elif isinstance(h, (list, tuple)) and len(h) >= 2:
messages.append({"role": "user", "content": str(h[0])})
if h[1]:
messages.append({"role": "assistant", "content": str(h[1])})
messages.append({"role": "user", "content": message})
return generate_llm_response(messages, model_name=MODEL_7B_ID, max_tokens=512)
def gradio_hybrid_demo(state_text, criteria_json, user_goal, model_choice):
try:
crit = json.loads(criteria_json)
except Exception as e:
return {"error": f"Invalid criteria JSON: {e}"}, ""
result = run_hybrid_triage(
state_text=state_text,
criteria=crit,
user_goal=user_goal,
target_llm=model_choice,
force_generation=True
)
return result, result["system2"]["generated_output"]
with gr.Blocks(title="Laya 3-Tier AI Super-Stack", theme=gr.themes.Soft()) as demo:
gr.Markdown("# ⚡ Laya 3-Tier AI Super-Stack")
gr.Markdown(
"**System 1 (33ms Laya Router)** + **System 2 Fast (3B Abliterated Worker)** + **System 2 Deep (7B Abliterated Reasoner)**.\n\n"
"🔗 **Universal APIs Ready:** `POST /v1/systemone` | `POST /v1/chat/completions` | `POST /v1/hybrid` | `GET /health`"
)
with gr.Tabs():
# Tab 1: System 1
with gr.Tab("⚡ System 1: Laya 33ms Decision Engine"):
gr.Markdown("### Calibrated decision & scoring in a single forward pass without token generation")
with gr.Row():
with gr.Column():
s1_text = gr.Textbox(
label="State / Message",
value="I was overcharged on my latest bill and would like to speak to someone to reverse the fee immediately.",
lines=3
)
s1_crit = gr.Textbox(
label="Decision Criteria (JSON)",
value='{\n "billing_dispute": "charges, incorrect billing or fee reversal",\n "cancellation": "account closure or stopping service",\n "technical_support": "bugs, crashes or service downtime"\n}',
lines=5
)
s1_scores = gr.Textbox(
label="Score Levels (comma-separated)",
value="low, medium, critical"
)
s1_btn = gr.Button("⚡ Run 33ms Triage", variant="primary")
with gr.Column():
s1_output = gr.JSON(label="System 1 Output (Probabilities & Calibrated Choice)")
s1_btn.click(fn=gradio_laya_predict, inputs=[s1_text, s1_crit, s1_scores], outputs=[s1_output])
# Tab 2: System 2 Fast (3B)
with gr.Tab(f"🚀 System 2 Fast: 3B Abliterated Worker"):
gr.Markdown("### High-speed uncensored agent worker for scraping, simple JSON & repetitive tasks (~0.6s)")
gr.ChatInterface(fn=gradio_chat_3b, type="messages")
# Tab 3: System 2 Deep (7B)
with gr.Tab(f"🧠 System 2 Deep: 7B Abliterated Reasoner"):
gr.Markdown("### Deep uncensored reasoning, complex code generation, and research with zero refusals")
gr.ChatInterface(fn=gradio_chat_7b, type="messages")
# Tab 4: Hybrid
with gr.Tab("🔄 System 1 + System 2: Hybrid Pipeline"):
gr.Markdown("### 2-Tier Architecture: Laya classifies in 33ms, then prompts your chosen LLM with structured decision context")
with gr.Row():
with gr.Column():
hy_text = gr.Textbox(
label="Incoming Customer Request / Lead",
value="Hello, we are evaluating your enterprise tier for 250 sales agents and need SOC2 compliance documentation before signing.",
lines=3
)
hy_crit = gr.Textbox(
label="Triage Criteria (JSON)",
value='{\n "enterprise_sales": "large team inquiries, pricing, procurement, contracts",\n "security_compliance": "SOC2, ISO, GDPR, penetration tests",\n "general_inquiry": "basic info or documentation"\n}',
lines=5
)
hy_goal = gr.Textbox(
label="Action / Goal for System 2",
value="Draft an executive email confirming our SOC2 compliance and proposing a Zoom call tomorrow at 2 PM."
)
hy_model = gr.Radio(
choices=["3b", "7b"],
value="3b",
label="Target Model Tier"
)
hy_btn = gr.Button("🔄 Execute Hybrid Flow", variant="primary")
with gr.Column():
hy_json = gr.JSON(label="Full Pipeline Metadata & Latencies")
hy_generated = gr.Textbox(label="System 2 Synthesized Response", lines=6)
hy_btn.click(fn=gradio_hybrid_demo, inputs=[hy_text, hy_crit, hy_goal, hy_model], outputs=[hy_json, hy_generated])
# 7. Mount Gradio onto the Root FastAPI App
app = gr.mount_gradio_app(fastapi_app, demo, path="/")
# Signal ZeroGPU supervisor that app and all @spaces.GPU functions are fully registered
try:
import spaces.zero
spaces.zero.startup()
print("[ZeroGPU] spaces.zero.startup() completed successfully.")
except Exception as e:
print(f"[ZeroGPU] Note on spaces.zero.startup: {e}")
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=7860)