Initial commit: LocalPilot:本地模型运行时,复用 llama-server 并提供 Ollama 兼容 provider
This commit is contained in:
@@ -0,0 +1,135 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
|
||||
class RuntimeConfig(BaseModel):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
llama_server: str = r"E:\ollama\模型\llama-b8642-bin-win-cuda-13.1-x64\llama-server.exe"
|
||||
ollama_url: str = "http://127.0.0.1:11434"
|
||||
host: str = "127.0.0.1"
|
||||
port: int = 8787
|
||||
default_model: str | None = "apex-qwen3"
|
||||
model_dir: str = r"E:\model"
|
||||
runtime_dir: str = r"E:\llm-backend\runtime"
|
||||
cache_dir: str = r"E:\llm-backend\cache"
|
||||
n_gpu_layers: str | int = "auto"
|
||||
ctx_size: int = 32768
|
||||
batch_size: int = 2048
|
||||
ubatch_size: int = 512
|
||||
flash_attn: str = "auto"
|
||||
fit_vram: bool = True
|
||||
auto_load: bool = False
|
||||
api_key: str | None = None
|
||||
max_loaded_models: int = 1
|
||||
keep_alive_seconds: float = 300.0
|
||||
|
||||
|
||||
class ModelConfig(BaseModel):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
id: str
|
||||
kind: str = "auto"
|
||||
path: str | None = None
|
||||
enabled: bool = True
|
||||
description: str = ""
|
||||
cloud_profile: str | None = None
|
||||
remote_model: str | None = None
|
||||
template: str | None = None
|
||||
system: str | None = None
|
||||
messages: list[dict[str, Any]] = Field(default_factory=list)
|
||||
parameters: dict[str, Any] = Field(default_factory=dict)
|
||||
options: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class CloudProfile(BaseModel):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
id: str
|
||||
base_url: str
|
||||
api_key_env: str | None = None
|
||||
api_key: str | None = None
|
||||
default_model: str | None = None
|
||||
timeout: float = 120.0
|
||||
headers: dict[str, str] = Field(default_factory=dict)
|
||||
|
||||
def resolved_api_key(self) -> str | None:
|
||||
if self.api_key:
|
||||
return self.api_key
|
||||
if self.api_key_env:
|
||||
return os.getenv(self.api_key_env)
|
||||
return None
|
||||
|
||||
|
||||
class AppConfig(BaseModel):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
runtime: RuntimeConfig = Field(default_factory=RuntimeConfig)
|
||||
models: list[ModelConfig] = Field(default_factory=list)
|
||||
cloud_profiles: list[CloudProfile] = Field(default_factory=list)
|
||||
|
||||
def model_by_id(self, model_id: str) -> ModelConfig:
|
||||
for model in self.models:
|
||||
if model.id == model_id and model.enabled:
|
||||
return model
|
||||
raise KeyError(f"模型未配置或已禁用: {model_id}")
|
||||
|
||||
def cloud_by_id(self, profile_id: str) -> CloudProfile:
|
||||
for profile in self.cloud_profiles:
|
||||
if profile.id == profile_id:
|
||||
return profile
|
||||
raise KeyError(f"云端配置不存在: {profile_id}")
|
||||
|
||||
|
||||
def infer_kind(model: ModelConfig) -> str:
|
||||
if model.kind != "auto":
|
||||
return {
|
||||
"safetensors": "transformers",
|
||||
"hf": "transformers",
|
||||
"onnxruntime": "onnx",
|
||||
}.get(model.kind, model.kind)
|
||||
if not model.path:
|
||||
return "cloud" if model.cloud_profile else "unknown"
|
||||
path = Path(model.path)
|
||||
suffix = path.suffix.lower()
|
||||
if suffix == ".gguf":
|
||||
return "gguf"
|
||||
if suffix == ".onnx":
|
||||
return "onnx"
|
||||
if suffix in {".safetensors", ".bin", ".pt", ".pth"}:
|
||||
return "transformers"
|
||||
if path.is_dir():
|
||||
if (path / "config.json").exists() and (
|
||||
list(path.glob("*.safetensors")) or (path / "model.safetensors.index.json").exists()
|
||||
):
|
||||
return "transformers"
|
||||
if list(path.glob("*.onnx")):
|
||||
return "onnx"
|
||||
return "unknown"
|
||||
|
||||
|
||||
def load_config(path: str | Path) -> AppConfig:
|
||||
config_path = Path(path)
|
||||
if not config_path.exists():
|
||||
return AppConfig()
|
||||
with config_path.open("r", encoding="utf-8") as handle:
|
||||
return AppConfig.model_validate(json.load(handle))
|
||||
|
||||
|
||||
def save_config(path: str | Path, config: AppConfig) -> None:
|
||||
config_path = Path(path)
|
||||
config_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with config_path.open("w", encoding="utf-8") as handle:
|
||||
json.dump(config.model_dump(exclude_none=True), handle, ensure_ascii=False, indent=2)
|
||||
handle.write("\n")
|
||||
|
||||
|
||||
def ensure_runtime_dirs(config: AppConfig) -> None:
|
||||
Path(config.runtime.runtime_dir).mkdir(parents=True, exist_ok=True)
|
||||
Path(config.runtime.cache_dir).mkdir(parents=True, exist_ok=True)
|
||||
Reference in New Issue
Block a user