permanence-training / permanence /task_manager.py
chane35's picture
PERMANENCE: reversibility-aware RL environment for training LLM agents
796da7c verified
Raw
History Blame
1.39 kB
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)