Download adam/model_plugins.py from SyntheticMDProductions/AI_Development_Automation_Manager: direct link, hf CLI and curl.
- Browser
- Download file 19.2 kB
-
https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/refs%2Fpr%2F1/adam/model_plugins.py
- Command line
-
hf download hf://SyntheticMDProductions/AI_Development_Automation_Manager@refs/pr/1/adam/model_plugins.py
-
curl -L -o model_plugins.py https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/refs%2Fpr%2F1/adam/model_plugins.py
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 | |
| 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 | |
| def name(self) -> str: | |
| return str(self.info.get("name", self.id)) | |
| def trainer_id(self) -> str: | |
| return str(self.training_tool.get("id") or f"{self.id}_trainer") | |
| 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, | |
| ) | |
| 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 | |
| ] | |
| 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 {} | |
| ''' | |