| from __future__ import annotations |
|
|
| from typing import Any |
|
|
| from adam.executor import ToolContext, ToolExecutionError |
| from adam.model_plugins import ModelPluginRegistry, plugin_function, validate_settings |
|
|
|
|
| def _plugin_for_tool(context: ToolContext, mode: str): |
| registry = ModelPluginRegistry(context.root) |
| for plugin in registry.all(): |
| if mode == "training" and plugin.trainer_id == context.tool.id: |
| return plugin |
| if mode == "generation" and plugin.generator_id == context.tool.id: |
| return plugin |
| raise ToolExecutionError(f"No model plugin owns {context.tool.id}.") |
|
|
|
|
| def train(context: ToolContext, **settings: Any) -> dict[str, Any]: |
| plugin = _plugin_for_tool(context, "training") |
| errors = validate_settings(plugin.training_settings, settings) |
| if errors: |
| raise ToolExecutionError(" ".join(errors)) |
| function = plugin_function(plugin, "train") |
| if function is None: |
| raise ToolExecutionError(f"{plugin.name} does not implement train().") |
| return function(settings=settings, callbacks=context) or {} |
|
|
|
|
| def generate(context: ToolContext, **settings: Any) -> dict[str, Any]: |
| plugin = _plugin_for_tool(context, "generation") |
| errors = validate_settings(plugin.generation_settings, settings) |
| if errors: |
| raise ToolExecutionError(" ".join(errors)) |
| function = plugin_function(plugin, "generate") |
| if function is None: |
| raise ToolExecutionError(f"{plugin.name} does not implement generate().") |
| return function(settings=settings, callbacks=context) or {} |
|
|