Spaces:
Sleeping
Sleeping
| 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) | |