108 lines
3.4 KiB
Python
108 lines
3.4 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import re
|
|
from dataclasses import dataclass, field
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from .config import ModelConfig
|
|
|
|
|
|
@dataclass
|
|
class ModelfileSpec:
|
|
source: str = ""
|
|
parameters: dict[str, Any] = field(default_factory=dict)
|
|
system: str | None = None
|
|
template: str | None = None
|
|
messages: list[dict[str, str]] = field(default_factory=list)
|
|
adapters: list[str] = field(default_factory=list)
|
|
licenses: list[str] = field(default_factory=list)
|
|
|
|
|
|
def _value(raw: str) -> Any:
|
|
try:
|
|
return json.loads(raw)
|
|
except json.JSONDecodeError:
|
|
return raw
|
|
|
|
|
|
def parse_modelfile(text: str) -> ModelfileSpec:
|
|
spec = ModelfileSpec()
|
|
block_pattern = re.compile(r"(?ms)^\s*(SYSTEM|TEMPLATE|LICENSE)\s+\"\"\"(.*?)\"\"\"\s*$")
|
|
|
|
def take_block(match: re.Match[str]) -> str:
|
|
key, value = match.group(1), match.group(2).strip("\r\n")
|
|
if key == "SYSTEM":
|
|
spec.system = value
|
|
elif key == "TEMPLATE":
|
|
spec.template = value
|
|
else:
|
|
spec.licenses.append(value)
|
|
return ""
|
|
|
|
remaining = block_pattern.sub(take_block, text)
|
|
for line in remaining.splitlines():
|
|
line = line.strip()
|
|
if not line or line.startswith("#"):
|
|
continue
|
|
parts = line.split(maxsplit=2)
|
|
instruction = parts[0].upper()
|
|
if instruction == "FROM" and len(parts) >= 2:
|
|
source = line[len(parts[0]):].strip()
|
|
spec.source = source.strip('"\'')
|
|
elif instruction == "PARAMETER" and len(parts) >= 3:
|
|
key, raw = parts[1], parts[2]
|
|
value = _value(raw)
|
|
if key == "stop":
|
|
spec.parameters.setdefault("stop", []).append(str(value))
|
|
else:
|
|
spec.parameters[key] = value
|
|
elif instruction == "MESSAGE" and len(parts) >= 3:
|
|
spec.messages.append({"role": parts[1].lower(), "content": parts[2]})
|
|
elif instruction == "ADAPTER" and len(parts) >= 2:
|
|
spec.adapters.append(parts[1])
|
|
return spec
|
|
|
|
|
|
def model_config_from_modelfile(model_id: str, text: str) -> ModelConfig:
|
|
parsed = parse_modelfile(text)
|
|
if not parsed.source:
|
|
raise ValueError("Modelfile 缺少 FROM")
|
|
source_path = Path(parsed.source)
|
|
if source_path.exists() or source_path.suffix.lower() in {".gguf", ".onnx", ".safetensors", ".bin", ".pt", ".pth"}:
|
|
path: str | None = parsed.source
|
|
kind = "auto"
|
|
remote_model = None
|
|
else:
|
|
path = None
|
|
kind = "ollama"
|
|
remote_model = parsed.source
|
|
parameters: dict[str, Any] = {}
|
|
options: dict[str, Any] = {}
|
|
aliases = {
|
|
"num_predict": "max_tokens",
|
|
"num_ctx": "ctx_size",
|
|
"num_batch": "batch_size",
|
|
"num_gpu": "n_gpu_layers",
|
|
}
|
|
for key, value in parsed.parameters.items():
|
|
target = aliases.get(key, key)
|
|
if target in {"ctx_size", "batch_size", "n_gpu_layers"}:
|
|
options[target] = value
|
|
else:
|
|
parameters[target] = value
|
|
if parsed.adapters:
|
|
options["adapters"] = parsed.adapters
|
|
return ModelConfig(
|
|
id=model_id,
|
|
kind=kind,
|
|
path=path,
|
|
remote_model=remote_model,
|
|
template=parsed.template,
|
|
system=parsed.system,
|
|
messages=parsed.messages,
|
|
parameters=parameters,
|
|
options=options,
|
|
)
|