From fd84d66eb74e4ecdc12cc3e3c3555547e3d64f08 Mon Sep 17 00:00:00 2001 From: blepping Date: Mon, 30 Dec 2024 05:09:45 -0700 Subject: [PATCH] Support kl_optimal schedule in schedule overrides when available --- py/schedule.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/py/schedule.py b/py/schedule.py index 09eaa71..b2ac945 100644 --- a/py/schedule.py +++ b/py/schedule.py @@ -28,6 +28,7 @@ class Schedule: "sgm_uniform": {}, "simple": {}, "vp": {"beta_d": 19.9, "beta_min": 0.1, "eps_s": 0.001}, + "kl_optimal": {"sigma_min": -1, "sigma_max": -1}, } def __init__( @@ -63,12 +64,11 @@ class Schedule: schedule_name: str, schedule_kwargs: dict[str, Any] | None = None, ) -> dict[str, int | float | str]: - schedule_default_kwargs = self.schedule_default_kwargs.get( - schedule_name.lower(), - ) - if schedule_default_kwargs is None: + schedule_name = schedule_name.lower().strip() + if schedule_name not in self._make_sigmas_handlers: errstr = f"Unknown schedule: {schedule_name}" raise ValueError(errstr) + schedule_default_kwargs = self.schedule_default_kwargs.get(schedule_name, {}) schedule_kwargs = fallback(schedule_kwargs, {}) bad_keys = ",".join({ k for k in schedule_kwargs if k not in schedule_default_kwargs