| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| import datetime as dt |
| from dataclasses import dataclass, field |
| from logging import getLogger |
| from pathlib import Path |
|
|
| from lerobot import envs, policies |
|
|
| from . import parser |
| from .default import EvalConfig |
| from .policies import PreTrainedConfig |
|
|
| logger = getLogger(__name__) |
|
|
|
|
| @dataclass |
| class EvalPipelineConfig: |
| |
| |
| |
| env: envs.EnvConfig |
| eval: EvalConfig = field(default_factory=EvalConfig) |
| policy: PreTrainedConfig | None = None |
| output_dir: Path | None = None |
| job_name: str | None = None |
| seed: int | None = 1000 |
| |
| rename_map: dict[str, str] = field(default_factory=dict) |
| |
| trust_remote_code: bool = False |
|
|
| def __post_init__(self) -> None: |
| |
| policy_path = parser.get_path_arg("policy") |
| if policy_path: |
| yaml_overrides = parser.get_yaml_overrides("policy") |
| cli_overrides = parser.get_cli_overrides("policy") or [] |
| self.policy = PreTrainedConfig.from_pretrained( |
| policy_path, cli_overrides=yaml_overrides + cli_overrides |
| ) |
| self.policy.pretrained_path = Path(policy_path) |
|
|
| else: |
| logger.warning( |
| "No pretrained path was provided, evaluated policy will be built from scratch (random weights)." |
| ) |
|
|
| if not self.job_name: |
| if self.env is None: |
| self.job_name = f"{self.policy.type if self.policy is not None else 'scratch'}" |
| else: |
| self.job_name = ( |
| f"{self.env.type}_{self.policy.type if self.policy is not None else 'scratch'}" |
| ) |
|
|
| logger.warning(f"No job name provided, using '{self.job_name}' as job name.") |
|
|
| if not self.output_dir: |
| now = dt.datetime.now() |
| eval_dir = f"{now:%Y-%m-%d}/{now:%H-%M-%S}_{self.job_name}" |
| self.output_dir = Path("outputs/eval") / eval_dir |
|
|
| @classmethod |
| def __get_path_fields__(cls) -> list[str]: |
| """This enables the parser to load config from the policy using `--policy.path=local/dir`""" |
| return ["policy"] |
|
|