added the inference fetching from model mappers

This commit is contained in:
sshah229 2026-02-15 02:48:53 -07:00
commit 9e50e167d9
4 changed files with 81 additions and 0 deletions

View file

@ -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):

View file

@ -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:

View file

@ -0,0 +1,7 @@
"""
Inference utility functions
"""
from utils.inference.inference_config import load_inference_config
__all__ = ["load_inference_config"]

View 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