DeepXR/Helion-OSC
465
1"""2Helion-OSC API Server3FastAPI-based REST API for serving Helion-OSC model4"""5 6from fastapi import FastAPI, HTTPException, BackgroundTasks7from fastapi.middleware.cors import CORSMiddleware8from fastapi.responses import StreamingResponse9from pydantic import BaseModel, Field10from typing import Optional, List, Dict, Any, AsyncGenerator11import torch12from transformers import AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer13from threading import Thread14import uvicorn15import logging16import time17import json18from queue import Queue19import asyncio20 21logging.basicConfig(level=logging.INFO)22logger = logging.getLogger(__name__)23 24# Initialize FastAPI app25app = FastAPI(26 title="Helion-OSC API",27 description="REST API for Helion-OSC Code Generation Model",28 version="1.0.0"29)30 31# Add CORS middleware32app.add_middleware(33 CORSMiddleware,34 allow_origins=["*"],35 allow_credentials=True,36 allow_methods=["*"],37 allow_headers=["*"],38)39 40# Global model variables41model = None42tokenizer = None43device = None44 45 46class GenerationRequest(BaseModel):47 """Request model for text generation"""48 prompt: str = Field(..., description="Input prompt for generation")49 max_length: int = Field(2048, ge=1, le=16384, description="Maximum length of generation")50 temperature: float = Field(0.7, ge=0.0, le=2.0, description="Sampling temperature")51 top_p: float = Field(0.95, ge=0.0, le=1.0, description="Nucleus sampling parameter")52 top_k: int = Field(50, ge=0, le=200, description="Top-k sampling parameter")53 repetition_penalty: float = Field(1.05, ge=1.0, le=2.0, description="Repetition penalty")54 do_sample: bool = Field(True, description="Whether to use sampling")55 num_return_sequences: int = Field(1, ge=1, le=10, description="Number of sequences to generate")56 stop_sequences: Optional[List[str]] = Field(None, description="Stop generation at these sequences")57 stream: bool = Field(False, description="Stream the response")58 task_type: Optional[str] = Field("code_generation", description="Task type for optimized parameters")59 60 61class GenerationResponse(BaseModel):62 """Response model for text generation"""63 generated_text: str64 prompt: str65 model: str66 generation_time: float67 tokens_generated: int68 69 70class ModelInfo(BaseModel):71 """Model information"""72 model_name: str73 model_type: str74 vocabulary_size: int75 hidden_size: int76 num_layers: int77 device: str78 dtype: str79 max_position_embeddings: int80 81 82class HealthResponse(BaseModel):83 """Health check response"""84 status: str85 model_loaded: bool86 device: str87 timestamp: float88 89 90@app.on_event("startup")91async def load_model():92 """Load model on startup"""93 global model, tokenizer, device94 95 logger.info("Loading Helion-OSC model...")96 97 model_name = "DeepXR/Helion-OSC"98 device = "cuda" if torch.cuda.is_available() else "cpu"99 100 try:101 # Load tokenizer102 tokenizer = AutoTokenizer.from_pretrained(103 model_name,104 trust_remote_code=True105 )106 107 if tokenizer.pad_token is None:108 tokenizer.pad_token = tokenizer.eos_token109 110 # Load model111 model = AutoModelForCausalLM.from_pretrained(112 model_name,113 torch_dtype=torch.bfloat16 if device == "cuda" else torch.float32,114 device_map="auto" if device == "cuda" else None,115 trust_remote_code=True,116 low_cpu_mem_usage=True117 )118 119 if device == "cpu":120 model = model.to(device)121 122 model.eval()123 124 logger.info(f"Model loaded successfully on {device}")125 126 except Exception as e:127 logger.error(f"Failed to load model: {e}")128 raise129 130 131@app.get("/", response_model=Dict[str, str])132async def root():133 """Root endpoint"""134 return {135 "message": "Helion-OSC API Server",136 "version": "1.0.0",137 "documentation": "/docs"138 }139 140 141@app.get("/health", response_model=HealthResponse)142async def health_check():143 """Health check endpoint"""144 return HealthResponse(145 status="healthy" if model is not None else "unhealthy",146 model_loaded=model is not None,147 device=device,148 timestamp=time.time()149 )150 151 152@app.get("/info", response_model=ModelInfo)153async def model_info():154 """Get model information"""155 if model is None:156 raise HTTPException(status_code=503, detail="Model not loaded")157 158 config = model.config159 160 return ModelInfo(161 model_name="DeepXR/Helion-OSC",162 model_type=config.model_type,163 vocabulary_size=config.vocab_size,164 hidden_size=config.hidden_size,165 num_layers=config.num_hidden_layers,166 device=device,167 dtype=str(next(model.parameters()).dtype),168 max_position_embeddings=config.max_position_embeddings169 )170 171 172@app.post("/generate", response_model=GenerationResponse)173async def generate(request: GenerationRequest):174 """Generate text based on prompt"""175 if model is None or tokenizer is None:176 raise HTTPException(status_code=503, detail="Model not loaded")177 178 if request.stream:179 raise HTTPException(180 status_code=400,181 detail="Use /generate/stream endpoint for streaming responses"182 )183 184 start_time = time.time()185 186 try:187 # Tokenize input188 inputs = tokenizer(request.prompt, return_tensors="pt").to(device)189 input_length = inputs.input_ids.shape[1]190 191 # Generate192 with torch.no_grad():193 outputs = model.generate(194 **inputs,195 max_length=request.max_length,196 temperature=request.temperature,197 top_p=request.top_p,198 top_k=request.top_k,199 repetition_penalty=request.repetition_penalty,200 do_sample=request.do_sample,201 num_return_sequences=request.num_return_sequences,202 pad_token_id=tokenizer.pad_token_id,203 eos_token_id=tokenizer.eos_token_id204 )205 206 # Decode output207 generated_text = tokenizer.decode(outputs[0], skip_special_tokens=True)208 209 # Remove prompt from output210 generated_text = generated_text[len(request.prompt):].strip()211 212 generation_time = time.time() - start_time213 tokens_generated = outputs.shape[1] - input_length214 215 return GenerationResponse(216 generated_text=generated_text,217 prompt=request.prompt,218 model="DeepXR/Helion-OSC",219 generation_time=generation_time,220 tokens_generated=tokens_generated221 )222 223 except Exception as e:224 logger.error(f"Generation error: {e}")225 raise HTTPException(status_code=500, detail=str(e))226 227 228@app.post("/generate/stream")229async def generate_stream(request: GenerationRequest):230 """Generate text with streaming response"""231 if model is None or tokenizer is None:232 raise HTTPException(status_code=503, detail="Model not loaded")233 234 async def stream_generator() -> AsyncGenerator[str, None]:235 try:236 # Tokenize input237 inputs = tokenizer(request.prompt, return_tensors="pt").to(device)238 239 # Setup streamer240 streamer = TextIteratorStreamer(241 tokenizer,242 skip_prompt=True,243 skip_special_tokens=True244 )245 246 # Generation kwargs247 generation_kwargs = {248 **inputs,249 "max_length": request.max_length,250 "temperature": request.temperature,251 "top_p": request.top_p,252 "top_k": request.top_k,253 "repetition_penalty": request.repetition_penalty,254 "do_sample": request.do_sample,255 "pad_token_id": tokenizer.pad_token_id,256 "eos_token_id": tokenizer.eos_token_id,257 "streamer": streamer258 }259 260 # Start generation in separate thread261 thread = Thread(target=model.generate, kwargs=generation_kwargs)262 thread.start()263 264 # Stream tokens265 for text in streamer:266 yield f"data: {json.dumps({'text': text})}\n\n"267 await asyncio.sleep(0) # Allow other tasks to run268 269 yield f"data: {json.dumps({'done': True})}\n\n"270 271 except Exception as e:272 logger.error(f"Streaming error: {e}")273 yield f"data: {json.dumps({'error': str(e)})}\n\n"274 275 return StreamingResponse(276 stream_generator(),277 media_type="text/event-stream"278 )279 280 281@app.post("/code/complete")282async def code_complete(283 code: str,284 language: Optional[str] = "python",285 max_length: int = 1024286):287 """Code completion endpoint"""288 if model is None or tokenizer is None:289 raise HTTPException(status_code=503, detail="Model not loaded")290 291 request = GenerationRequest(292 prompt=code,293 max_length=max_length,294 temperature=0.6,295 top_p=0.92,296 do_sample=True,297 task_type="code_completion"298 )299 300 return await generate(request)301 302 303@app.post("/code/explain")304async def code_explain(code: str, language: Optional[str] = "python"):305 """Code explanation endpoint"""306 if model is None or tokenizer is None:307 raise HTTPException(status_code=503, detail="Model not loaded")308 309 prompt = f"Explain the following {language} code in detail:\n\n```{language}\n{code}\n```\n\nExplanation:"310 311 request = GenerationRequest(312 prompt=prompt,313 max_length=2048,314 temperature=0.6,315 top_p=0.9,316 do_sample=True,317 task_type="code_explanation"318 )319 320 return await generate(request)321 322 323@app.post("/code/debug")324async def code_debug(325 code: str,326 error_message: Optional[str] = None,327 language: Optional[str] = "python"328):329 """Code debugging endpoint"""330 if model is None or tokenizer is None:331 raise HTTPException(status_code=503, detail="Model not loaded")332 333 prompt = f"Debug the following {language} code:\n\n```{language}\n{code}\n```"334 if error_message:335 prompt += f"\n\nError message: {error_message}"336 prompt += "\n\nProvide a detailed analysis and fixed code:"337 338 request = GenerationRequest(339 prompt=prompt,340 max_length=2048,341 temperature=0.4,342 top_p=0.88,343 do_sample=False,344 task_type="debugging"345 )346 347 return await generate(request)348 349 350@app.post("/math/solve")351async def math_solve(problem: str):352 """Mathematical problem solving endpoint"""353 if model is None or tokenizer is None:354 raise HTTPException(status_code=503, detail="Model not loaded")355 356 prompt = f"Solve the following mathematical problem step by step:\n\n{problem}\n\nSolution:"357 358 request = GenerationRequest(359 prompt=prompt,360 max_length=2048,361 temperature=0.3,362 top_p=0.9,363 do_sample=False,364 task_type="mathematical_reasoning"365 )366 367 return await generate(request)368 369 370@app.post("/algorithm/design")371async def algorithm_design(372 problem: str,373 include_complexity: bool = True374):375 """Algorithm design endpoint"""376 if model is None or tokenizer is None:377 raise HTTPException(status_code=503, detail="Model not loaded")378 379 prompt = f"Design an efficient algorithm for the following problem:\n\n{problem}"380 if include_complexity:381 prompt += "\n\nInclude time and space complexity analysis."382 383 request = GenerationRequest(384 prompt=prompt,385 max_length=3072,386 temperature=0.5,387 top_p=0.93,388 do_sample=True,389 task_type="algorithm_design"390 )391 392 return await generate(request)393 394 395def main():396 """Run the API server"""397 import argparse398 399 parser = argparse.ArgumentParser(description="Helion-OSC API Server")400 parser.add_argument("--host", type=str, default="0.0.0.0", help="Host to bind to")401 parser.add_argument("--port", type=int, default=8000, help="Port to bind to")402 parser.add_argument("--reload", action="store_true", help="Enable auto-reload")403 parser.add_argument("--workers", type=int, default=1, help="Number of worker processes")404 405 args = parser.parse_args()406 407 logger.info(f"Starting Helion-OSC API Server on {args.host}:{args.port}")408 409 uvicorn.run(410 "api_server:app",411 host=args.host,412 port=args.port,413 reload=args.reload,414 workers=args.workers415 )416 417 418if __name__ == "__main__":419 main()