refactor/inference-api-routes-part-1
This commit is contained in:
parent
6df5b1eded
commit
b4ec0389f0
5 changed files with 334 additions and 100 deletions
|
|
@ -847,24 +847,16 @@ class InferenceBackend:
|
|||
formatted += "Assistant: "
|
||||
return formatted
|
||||
|
||||
def check_vision_model_compatibility(self, show_warning: bool = True) -> bool:
|
||||
def check_vision_model_compatibility(self) -> bool:
|
||||
"""
|
||||
Check if current model supports vision and optionally show warning if image uploaded to text-only model
|
||||
|
||||
Args:
|
||||
show_warning: Whether to show Gradio warning if vision not supported
|
||||
Check if current model supports vision.
|
||||
|
||||
Returns:
|
||||
bool: True if current model supports vision, False otherwise
|
||||
"""
|
||||
current_model = self.get_current_model()
|
||||
if current_model and current_model in self.models:
|
||||
is_vision = self.models[current_model].get("is_vision", False)
|
||||
if not is_vision and show_warning:
|
||||
import gradio as gr
|
||||
model_short = current_model.split('/')[-1] if '/' in current_model else current_model
|
||||
gr.Warning(f"Image uploaded, but {model_short} is a text-only model. Please select a vision model to analyze images.")
|
||||
return is_vision
|
||||
return self.models[current_model].get("is_vision", False)
|
||||
return False
|
||||
|
||||
def _reset_model_generation_state(self, model_name: str):
|
||||
|
|
@ -1066,94 +1058,7 @@ class InferenceBackend:
|
|||
logger.error(f"Error in load_model_simple: {e}")
|
||||
return False
|
||||
|
||||
def add_local_model_to_dropdown(self, model_path: str):
|
||||
"""Add successfully loaded local model to dropdown storage"""
|
||||
try:
|
||||
from pathlib import Path
|
||||
|
||||
path_obj = Path(model_path)
|
||||
display_name = f"{path_obj.name}"
|
||||
|
||||
# Check if already exists
|
||||
for existing_display, existing_path in self.loaded_local_models:
|
||||
if existing_path == model_path:
|
||||
logger.debug(f"Local model already in dropdown: {model_path}")
|
||||
return
|
||||
|
||||
# Add to beginning of list
|
||||
self.loaded_local_models.insert(0, (display_name, model_path))
|
||||
logger.info(f"Added local model to dropdown: {display_name} -> {model_path}")
|
||||
|
||||
# Keep only last 5
|
||||
if len(self.loaded_local_models) > 5:
|
||||
self.loaded_local_models.pop()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error adding local model to dropdown: {e}")
|
||||
|
||||
def get_model_dropdown_choices(self, models: list = None) -> list:
|
||||
"""Get model dropdown choices with status indicators"""
|
||||
if models is None:
|
||||
models = self.default_models
|
||||
|
||||
try:
|
||||
active_model = self.active_model_name
|
||||
loading_model = self.get_loading_model()
|
||||
|
||||
choices = []
|
||||
|
||||
# Add local models first
|
||||
for local_display, local_path in self.loaded_local_models:
|
||||
if local_path == active_model:
|
||||
choices.append((f"{local_display} (Active)", local_path))
|
||||
else:
|
||||
choices.append((local_display, local_path))
|
||||
|
||||
# Add default models
|
||||
for model in models:
|
||||
short_name = model.split('/')[-1] if '/' in model else model
|
||||
|
||||
if model == active_model:
|
||||
display_name = f"{short_name} (Active)"
|
||||
elif model == loading_model:
|
||||
display_name = f"{short_name} (Loading...)"
|
||||
elif model in self.models and self.models[model].get("model"):
|
||||
# Model is loaded in memory
|
||||
display_name = f"{short_name} (Ready)"
|
||||
# elif model in self.models:
|
||||
# display_name = f"{short_name} (Ready)"
|
||||
elif is_model_cached(model):
|
||||
# Model is downloaded but not loaded
|
||||
display_name = f"{short_name} (Cached)"
|
||||
else:
|
||||
display_name = f"↓ {short_name}" # Not downloaded
|
||||
|
||||
choices.append((display_name, model))
|
||||
|
||||
return choices
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting model choices: {e}")
|
||||
return [(model.split('/')[-1], model) for model in models]
|
||||
|
||||
|
||||
def update_model_dropdown(self, models: list = None):
|
||||
"""Update model dropdown with current status"""
|
||||
try:
|
||||
import gradio as gr
|
||||
|
||||
choices = self.get_model_dropdown_choices(models)
|
||||
active_model = self.active_model_name
|
||||
|
||||
# Set value to active model if exists
|
||||
value = active_model if active_model else (choices[0][1] if choices else None)
|
||||
|
||||
return gr.update(choices=choices, value=value)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating model dropdown: {e}")
|
||||
import gradio as gr
|
||||
return gr.update()
|
||||
|
||||
def load_model_simple(self,
|
||||
model_path: str,
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from pathlib import Path
|
|||
from datetime import datetime
|
||||
|
||||
# Import routers
|
||||
from routes import training_router, models_router
|
||||
from routes import training_router, models_router, inference_router
|
||||
|
||||
# Create FastAPI app
|
||||
app = FastAPI(
|
||||
|
|
@ -32,6 +32,7 @@ app.add_middleware(
|
|||
# Register routers
|
||||
app.include_router(training_router, prefix="/api/train", tags=["training"])
|
||||
app.include_router(models_router, prefix="/api/models", tags=["models"])
|
||||
app.include_router(inference_router, prefix="/api/inference", tags=["inference"])
|
||||
|
||||
|
||||
# ============ Health and System Endpoints ============
|
||||
|
|
|
|||
|
|
@ -4,5 +4,6 @@ API Routes
|
|||
|
||||
from routes.training import router as training_router
|
||||
from routes.models import router as models_router
|
||||
from routes.inference import router as inference_router
|
||||
|
||||
__all__ = ["training_router", "models_router"]
|
||||
__all__ = ["training_router", "models_router", "inference_router"]
|
||||
|
|
@ -0,0 +1,267 @@
|
|||
"""
|
||||
Inference API routes for model loading and text generation.
|
||||
"""
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel, Field
|
||||
from typing import Optional, List
|
||||
import json
|
||||
import logging
|
||||
|
||||
# Add backend directory to path
|
||||
backend_path = Path(__file__).parent.parent.parent
|
||||
if str(backend_path) not in sys.path:
|
||||
sys.path.insert(0, str(backend_path))
|
||||
|
||||
# Import backend functions
|
||||
try:
|
||||
from core.inference import get_inference_backend
|
||||
from utils.models import ModelConfig
|
||||
except ImportError:
|
||||
parent_backend = backend_path.parent / "backend"
|
||||
if str(parent_backend) not in sys.path:
|
||||
sys.path.insert(0, str(parent_backend))
|
||||
from core.inference import get_inference_backend
|
||||
from utils.models import ModelConfig
|
||||
|
||||
router = APIRouter()
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Configure logger
|
||||
if not logger.handlers:
|
||||
handler = logging.StreamHandler()
|
||||
handler.setLevel(logging.INFO)
|
||||
formatter = logging.Formatter('%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||||
handler.setFormatter(formatter)
|
||||
logger.addHandler(handler)
|
||||
logger.setLevel(logging.INFO)
|
||||
|
||||
|
||||
# ============================================
|
||||
# Request/Response Models
|
||||
# ============================================
|
||||
|
||||
class LoadRequest(BaseModel):
|
||||
"""Request to load a model for inference"""
|
||||
model_path: str = Field(..., description="Model identifier or local path")
|
||||
hf_token: Optional[str] = Field(None, description="HuggingFace token for gated models")
|
||||
max_seq_length: int = Field(2048, ge=128, le=32768, description="Maximum sequence length")
|
||||
load_in_4bit: bool = Field(True, description="Load model in 4-bit quantization")
|
||||
is_lora: bool = Field(False, description="Whether this is a LoRA adapter")
|
||||
|
||||
|
||||
class UnloadRequest(BaseModel):
|
||||
"""Request to unload a model"""
|
||||
model_path: str = Field(..., description="Model identifier to unload")
|
||||
|
||||
|
||||
class GenerateRequest(BaseModel):
|
||||
"""Request for text generation"""
|
||||
messages: List[dict] = Field(..., description="Chat messages in OpenAI format")
|
||||
system_prompt: str = Field("You are a helpful AI assistant.", description="System prompt")
|
||||
temperature: float = Field(0.7, ge=0.0, le=2.0, description="Sampling temperature")
|
||||
top_p: float = Field(0.9, ge=0.0, le=1.0, description="Top-p sampling")
|
||||
top_k: int = Field(40, ge=1, le=100, description="Top-k sampling")
|
||||
max_new_tokens: int = Field(512, ge=1, le=4096, description="Maximum tokens to generate")
|
||||
repetition_penalty: float = Field(1.1, ge=1.0, le=2.0, description="Repetition penalty")
|
||||
image_base64: Optional[str] = Field(None, description="Base64 encoded image for vision models")
|
||||
|
||||
|
||||
class LoadResponse(BaseModel):
|
||||
"""Response after loading a model"""
|
||||
status: str
|
||||
model: str
|
||||
display_name: str
|
||||
is_vision: bool
|
||||
is_lora: bool
|
||||
|
||||
|
||||
class StatusResponse(BaseModel):
|
||||
"""Current inference backend status"""
|
||||
active_model: Optional[str]
|
||||
is_vision: bool
|
||||
loading: List[str]
|
||||
loaded: List[str]
|
||||
|
||||
|
||||
# ============================================
|
||||
# Routes
|
||||
# ============================================
|
||||
|
||||
@router.post("/load", response_model=LoadResponse)
|
||||
async def load_model(request: LoadRequest):
|
||||
"""
|
||||
Load a model for inference.
|
||||
|
||||
The model_path should be a clean identifier from GET /models/list.
|
||||
"""
|
||||
try:
|
||||
backend = get_inference_backend()
|
||||
|
||||
# Create config using clean factory method
|
||||
config = ModelConfig.from_identifier(
|
||||
model_id=request.model_path,
|
||||
hf_token=request.hf_token,
|
||||
is_lora=request.is_lora,
|
||||
)
|
||||
|
||||
if not config:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Invalid model identifier: {request.model_path}"
|
||||
)
|
||||
|
||||
# Load the model
|
||||
success = backend.load_model(
|
||||
config=config,
|
||||
max_seq_length=request.max_seq_length,
|
||||
load_in_4bit=request.load_in_4bit,
|
||||
hf_token=request.hf_token,
|
||||
)
|
||||
|
||||
if not success:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"Failed to load model: {config.display_name}"
|
||||
)
|
||||
|
||||
logger.info(f"Loaded model: {config.identifier}")
|
||||
|
||||
return LoadResponse(
|
||||
status="loaded",
|
||||
model=config.identifier,
|
||||
display_name=config.display_name,
|
||||
is_vision=config.is_vision,
|
||||
is_lora=config.is_lora,
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Error loading model: {e}", exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"Failed to load model: {str(e)}"
|
||||
)
|
||||
|
||||
|
||||
@router.post("/unload")
|
||||
async def unload_model(request: UnloadRequest):
|
||||
"""
|
||||
Unload a model from memory.
|
||||
"""
|
||||
try:
|
||||
backend = get_inference_backend()
|
||||
backend.unload_model(request.model_path)
|
||||
logger.info(f"Unloaded model: {request.model_path}")
|
||||
return {"status": "unloaded", "model": request.model_path}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error unloading model: {e}", exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"Failed to unload model: {str(e)}"
|
||||
)
|
||||
|
||||
|
||||
@router.post("/generate/stream")
|
||||
async def generate_stream(request: GenerateRequest):
|
||||
"""
|
||||
Generate a chat response with Server-Sent Events (SSE) streaming.
|
||||
|
||||
For vision models, provide image_base64 with the base64-encoded image.
|
||||
"""
|
||||
backend = get_inference_backend()
|
||||
|
||||
if not backend.active_model_name:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="No model loaded. Call POST /inference/load first."
|
||||
)
|
||||
|
||||
# Decode image if provided (for vision models)
|
||||
image = None
|
||||
if request.image_base64:
|
||||
try:
|
||||
import base64
|
||||
from PIL import Image
|
||||
from io import BytesIO
|
||||
|
||||
# Check if current model supports vision
|
||||
model_info = backend.models.get(backend.active_model_name, {})
|
||||
if not model_info.get("is_vision"):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Image provided but current model is text-only. Load a vision model."
|
||||
)
|
||||
|
||||
image_data = base64.b64decode(request.image_base64)
|
||||
image = Image.open(BytesIO(image_data))
|
||||
image = backend.resize_image(image)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Failed to decode image: {str(e)}"
|
||||
)
|
||||
|
||||
async def stream():
|
||||
try:
|
||||
for chunk in backend.generate_chat_response(
|
||||
messages=request.messages,
|
||||
system_prompt=request.system_prompt,
|
||||
image=image,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
max_new_tokens=request.max_new_tokens,
|
||||
repetition_penalty=request.repetition_penalty,
|
||||
):
|
||||
yield f"data: {json.dumps({'content': chunk})}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
except Exception as e:
|
||||
backend.reset_generation_state()
|
||||
logger.error(f"Error during generation: {e}", exc_info=True)
|
||||
yield f"data: {json.dumps({'error': str(e)})}\n\n"
|
||||
|
||||
return StreamingResponse(
|
||||
stream(),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Connection": "keep-alive",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.get("/status", response_model=StatusResponse)
|
||||
async def get_status():
|
||||
"""
|
||||
Get current inference backend status.
|
||||
"""
|
||||
try:
|
||||
backend = get_inference_backend()
|
||||
|
||||
is_vision = False
|
||||
if backend.active_model_name:
|
||||
model_info = backend.models.get(backend.active_model_name, {})
|
||||
is_vision = model_info.get("is_vision", False)
|
||||
|
||||
return StatusResponse(
|
||||
active_model=backend.active_model_name,
|
||||
is_vision=is_vision,
|
||||
loading=list(getattr(backend, 'loading_models', set())),
|
||||
loaded=list(backend.models.keys()),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting status: {e}", exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"Failed to get status: {str(e)}"
|
||||
)
|
||||
|
|
@ -625,6 +625,66 @@ class ModelConfig:
|
|||
logger.error(f"Error creating ModelConfig from LoRA path: {e}")
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def from_identifier(
|
||||
cls,
|
||||
model_id: str,
|
||||
hf_token: Optional[str] = None,
|
||||
is_lora: bool = False
|
||||
) -> Optional['ModelConfig']:
|
||||
"""
|
||||
Create ModelConfig from a clean model identifier.
|
||||
|
||||
For FastAPI routes where the frontend sends sanitized model paths.
|
||||
No Gradio dropdown parsing - expects clean identifiers like:
|
||||
- "unsloth/Meta-Llama-3.1-8B-Instruct-bnb-4bit"
|
||||
- "./outputs/my_lora_adapter"
|
||||
- "/absolute/path/to/model"
|
||||
|
||||
Args:
|
||||
model_id: Clean model identifier (HF repo name or local path)
|
||||
hf_token: Optional HF token for vision detection on gated models
|
||||
is_lora: Whether this is a LoRA adapter
|
||||
|
||||
Returns:
|
||||
ModelConfig or None if configuration cannot be created
|
||||
"""
|
||||
if not model_id or not model_id.strip():
|
||||
return None
|
||||
|
||||
identifier = model_id.strip()
|
||||
is_local = is_local_path(identifier)
|
||||
path = normalize_path(identifier) if is_local else identifier
|
||||
|
||||
# Add unsloth/ prefix for shorthand HF models
|
||||
if not is_local and "/" not in identifier:
|
||||
identifier = f"unsloth/{identifier}"
|
||||
path = identifier
|
||||
|
||||
# Handle LoRA adapters
|
||||
base_model = None
|
||||
if is_lora:
|
||||
base_model = get_base_model_from_lora(path)
|
||||
if not base_model:
|
||||
logger.warning(f"Could not determine base model for LoRA '{path}'")
|
||||
return None
|
||||
vision = is_vision_model(base_model, hf_token=hf_token)
|
||||
else:
|
||||
vision = is_vision_model(identifier, hf_token=hf_token)
|
||||
|
||||
display_name = Path(path).name if is_local else identifier.split("/")[-1]
|
||||
|
||||
return cls(
|
||||
identifier=identifier,
|
||||
display_name=display_name,
|
||||
path=path,
|
||||
is_local=is_local,
|
||||
is_cached=is_model_cached(identifier) if not is_local else True,
|
||||
is_vision=vision,
|
||||
is_lora=is_lora,
|
||||
base_model=base_model,
|
||||
)
|
||||
|
||||
|
||||
@classmethod
|
||||
def from_ui_selection(cls,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue