SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw
History Blame Contribute Delete
19.2 kB
from __future__ import annotations
import importlib
import importlib.util
import json
import logging
import pkgutil
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Callable
REQUIRED_INFO_FIELDS = {"name", "version", "category", "description"}
SUPPORTED_SETTING_TYPES = {
"int",
"float",
"bool",
"choice",
"text",
"multiline_text",
"path",
"folder",
"slider",
}
class ModelPluginError(RuntimeError):
pass
@dataclass(frozen=True, slots=True)
class ModelPlugin:
id: str
info: dict[str, Any]
training_settings: dict[str, dict[str, Any]] = field(default_factory=dict)
generation_settings: dict[str, dict[str, Any]] = field(default_factory=dict)
training_tool: dict[str, Any] = field(default_factory=dict)
generation_tool: dict[str, Any] = field(default_factory=dict)
module_name: str = ""
plugin_path: Path | None = None
@property
def name(self) -> str:
return str(self.info.get("name", self.id))
@property
def trainer_id(self) -> str:
return str(self.training_tool.get("id") or f"{self.id}_trainer")
@property
def generator_id(self) -> str:
return str(self.generation_tool.get("id") or f"{self.id}_generator")
class ModelPluginRegistry:
"""Discovers model plugins and validates their setting schemas."""
def __init__(self, root: Path, logger: logging.Logger | None = None) -> None:
self.root = root.resolve()
self.logger = logger or logging.getLogger(__name__)
self.plugins: dict[str, ModelPlugin] = {}
self.errors: list[str] = []
self.discover()
def discover(self) -> None:
self.plugins = {}
self.errors = []
for module_name in self._candidate_modules():
try:
plugin = self._load_module_plugin(module_name)
except Exception as exc:
message = f"{module_name}: {exc}"
self.errors.append(message)
self.logger.warning("Model plugin failed to load: %s", message)
continue
if plugin.id in self.plugins:
self.errors.append(f"{module_name}: duplicate model plugin id {plugin.id}")
continue
self.plugins[plugin.id] = plugin
def _candidate_modules(self) -> list[str | Path]:
modules: list[str | Path] = []
try:
package = importlib.import_module("adam.model_plugins_builtin")
for item in pkgutil.iter_modules(package.__path__, package.__name__ + "."):
if not item.ispkg:
continue
modules.append(item.name + ".manifest")
except Exception as exc:
self.errors.append(f"adam.model_plugins_builtin: {exc}")
models_dir = self.root / "models"
if models_dir.is_dir():
for folder in sorted(models_dir.iterdir()):
manifest = folder / "manifest.py"
if not folder.is_dir() or not manifest.is_file():
continue
modules.append(manifest)
return modules
def _load_module_plugin(self, module_name: str | Path) -> ModelPlugin:
if isinstance(module_name, Path):
fallback_id = module_name.parent.name
unique_name = f"adam_user_model_{fallback_id}_{abs(hash(str(module_name.resolve())))}"
spec = importlib.util.spec_from_file_location(unique_name, module_name)
if spec is None or spec.loader is None:
raise ModelPluginError(f"Could not load manifest file: {module_name}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
module_label = str(module_name)
else:
module = importlib.import_module(module_name)
fallback_id = module_name.split(".")[-2]
module_label = module_name
plugin_id = str(getattr(module, "PLUGIN_ID", "") or fallback_id)
info = dict(getattr(module, "MODEL_INFO", {}))
missing = REQUIRED_INFO_FIELDS - set(info)
if missing:
raise ModelPluginError(
"MODEL_INFO is missing " + ", ".join(sorted(missing))
)
training_settings = self._validate_schema(
dict(getattr(module, "TRAINING_SETTINGS", {})),
f"{plugin_id} training",
)
generation_settings = self._validate_schema(
dict(getattr(module, "GENERATION_SETTINGS", {})),
f"{plugin_id} generation",
)
plugin_path = Path(getattr(module, "__file__", "")).resolve().parent
return ModelPlugin(
id=plugin_id,
info=info,
training_settings=training_settings,
generation_settings=generation_settings,
training_tool=dict(getattr(module, "TRAINING_TOOL", {})),
generation_tool=dict(getattr(module, "GENERATION_TOOL", {})),
module_name=module_label,
plugin_path=plugin_path,
)
@staticmethod
def _validate_schema(
schema: dict[str, Any],
label: str,
) -> dict[str, dict[str, Any]]:
clean: dict[str, dict[str, Any]] = {}
for key, raw in schema.items():
if not isinstance(raw, dict):
raise ModelPluginError(f"{label} setting {key} must be an object")
spec = dict(raw)
setting_type = str(spec.get("type", "text"))
if setting_type not in SUPPORTED_SETTING_TYPES:
raise ModelPluginError(
f"{label} setting {key} has unsupported type {setting_type}"
)
spec["type"] = setting_type
spec.setdefault("label", key.replace("_", " ").title())
spec.setdefault("group", "Basic")
if setting_type == "choice":
options = spec.get("options", [])
if not isinstance(options, (list, tuple)) or not options:
raise ModelPluginError(f"{label} setting {key} needs options")
spec["options"] = list(options)
spec.setdefault("default", spec["options"][0])
clean[str(key)] = spec
return clean
def get(self, plugin_id: str) -> ModelPlugin:
return self.plugins[plugin_id]
def all(self) -> list[ModelPlugin]:
return list(self.plugins.values())
def by_trainer(self, trainer: str) -> ModelPlugin | None:
return next((plugin for plugin in self.plugins.values() if plugin.id == trainer), None)
def training_schema(self, trainer: str) -> dict[str, dict[str, Any]]:
plugin = self.by_trainer(trainer)
return plugin.training_settings if plugin else {}
def generation_schema_for_tool(self, tool_id: str) -> dict[str, dict[str, Any]]:
for plugin in self.plugins.values():
if plugin.generator_id == tool_id:
return plugin.generation_settings
return {}
def training_tool_specs(self) -> list[dict[str, Any]]:
return [
self._tool_spec(plugin, mode="training")
for plugin in self.plugins.values()
if plugin.training_tool
]
def generation_tool_specs(self) -> list[dict[str, Any]]:
return [
self._tool_spec(plugin, mode="generation")
for plugin in self.plugins.values()
if plugin.generation_tool
]
@staticmethod
def _tool_spec(plugin: ModelPlugin, *, mode: str) -> dict[str, Any]:
tool = dict(plugin.training_tool if mode == "training" else plugin.generation_tool)
schema = plugin.training_settings if mode == "training" else plugin.generation_settings
core_arguments = (
["dataset_dir", "model_name", "epochs", "output_dir", "resume_from"]
if mode == "training"
else [
"model_name", "model_path", "prompt", "image_count", "steps",
"seed", "sampler", "aspect_ratio",
]
)
core_required = (
["dataset_dir", "model_name", "epochs", "output_dir"]
if mode == "training"
else ["model_name", "model_path", "image_count", "steps", "seed"]
)
defaults = {
"id": plugin.trainer_id if mode == "training" else plugin.generator_id,
"name": f"{plugin.name} {'Trainer' if mode == 'training' else 'Generator'}",
"description": plugin.info.get("description", ""),
"category": "Training" if mode == "training" else "Output",
"entry_function": "train" if mode == "training" else "generate",
"arguments": [*core_arguments, *list(schema)],
"required_arguments": [
*core_required,
*[key for key, spec in schema.items() if bool(spec.get("required"))],
],
"capabilities": (
["fresh_training", "progress", "pause", "cancel"]
if mode == "training"
else ["image_generation", "progress", "cancel"]
),
"requires_confirmation": mode == "training",
"enabled": True,
"demo": False,
}
defaults.update(tool)
defaults["arguments"] = list(defaults.get("arguments") or [*core_arguments, *list(schema)])
defaults["required_arguments"] = list(defaults.get("required_arguments") or [])
return defaults
def validate_settings(
self,
trainer: str,
values: dict[str, Any],
*,
mode: str = "training",
) -> list[str]:
plugin = self.by_trainer(trainer)
if not plugin:
return [f"Unknown model plugin: {trainer}"]
schema = plugin.training_settings if mode == "training" else plugin.generation_settings
return validate_settings(schema, values)
def validate_settings(schema: dict[str, dict[str, Any]], values: dict[str, Any]) -> list[str]:
errors: list[str] = []
for key, spec in schema.items():
value = values.get(key, spec.get("default"))
label = str(spec.get("label", key))
if spec.get("required") and (value is None or str(value).strip() == ""):
errors.append(f"{label} is required.")
continue
if value in (None, "") and not spec.get("required"):
continue
setting_type = str(spec.get("type", "text"))
try:
if setting_type in {"int", "slider"}:
if isinstance(value, bool):
raise ValueError
numeric = int(value)
elif setting_type == "float":
if isinstance(value, bool):
raise ValueError
numeric = float(value)
else:
numeric = None
except (TypeError, ValueError):
errors.append(f"{label} must be a number.")
continue
if numeric is not None:
if "min" in spec and numeric < float(spec["min"]):
errors.append(f"{label} must be at least {spec['min']}.")
if "max" in spec and numeric > float(spec["max"]):
errors.append(f"{label} must be at most {spec['max']}.")
if setting_type == "choice" and "options" in spec and value not in spec["options"]:
errors.append(f"{label} must be one of: {', '.join(map(str, spec['options']))}.")
if setting_type == "path" and spec.get("must_exist") and not Path(str(value)).expanduser().is_file():
errors.append(f"{label} must point to an existing file.")
if setting_type == "folder" and spec.get("must_exist") and not Path(str(value)).expanduser().is_dir():
errors.append(f"{label} must point to an existing folder.")
return errors
def load_presets(root: Path, plugin_id: str, mode: str) -> dict[str, dict[str, Any]]:
path = root.resolve() / "config" / "model_presets.json"
try:
payload = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
return {}
presets = payload.get(plugin_id, {}).get(mode, {})
return dict(presets) if isinstance(presets, dict) else {}
def save_preset(
root: Path,
plugin_id: str,
mode: str,
name: str,
settings: dict[str, Any],
) -> None:
path = root.resolve() / "config" / "model_presets.json"
try:
payload = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
payload = {}
payload.setdefault(plugin_id, {}).setdefault(mode, {})[name] = settings
temporary = path.with_suffix(".tmp")
temporary.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8")
temporary.replace(path)
def plugin_function(plugin: ModelPlugin, function_name: str) -> Callable[..., Any] | None:
if plugin.module_name.endswith("manifest.py"):
spec = importlib.util.spec_from_file_location(
f"adam_user_model_{plugin.id}_{abs(hash(plugin.module_name))}",
plugin.module_name,
)
if spec is None or spec.loader is None:
return None
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
else:
module = importlib.import_module(plugin.module_name)
function = getattr(module, function_name, None)
return function if callable(function) else None
def safe_plugin_id(name: str) -> str:
cleaned = "".join(
character.lower() if character.isalnum() else "_"
for character in name.strip()
)
cleaned = "_".join(part for part in cleaned.split("_") if part)
return cleaned[:48] or "my_model"
def scaffold_model_plugin(
root: Path,
*,
plugin_id: str,
name: str,
architecture: str = "custom",
output_type: str = "image",
include_training: bool = True,
include_generation: bool = True,
) -> Path:
"""Create a simple user-editable model plugin folder."""
plugin_id = safe_plugin_id(plugin_id)
if plugin_id in {"ddpm", "flow", "lora", "model_template"}:
raise ModelPluginError("Choose a plugin id that does not conflict with a built-in model.")
folder = root.resolve() / "models" / plugin_id
if folder.exists():
raise ModelPluginError(f"A model plugin folder already exists: {folder}")
folder.mkdir(parents=True)
(folder / "__init__.py").write_text(
f'"""ADAM model plugin: {name}."""\n',
encoding="utf-8",
)
(folder / "manifest.py").write_text(
_manifest_template(
plugin_id=plugin_id,
name=name,
architecture=architecture,
output_type=output_type,
include_training=include_training,
include_generation=include_generation,
),
encoding="utf-8",
)
(folder / "model.py").write_text(_model_template(), encoding="utf-8")
if include_training:
(folder / "trainer.py").write_text(_trainer_template(), encoding="utf-8")
if include_generation:
(folder / "generator.py").write_text(_generator_template(), encoding="utf-8")
return folder
def _manifest_template(
*,
plugin_id: str,
name: str,
architecture: str,
output_type: str,
include_training: bool,
include_generation: bool,
) -> str:
plugin_id_json = json.dumps(plugin_id)
name_json = json.dumps(name)
architecture_json = json.dumps(architecture)
output_type_json = json.dumps(output_type)
training_tool = (
"{\n"
f' "id": "{plugin_id}_trainer",\n'
f' "name": {json.dumps(name + " Trainer")},\n'
f' "backend": {{"type": "python", "module": "models.{plugin_id}.trainer", "function": "train"}},\n'
"}"
if include_training else "{}"
)
generation_tool = (
"{\n"
f' "id": "{plugin_id}_generator",\n'
f' "name": {json.dumps(name + " Generator")},\n'
f' "model_trainers": ["{plugin_id}"],\n'
f' "backend": {{"type": "python", "module": "models.{plugin_id}.generator", "function": "generate"}},\n'
"}"
if include_generation else "{}"
)
return f'''PLUGIN_ID = {plugin_id_json}
MODEL_INFO = {{
"name": {name_json},
"version": "0.1",
"category": "Image Generation",
"description": {json.dumps("Describe what " + name + " trains or generates.")},
"architecture": {architecture_json},
"status": "experimental",
"output_type": {output_type_json},
}}
TRAINING_SETTINGS = {{
"resolution": {{"label": "Resolution", "type": "choice", "options": [64, 128, 256, 384, 512], "default": 256, "group": "Basic"}},
"batch_size": {{"label": "Batch size", "type": "int", "default": 1, "min": 1, "max": 64, "group": "Basic"}},
"learning_rate": {{"label": "Learning rate", "type": "float", "default": 0.0001, "min": 0.0000001, "max": 0.1, "decimals": 7, "group": "Optimization"}},
"mixed_precision": {{"label": "Precision", "type": "choice", "options": ["fp16", "no"], "default": "fp16", "group": "Optimization"}},
"preview_enabled": {{"label": "Generate previews while training", "type": "bool", "default": True, "group": "Preview"}},
"preview_every": {{"label": "Preview interval", "type": "int", "default": 5, "min": 1, "max": 100000, "group": "Preview"}},
"preview_prompt": {{"label": "Preview prompt", "type": "text", "default": "", "group": "Preview"}},
"preview_seed": {{"label": "Preview seed", "type": "int", "default": 123456789, "min": 0, "max": 2147483647, "group": "Preview"}},
}}
GENERATION_SETTINGS = {{
"prompt": {{"label": "Prompt", "type": "multiline_text", "default": "", "group": "Prompt"}},
"image_count": {{"label": "Images", "type": "int", "default": 1, "min": 1, "max": 48, "group": "Generation"}},
"steps": {{"label": "Steps", "type": "int", "default": 30, "min": 1, "max": 500, "group": "Generation"}},
"seed": {{"label": "Seed", "type": "int", "default": 0, "min": 0, "max": 2147483647, "group": "Generation"}},
}}
TRAINING_TOOL = {training_tool}
GENERATION_TOOL = {generation_tool}
'''
def _model_template() -> str:
return '''from __future__ import annotations
from typing import Any
def load_model(model_path: str, settings: dict[str, Any] | None = None) -> Any:
"""Load your model or inference pipeline here."""
raise NotImplementedError("Add your model loading code.")
'''
def _trainer_template() -> str:
return '''from __future__ import annotations
from typing import Any
def train(context, **settings: Any) -> dict[str, Any]:
"""Train the model and report progress back to ADAM."""
context.log("Replace this with real training code.")
context.progress(100, "Training placeholder complete")
return {}
'''
def _generator_template() -> str:
return '''from __future__ import annotations
from typing import Any
def generate(context, **settings: Any) -> dict[str, Any]:
"""Generate outputs and report progress back to ADAM."""
context.log("Replace this with real generation code.")
context.progress(100, "Generation placeholder complete")
return {}
'''