nyara-phi / server.py
mssonya72's picture
Update server.py
81fe34d verified
Raw
History Blame Contribute Delete
7.77 kB
"""
Nyxen Engine — Fiction
Text generation backend for Cantrell Creatives Publisher Workspace.
Handles creative writing: scene work, voice-preserving rewrites, chapter extensions.
Model: Phi-3.5-mini-instruct (MIT license, no gating)
Hardware: HF Free CPU tier
Endpoints: /health, /generate, /chat, /extend, /rewrite
"""
import logging
from typing import List, Optional
import torch
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field
from transformers import AutoModelForCausalLM, AutoTokenizer
# ----- Logging -----
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(message)s",
)
log = logging.getLogger("nyxen-engine-fiction")
# ----- Config -----
MODEL_ID = "microsoft/Phi-3.5-mini-instruct"
DEVICE = "cpu"
MAX_NEW_TOKENS_DEFAULT = 800
TEMPERATURE_DEFAULT = 0.85
TOP_P_DEFAULT = 0.95
tokenizer = None
model = None
# ----- App -----
app = FastAPI(
title="Nyxen Engine — Fiction",
description="Text generation backend for Cantrell Creatives Publisher Workspace.",
version="1.0.0",
)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=False,
allow_methods=["*"],
allow_headers=["*"],
)
# ----- Request and response models -----
class GenerateRequest(BaseModel):
prompt: str = Field(..., description="The prompt to complete")
system: Optional[str] = Field(None, description="Optional system prompt")
max_new_tokens: int = Field(MAX_NEW_TOKENS_DEFAULT, ge=1, le=2048)
temperature: float = Field(TEMPERATURE_DEFAULT, ge=0.0, le=2.0)
top_p: float = Field(TOP_P_DEFAULT, ge=0.0, le=1.0)
class ChatMessage(BaseModel):
role: str = Field(..., description="user, assistant, or system")
content: str
class ChatRequest(BaseModel):
messages: List[ChatMessage]
max_new_tokens: int = Field(MAX_NEW_TOKENS_DEFAULT, ge=1, le=2048)
temperature: float = Field(TEMPERATURE_DEFAULT, ge=0.0, le=2.0)
top_p: float = Field(TOP_P_DEFAULT, ge=0.0, le=1.0)
class ExtendRequest(BaseModel):
text: str = Field(..., description="Prose to continue in the writer's voice")
instruction: Optional[str] = Field(
None, description="Optional direction for the continuation"
)
max_new_tokens: int = Field(MAX_NEW_TOKENS_DEFAULT, ge=1, le=2048)
temperature: float = Field(0.75, ge=0.0, le=2.0)
top_p: float = Field(TOP_P_DEFAULT, ge=0.0, le=1.0)
class RewriteRequest(BaseModel):
text: str = Field(..., description="Text to rewrite")
instruction: str = Field(..., description="How to rewrite (tone, length, focus)")
max_new_tokens: int = Field(MAX_NEW_TOKENS_DEFAULT, ge=1, le=2048)
temperature: float = Field(0.75, ge=0.0, le=2.0)
top_p: float = Field(TOP_P_DEFAULT, ge=0.0, le=1.0)
class TextResponse(BaseModel):
text: str
model: str
tokens_generated: int
# ----- Startup -----
@app.on_event("startup")
async def load_model():
global tokenizer, model
log.info(f"Loading model: {MODEL_ID} on {DEVICE}")
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, trust_remote_code=True)
# Phi-3.5-mini loads cleanly in float32 on free CPU (3.8B params, ~7GB)
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
torch_dtype=torch.float32,
device_map=DEVICE,
trust_remote_code=True,
low_cpu_mem_usage=True,
)
model.eval()
log.info("Model loaded and ready.")
# ----- Core generation -----
def _generate(
messages: List[dict],
max_new_tokens: int,
temperature: float,
top_p: float,
) -> tuple[str, int]:
"""Apply chat template, run generation, return (text, token_count)."""
if model is None or tokenizer is None:
raise HTTPException(status_code=503, detail="Model not loaded yet")
prompt_text = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
inputs = tokenizer(prompt_text, return_tensors="pt").to(DEVICE)
input_len = inputs["input_ids"].shape[1]
with torch.no_grad():
output = model.generate(
**inputs,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
do_sample=temperature > 0.0,
pad_token_id=tokenizer.eos_token_id,
use_cache=False,
)
generated_tokens = output[0][input_len:]
text = tokenizer.decode(generated_tokens, skip_special_tokens=True).strip()
return text, len(generated_tokens)
# ----- Endpoints -----
@app.get("/")
def root():
return {
"service": "Nyxen Engine — Fiction",
"model": MODEL_ID,
"endpoints": ["/health", "/generate", "/chat", "/extend", "/rewrite"],
"status": "online" if model is not None else "loading",
}
@app.get("/health")
def health():
return {
"status": "ok" if model is not None else "loading",
"model": MODEL_ID,
"device": DEVICE,
}
@app.post("/generate", response_model=TextResponse)
def generate(req: GenerateRequest):
"""Single-turn generation. Optional system prompt + user prompt."""
messages = []
if req.system:
messages.append({"role": "system", "content": req.system})
messages.append({"role": "user", "content": req.prompt})
text, tokens = _generate(
messages, req.max_new_tokens, req.temperature, req.top_p
)
return TextResponse(text=text, model=MODEL_ID, tokens_generated=tokens)
@app.post("/chat", response_model=TextResponse)
def chat(req: ChatRequest):
"""Multi-turn conversation. Caller manages message history."""
if not req.messages:
raise HTTPException(status_code=400, detail="messages must not be empty")
messages = [{"role": m.role, "content": m.content} for m in req.messages]
text, tokens = _generate(
messages, req.max_new_tokens, req.temperature, req.top_p
)
return TextResponse(text=text, model=MODEL_ID, tokens_generated=tokens)
@app.post("/extend", response_model=TextResponse)
def extend(req: ExtendRequest):
"""Continue prose in the writer's existing voice. Voice-preserving."""
system = (
"You are a voice-preserving writing assistant. Continue the user's prose "
"in their exact voice and style. Match their sentence rhythm, word choice, "
"tone, and pacing. Do not introduce new stylistic elements. Do not "
"summarize, explain, or add commentary. Output only the continuation."
)
user_content = req.text
if req.instruction:
user_content = f"{req.text}\n\n[Direction for continuation: {req.instruction}]"
messages = [
{"role": "system", "content": system},
{"role": "user", "content": user_content},
]
text, tokens = _generate(
messages, req.max_new_tokens, req.temperature, req.top_p
)
return TextResponse(text=text, model=MODEL_ID, tokens_generated=tokens)
@app.post("/rewrite", response_model=TextResponse)
def rewrite(req: RewriteRequest):
"""Rewrite text per an instruction (tone, length, focus)."""
system = (
"You are a precise editorial assistant. Rewrite the user's text according "
"to their instruction. Preserve meaning. Output only the rewritten text — "
"no preamble, no commentary, no explanations."
)
user_content = f"Instruction: {req.instruction}\n\nText:\n{req.text}"
messages = [
{"role": "system", "content": system},
{"role": "user", "content": user_content},
]
text, tokens = _generate(
messages, req.max_new_tokens, req.temperature, req.top_p
)
return TextResponse(text=text, model=MODEL_ID, tokens_generated=tokens)