feat: add skill learning runtime
This commit is contained in:
+68
-2
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, replace
|
||||
from inspect import Parameter, signature
|
||||
@@ -9,6 +10,15 @@ from runtime.context import TaskContext
|
||||
from runtime.executor import Executor
|
||||
from runtime.planner import PlannedStep, Planner
|
||||
from semantic.models import SemanticScene
|
||||
from skills_learning.config import (
|
||||
SkillAuthoringConfig,
|
||||
load_config as load_skill_authoring_config,
|
||||
)
|
||||
from skills_learning.embeddings import EmbeddingClient, embed_skill_text
|
||||
from skills_learning.models import skill_embedding_text
|
||||
from skills_learning.store import SkillStore, get_default_store
|
||||
from skills_learning.synthesis import synthesize_flow_skill
|
||||
from skills_learning.versioning import store_synthesized_skill
|
||||
from storage.task_metadata import TaskMetadataStore
|
||||
from storage.timeline import Timeline
|
||||
from tools.describe_screen import describe_screen
|
||||
@@ -16,6 +26,8 @@ from tools.screenshot import take_screenshot
|
||||
from world.config import WorldConfig, load_config as load_world_config
|
||||
from world.model import WorldModel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TaskRunnerConfig:
|
||||
@@ -24,6 +36,7 @@ class TaskRunnerConfig:
|
||||
|
||||
Observer = Callable[[str], Scene]
|
||||
ScreenshotProvider = Callable[[str], bytes]
|
||||
TaskSucceededHook = Callable[[str, str, Timeline], None]
|
||||
|
||||
|
||||
class TaskRunner:
|
||||
@@ -39,6 +52,10 @@ class TaskRunner:
|
||||
screenshot_provider: ScreenshotProvider | None = None,
|
||||
world_model: WorldModel | None = None,
|
||||
world_config: WorldConfig | None = None,
|
||||
on_task_succeeded: TaskSucceededHook | None = None,
|
||||
skill_authoring_config: SkillAuthoringConfig | None = None,
|
||||
skill_store: SkillStore | None = None,
|
||||
skill_embedding_client: EmbeddingClient | None = None,
|
||||
) -> None:
|
||||
self.planner = planner or Planner()
|
||||
self.executor = executor or Executor()
|
||||
@@ -56,6 +73,17 @@ class TaskRunner:
|
||||
self.world_model = WorldModel(config=self.world_config)
|
||||
else:
|
||||
self.world_model = None
|
||||
self.skill_authoring_config = (
|
||||
skill_authoring_config or load_skill_authoring_config()
|
||||
)
|
||||
self.skill_store = skill_store
|
||||
self.skill_embedding_client = skill_embedding_client
|
||||
if on_task_succeeded is not None:
|
||||
self.on_task_succeeded = on_task_succeeded
|
||||
elif self.skill_authoring_config.enabled:
|
||||
self.on_task_succeeded = self._default_task_succeeded_hook
|
||||
else:
|
||||
self.on_task_succeeded = None
|
||||
|
||||
def run(self, task: Task) -> Task:
|
||||
context = TaskContext(task_id=task.id, goal=task.goal)
|
||||
@@ -72,8 +100,7 @@ class TaskRunner:
|
||||
scene=scene,
|
||||
context=context,
|
||||
):
|
||||
self._update_task(task, status="completed", completed=True)
|
||||
return task
|
||||
return self._complete_task(task)
|
||||
|
||||
for step in steps:
|
||||
executable_step = self._step_for_device(step, task.device_id)
|
||||
@@ -101,6 +128,45 @@ class TaskRunner:
|
||||
)
|
||||
return task
|
||||
|
||||
def _complete_task(self, task: Task) -> Task:
|
||||
self._update_task(task, status="completed", completed=True)
|
||||
self._notify_task_succeeded(task)
|
||||
return task
|
||||
|
||||
def _notify_task_succeeded(self, task: Task) -> None:
|
||||
if self.on_task_succeeded is None or self.timeline is None:
|
||||
return
|
||||
try:
|
||||
self.on_task_succeeded(task.id, task.goal, self.timeline)
|
||||
except Exception as exc:
|
||||
logger.info("task succeeded hook failed: %s", exc)
|
||||
|
||||
def _default_task_succeeded_hook(
|
||||
self,
|
||||
task_id: str,
|
||||
goal: str,
|
||||
timeline: Timeline,
|
||||
) -> None:
|
||||
store = self.skill_store or get_default_store()
|
||||
candidate = synthesize_flow_skill(
|
||||
goal,
|
||||
timeline,
|
||||
task_id=task_id,
|
||||
store=store,
|
||||
)
|
||||
stored = store_synthesized_skill(store, candidate).skill
|
||||
vector = embed_skill_text(
|
||||
skill_embedding_text(stored),
|
||||
client=self.skill_embedding_client,
|
||||
config=self.skill_authoring_config,
|
||||
)
|
||||
if vector is not None:
|
||||
store.store_embedding(
|
||||
stored,
|
||||
vector,
|
||||
model_name=self.skill_authoring_config.embedding_model,
|
||||
)
|
||||
|
||||
def _plan(
|
||||
self,
|
||||
goal: str,
|
||||
|
||||
Reference in New Issue
Block a user