from __future__ import annotations from typing import Dict, Optional, Tuple from .tasks.task_bank import CurriculumScheduler, TaskBank, TaskSpec, TaskTemplate from .world.state import WorldState class TaskManager: """Mediates between the env and the task bank. Supports a ``domain`` filter so the curriculum only samples from a single domain. Changing the ``domain`` parameter switches which registered domain the curriculum samples from. """ def __init__( self, task_bank: Optional[TaskBank] = None, domain: Optional[str] = "devtools", ) -> None: self.task_bank = task_bank or TaskBank() # Replace the default scheduler with a domain-aware one. self.task_bank._scheduler = CurriculumScheduler(domain=domain) def select_template(self, episode_index: int, force_task: Optional[str] = None) -> TaskTemplate: if force_task is not None: return self.task_bank.get(force_task) return self.task_bank.get_for_episode(episode_index) def instantiate( self, episode_index: int, seed: int, force_task: Optional[str] = None, difficulty: float = 0.5, ) -> Tuple[TaskSpec, WorldState, Dict[str, float]]: template = self.select_template(episode_index, force_task) return template.instantiate(seed, difficulty=difficulty)