added the inference fetching from model mappers
This commit is contained in:
parent
625bc1bbc6
commit
9e50e167d9
4 changed files with 81 additions and 0 deletions
|
|
@ -43,6 +43,7 @@ class LoadResponse(BaseModel):
|
|||
display_name: str = Field(..., description="Display name of the model")
|
||||
is_vision: bool = Field(False, description="Whether model is a vision model")
|
||||
is_lora: bool = Field(False, description="Whether model is a LoRA adapter")
|
||||
inference: dict = Field(..., description="Inference parameters (temperature, top_p, top_k, min_p)")
|
||||
|
||||
|
||||
class UnloadResponse(BaseModel):
|
||||
|
|
|
|||
|
|
@ -22,12 +22,14 @@ if str(backend_path) not in sys.path:
|
|||
try:
|
||||
from core.inference import get_inference_backend
|
||||
from utils.models import ModelConfig
|
||||
from utils.inference import load_inference_config
|
||||
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
|
||||
from utils.inference import load_inference_config
|
||||
|
||||
from models.inference import (
|
||||
LoadRequest,
|
||||
|
|
@ -64,6 +66,8 @@ async def load_model(request: LoadRequest):
|
|||
Load a model for inference.
|
||||
|
||||
The model_path should be a clean identifier from GET /models/list.
|
||||
Returns inference configuration parameters (temperature, top_p, top_k, min_p)
|
||||
from the model's YAML config, falling back to default.yaml for missing values.
|
||||
"""
|
||||
try:
|
||||
backend = get_inference_backend()
|
||||
|
|
@ -97,12 +101,16 @@ async def load_model(request: LoadRequest):
|
|||
|
||||
logger.info(f"Loaded model: {config.identifier}")
|
||||
|
||||
# Load inference configuration parameters
|
||||
inference_config = load_inference_config(config.identifier)
|
||||
|
||||
return LoadResponse(
|
||||
status="loaded",
|
||||
model=config.identifier,
|
||||
display_name=config.display_name,
|
||||
is_vision=config.is_vision,
|
||||
is_lora=config.is_lora,
|
||||
inference=inference_config,
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
|
|
|
|||
7
studio/backend/utils/inference/__init__.py
Normal file
7
studio/backend/utils/inference/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""
|
||||
Inference utility functions
|
||||
"""
|
||||
from utils.inference.inference_config import load_inference_config
|
||||
|
||||
__all__ = ["load_inference_config"]
|
||||
|
||||
65
studio/backend/utils/inference/inference_config.py
Normal file
65
studio/backend/utils/inference/inference_config.py
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
"""
|
||||
Inference configuration loading utilities.
|
||||
|
||||
This module provides functions to load inference parameters (temperature, top_p, top_k, min_p)
|
||||
from model YAML configuration files, with fallback to default.yaml.
|
||||
"""
|
||||
from pathlib import Path
|
||||
from typing import Dict, Any
|
||||
import yaml
|
||||
import logging
|
||||
|
||||
from utils.models.model_config import load_model_defaults
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def load_inference_config(model_identifier: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Load inference configuration parameters for a model.
|
||||
|
||||
This function loads inference parameters (temperature, top_p, top_k, min_p) from the
|
||||
model's YAML configuration file using the same mapping logic as the /config endpoint.
|
||||
If a parameter is missing from the model's config, it falls back to the value in
|
||||
default.yaml.
|
||||
|
||||
Args:
|
||||
model_identifier: Model identifier (e.g., "unsloth/llama-3-8b-bnb-4bit")
|
||||
|
||||
Returns:
|
||||
Dictionary containing inference parameters:
|
||||
{
|
||||
"temperature": float,
|
||||
"top_p": float,
|
||||
"top_k": int,
|
||||
"min_p": float
|
||||
}
|
||||
"""
|
||||
# Load model defaults to get inference parameters
|
||||
model_defaults = load_model_defaults(model_identifier)
|
||||
|
||||
# Load default.yaml for fallback values
|
||||
script_dir = Path(__file__).parent.parent.parent
|
||||
defaults_dir = script_dir / "assets" / "configs" / "model_defaults"
|
||||
default_config_path = defaults_dir / "default.yaml"
|
||||
|
||||
default_inference = {}
|
||||
if default_config_path.exists():
|
||||
try:
|
||||
with open(default_config_path, 'r', encoding='utf-8') as f:
|
||||
default_config = yaml.safe_load(f) or {}
|
||||
default_inference = default_config.get("inference", {})
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load default.yaml: {e}")
|
||||
|
||||
# Extract inference parameters from model config, fallback to defaults
|
||||
model_inference = model_defaults.get("inference", {})
|
||||
inference_config = {
|
||||
"temperature": model_inference.get("temperature", default_inference.get("temperature", 0.7)),
|
||||
"top_p": model_inference.get("top_p", default_inference.get("top_p", 0.95)),
|
||||
"top_k": model_inference.get("top_k", default_inference.get("top_k", -1)),
|
||||
"min_p": model_inference.get("min_p", default_inference.get("min_p", 0.01)),
|
||||
}
|
||||
|
||||
return inference_config
|
||||
|
||||
Loading…
Add table
Add a link
Reference in a new issue