SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw History Blame Contribute Delete
3.44 kB
from __future__ import annotations
from dataclasses import asdict, dataclass, field
from typing import Any
from adam.model_plugins import ModelPlugin, ModelPluginRegistry
@dataclass(frozen=True, slots=True)
class ModelProfile:
"""Normalized model profile built from ADAM's plugin manifests."""
id: str
name: str
category: str
architecture: str
version: str
description: str
status: str = "experimental"
output_type: str = "image"
training: dict[str, dict[str, Any]] = field(default_factory=dict)
generation: dict[str, dict[str, Any]] = field(default_factory=dict)
trainer_module: str = ""
generator_module: str = ""
trainer_tool: str = ""
generator_tool: str = ""
capabilities: list[str] = field(default_factory=list)
hardware: dict[str, Any] = field(default_factory=dict)
vram_behavior: dict[str, Any] = field(default_factory=dict)
input_formats: list[str] = field(default_factory=list)
output_formats: list[str] = field(default_factory=list)
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def profile_from_plugin(plugin: ModelPlugin) -> ModelProfile:
info = dict(plugin.info)
training_tool = dict(plugin.training_tool)
generation_tool = dict(plugin.generation_tool)
capabilities = list(
dict.fromkeys(
[
*info.get("capabilities", []),
*training_tool.get("capabilities", []),
*generation_tool.get("capabilities", []),
]
)
)
trainer_backend = dict(training_tool.get("backend", {}))
generator_backend = dict(generation_tool.get("backend", {}))
return ModelProfile(
id=plugin.id,
name=str(info.get("name", plugin.name)),
category=str(info.get("category", "")),
architecture=str(info.get("architecture", plugin.id)),
version=str(info.get("version", "")),
description=str(info.get("description", "")),
status=str(info.get("status", "experimental")),
output_type=str(info.get("output_type", "image")),
training=plugin.training_settings,
generation=plugin.generation_settings,
trainer_module=str(trainer_backend.get("module", "")),
generator_module=str(generator_backend.get("module", "")),
trainer_tool=plugin.trainer_id if plugin.training_settings else "",
generator_tool=plugin.generator_id if plugin.generation_settings else "",
capabilities=capabilities,
hardware=dict(info.get("hardware", {})),
vram_behavior=dict(info.get("vram_behavior", {})),
input_formats=list(info.get("input_formats", [])),
output_formats=list(info.get("output_formats", [])),
)
class ModelProfileRegistry:
"""Read-only view over plugin manifests for UI and automation features."""
def __init__(self, plugins: ModelPluginRegistry) -> None:
self.plugins = plugins
def all(self) -> list[ModelProfile]:
return [
profile_from_plugin(plugin)
for plugin in self.plugins.all()
if plugin.info.get("category") != "Template"
]
def get(self, profile_id: str) -> ModelProfile | None:
plugin = self.plugins.by_trainer(profile_id)
return profile_from_plugin(plugin) if plugin else None
def as_catalog(self) -> list[dict[str, Any]]:
return [profile.to_dict() for profile in self.all()]