From 9e50e167d9b4074cbe1246447f53c2cdd04f445a Mon Sep 17 00:00:00 2001 From: sshah229 Date: Sun, 15 Feb 2026 02:48:53 -0700 Subject: [PATCH] added the inference fetching from model mappers --- studio/backend/models/inference.py | 1 + studio/backend/routes/inference.py | 8 +++ studio/backend/utils/inference/__init__.py | 7 ++ .../utils/inference/inference_config.py | 65 +++++++++++++++++++ 4 files changed, 81 insertions(+) create mode 100644 studio/backend/utils/inference/__init__.py create mode 100644 studio/backend/utils/inference/inference_config.py diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index ada8bdd539..b3924c5569 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -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): diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index b3ea1fa1be..ae8fc31a46 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -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: diff --git a/studio/backend/utils/inference/__init__.py b/studio/backend/utils/inference/__init__.py new file mode 100644 index 0000000000..660643a436 --- /dev/null +++ b/studio/backend/utils/inference/__init__.py @@ -0,0 +1,7 @@ +""" +Inference utility functions +""" +from utils.inference.inference_config import load_inference_config + +__all__ = ["load_inference_config"] + diff --git a/studio/backend/utils/inference/inference_config.py b/studio/backend/utils/inference/inference_config.py new file mode 100644 index 0000000000..d6de562d33 --- /dev/null +++ b/studio/backend/utils/inference/inference_config.py @@ -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 +