Files

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,
)