File size: 1,571 Bytes
c61c435
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
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 {}