Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 26 additions & 2 deletions predicators/agent_sdk/tools/synthesis.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import dataclasses
import hashlib
import os
import time
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, Union

import numpy as np
Expand Down Expand Up @@ -336,6 +337,7 @@ def create_synthesis_tools(

from predicators.code_sim_learning.base_simulator import \
base_simulator_class
from predicators.code_sim_learning.config import SysIdConfig
from predicators.code_sim_learning.evidence import format_evidence_lines
from predicators.code_sim_learning.fit_space import ParamSpec
from predicators.code_sim_learning.fitting import compute_sse, \
Expand Down Expand Up @@ -527,8 +529,16 @@ def _evaluate_rollout_fit(rules: list,
for s in list(physical_specs) + list(rule_specs)
}
anchors = approach.fit_prior_anchors(physical_specs)
# pylint: disable-next=import-outside-toplevel
from predicators.agent_sdk.belief_probe import ProbeBudgetExceeded
fit_timeout = float(CFG.agent_sdk_fit_call_timeout)
# The fit's own search yields to the budget the watchdog enforces.
fit_config = dataclasses.replace(
SysIdConfig.from_cfg(),
deadline=(time.monotonic() +
fit_timeout if fit_timeout > 0 else None))
try:
with suspend_budget_watchdog(CFG.agent_sdk_fit_call_timeout):
with suspend_budget_watchdog(fit_timeout):
outcome = run_rollout_sysid(
fit_env,
rollouts,
Expand All @@ -548,7 +558,19 @@ def _evaluate_rollout_fit(rules: list,
approach, "_explainability_cache", None)),
fit_cache=(None if exploratory else getattr(
approach, "_sysid_fit_cache", None)),
fit_cache_key=version_tag)
fit_cache_key=version_tag,
config=fit_config)
except ProbeBudgetExceeded:
return (
f"[{version_tag}] Error: the rollout system-ID fit ran past "
f"its {fit_timeout:.0f} s limit and was stopped; nothing was "
"fitted or applied, and the same fit will stop again. A fit "
"replays every scored recording once per candidate parameter "
"point, so its time grows with the declared parameters and "
"with the steps each replay runs. Recordings are cut where "
"every scored feature rests, so a scored feature that keeps "
"changing (a robot that is always moving, a running counter) "
"keeps a recording whole.")
except Exception as e: # pylint: disable=broad-except
return (
f"[{version_tag}] Error: rollout system-ID fit failed:\n{e}")
Expand Down Expand Up @@ -725,6 +747,8 @@ def _evaluate_rollout_fit(rules: list,

if outcome.belief is not None:
lines.extend(["", *_format_parameter_belief(outcome.belief), ""])
if outcome.fit_result.search_note:
lines.extend([outcome.fit_result.search_note, ""])
if exploratory:
lines.append(
"EXPLORATORY subset fit: nothing was applied or recorded "
Expand Down
4 changes: 4 additions & 0 deletions predicators/code_sim_learning/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,10 @@ class SysIdConfig:
# code_sim_learning_rollout_global_design_points / _global_starts).
global_design_points: int = 0
global_starts: int = 0
# The time.monotonic() time by which the fit call must have returned
# (the sim.fit tool's budget), or None for no limit. Not a flag: the
# caller that owns the budget sets it (run_rollout_sysid).
deadline: Optional[float] = None

@property
def flat_band_frac(self) -> float:
Expand Down
4 changes: 4 additions & 0 deletions predicators/code_sim_learning/fit_space.py
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,10 @@ class FitResult:
# error that persists through the recordings. None = no multi-start
# ran.
misfit: Optional[float] = None
# How the whole-box search stopped short of its full design and starts
# to keep the fit call inside its time budget, for the report; empty
# when it ran in full or did not run.
search_note: str = ""

@property
def point_estimate(self) -> Dict[str, float]:
Expand Down
96 changes: 88 additions & 8 deletions predicators/code_sim_learning/physical_sysid.py
Original file line number Diff line number Diff line change
Expand Up @@ -293,6 +293,7 @@ def fit_params_rollout(
# info-seeking explorer's calibrated ensemble.
lm_notes: List[str] = []
lm_flat: List[str] = []
lm_t0 = time.monotonic()
lm_theta, lm_jac = fit_map_lm_rollout(base_env,
trajectories,
lm_physical_specs,
Expand All @@ -306,6 +307,7 @@ def fit_params_rollout(
noise_sigma=noise_sigma,
notes_out=lm_notes,
flat_params_out=lm_flat)
lm_seconds = time.monotonic() - lm_t0
if (config.log_hessian_identifiability and lm_jac is not None
and lm_jac.size > 0):
log_hessian_identifiability(lm_jac, names, noise_sigma, prior_sigma)
Expand Down Expand Up @@ -335,13 +337,15 @@ def fit_params_rollout(
}
basins: Optional[List[Tuple[Dict[str, float], float, float]]] = None
misfit: Optional[float] = None
search_note = ""
if (config.anchors_are_guesses and config.global_design_points > 0
and config.global_starts > 0 and trajectories):
lm_theta, lm_jac, basins, misfit = _global_multistart(
(lm_theta, lm_jac, basins, misfit, search_note) = _global_multistart(
base_env, trajectories, all_specs, physical_specs, rule_specs,
residual_features, rules, latent_init, scaling, center_int,
prior_sigma, noise_sigma, config, (lm_theta, lm_jac), noise_sse,
sigma_tol)
sigma_tol,
_SearchClock(config.deadline, fit_t0, n_start, lm_seconds))
elif config.anchors_are_guesses and trajectories and noise_sse > 0.0:
# The same misfit the multi-start measures, at the local fit.
map_sse = compute_rollout_sse(
Expand All @@ -359,7 +363,8 @@ def fit_params_rollout(
sensitivity=sensitivity,
lm_notes=lm_notes,
basins=basins,
misfit=misfit)
misfit=misfit,
search_note=search_note)
n_lm = num_rollouts_run() - n_start - n_grid
# The ablation pins moved parameters back to their anchors, which
# only calibrated baselines warrant; a guess has no claim beyond the
Expand Down Expand Up @@ -398,6 +403,42 @@ def fit_params_rollout(
# fit-space box width are one basin.
_SAME_BASIN_FRACTION = 0.05

# The share of the time a fit call has left when its fit starts that the
# whole-box search may use; the rest is the report's and the belief's,
# which trace every basin.
_SEARCH_BUDGET_SHARE = 0.5


@dataclasses.dataclass(frozen=True)
class _SearchClock:
"""The whole-box search's time budget: ``deadline`` (monotonic, or None for
no limit) is when the fit call must have returned, ``fit_t0`` and
``rollouts_before`` when and after how many rollouts the fit started, and
``lm_seconds`` how long the local LM took."""

deadline: Optional[float]
fit_t0: float
rollouts_before: int
lm_seconds: float

@property
def budget(self) -> float:
"""The seconds the fit had when it started (inf without a limit)."""
if self.deadline is None:
return float("inf")
return max(0.0, self.deadline - self.fit_t0)

def allows(self, seconds: float) -> bool:
"""Whether a phase projected to take ``seconds`` ends inside the
search's share of the budget."""
return (time.monotonic() + seconds <=
self.fit_t0 + _SEARCH_BUDGET_SHARE * self.budget)

def seconds_per_rollout(self) -> float:
"""The fit's wall time per rollout so far, waves included."""
done = max(1, num_rollouts_run() - self.rollouts_before)
return (time.monotonic() - self.fit_t0) / done


def _global_multistart(
base_env: Any,
Expand All @@ -416,8 +457,9 @@ def _global_multistart(
incumbent: Tuple[np.ndarray, Optional[np.ndarray]],
noise_sse: float,
sigma_tol: float,
) -> Tuple[np.ndarray, Optional[np.ndarray], List[Tuple[Dict[
str, float], float, float]], Optional[float]]:
clock: _SearchClock,
) -> Tuple[np.ndarray, Optional[np.ndarray], Optional[List[Tuple[Dict[
str, float], float, float]]], Optional[float], str]:
"""Search the whole box for the fit when the starting values are guesses.

A scrambled Sobol design of ``config.global_design_points`` over
Expand All @@ -438,9 +480,36 @@ def _global_multistart(
is returned too: the belief's temperature counts it the same way.
Without a declared noise channel (``noise_sse`` 0) the misfit is
unknown: the point estimate is the only basin and the misfit None.

The search improves on a fit that already exists, so it never makes
the fit call overrun its budget (``clock``): the design and each LM
start run only if they are projected to end within the search's
share of it, timed by the fit so far. A search cut short says so in
the returned note (empty when it ran in full); one that cannot
afford the design and a start keeps the incumbent, whose misfit is
still measured, and returns no basins.
"""
names = [s.name for s in all_specs]
physical_names = [s.name for s in physical_specs]
design_seconds = (clock.seconds_per_rollout() *
config.global_design_points * len(trajectories))
if not clock.allows(design_seconds + clock.lm_seconds):
point = dict(zip(names, (float(v) for v in incumbent[0])))
sse = rollout_sse_by_point(base_env, trajectories, [point],
residual_features, physical_names, rules,
latent_init, scaling)[0]
incumbent_misfit = (max(sse - noise_sse, sigma_tol)
if noise_sse > 0.0 else None)
note = (f"The whole-box search did not run: the fit had used "
f"{time.monotonic() - clock.fit_t0:.0f} s of its "
f"{clock.budget:.0f} s budget, and the design and one LM "
f"start (about {design_seconds + clock.lm_seconds:.0f} s) "
"would not have ended within half of it. The point estimate "
"is the fit from the starting values; other parameter sets "
"that fit the recordings about as well were not looked for.")
logger.info("Rollout sysID global multi-start skipped: %s", note)
return (np.asarray(incumbent[0], dtype=float), incumbent[1], None,
incumbent_misfit, note)
lo, hi = fit_space_bounds(list(all_specs))
lo = np.where(np.isfinite(lo), lo, center_int - 2.0 * prior_sigma)
hi = np.where(np.isfinite(hi), hi, center_int + 2.0 * prior_sigma)
Expand Down Expand Up @@ -471,9 +540,19 @@ def started_at(specs: Sequence[ParamSpec],

thetas = [np.asarray(incumbent[0], dtype=float)]
jacobians = [incumbent[1]]
for index in np.argsort(np.asarray(sses),
kind="stable")[:config.global_starts]:
starts = np.argsort(np.asarray(sses), kind="stable")[:config.global_starts]
lm_seconds = clock.lm_seconds
note = ""
for ran, index in enumerate(starts):
if not clock.allows(lm_seconds):
note = (f"The whole-box search ran the LM from {ran} of its "
f"{len(starts)} best design points: another start (about "
f"{lm_seconds:.0f} s) would not have ended within half "
f"of the fit's {clock.budget:.0f} s budget.")
logger.info("Rollout sysID global multi-start cut short: %s", note)
break
start = design[int(index)]
lm_t0 = time.monotonic()
theta, jac = fit_map_lm_rollout(base_env,
trajectories,
started_at(physical_specs, start),
Expand All @@ -485,6 +564,7 @@ def started_at(specs: Sequence[ParamSpec],
prior_centers=center_int,
prior_sigmas=prior_sigma,
noise_sigma=noise_sigma)
lm_seconds = max(lm_seconds, time.monotonic() - lm_t0)
thetas.append(np.asarray(theta, dtype=float))
jacobians.append(jac)
points = [dict(zip(names, (float(v) for v in theta))) for theta in thetas]
Expand Down Expand Up @@ -525,7 +605,7 @@ def started_at(specs: Sequence[ParamSpec],
float(np.min(sses)),
len(thetas) - 1, finals[0], [round(float(s), 4) for s in finals[1:]],
len(basins), band)
return thetas[best], jacobians[best], basins, misfit
return thetas[best], jacobians[best], basins, misfit, note


# A MAP within this fit-space distance of its anchor counts as unmoved
Expand Down
53 changes: 53 additions & 0 deletions tests/approaches/test_agent_continual_approach.py
Original file line number Diff line number Diff line change
Expand Up @@ -298,6 +298,59 @@ def fake_query(*_args: Any, **_kwargs: Any) -> List[Dict[str, Any]]:
assert card.levels[0].sandbox.get("fits", 0) == int(fit_then_edit)


@pytest.mark.slow
def test_fit_that_overruns_its_budget_says_so(tmp_path: Any,
monkeypatch: Any) -> None:
"""A fit stopped by its time limit tells the agent what happened and what
makes fits long, where the watchdog's bare exception used to print an empty
"fit failed:"."""
# pylint: disable=import-outside-toplevel,protected-access
from predicators.agent_sdk.belief_probe import ProbeBudgetExceeded
from predicators.code_sim_learning import orchestrator

seen: Dict[str, Any] = {}

def overrun(*_args: Any, **kwargs: Any) -> Any:
seen["deadline"] = kwargs["config"].deadline
raise ProbeBudgetExceeded

# The fit tool binds run_rollout_sysid when the approach builds it.
monkeypatch.setattr(orchestrator, "run_rollout_sysid", overrun)
_config(tmp_path, agent_sdk_fit_call_timeout=1200.0)
env, approach = _make_approach()
source = '''
class Counter(BaseSimulator):
AGENT_PARAM_SPECS = [ParamSpec("rate", .25, lo=0.0, hi=1.0)]
MODEL_STATE_INIT = {"charge": 0.0}
RESIDUAL_FEATURES = {}

@classmethod
def update_model_state(cls, observation, model_state, params, action):
model_state["charge"] += params["rate"]

RESIDUAL_ENV = Counter
'''

def fake_query(*_args: Any, **_kwargs: Any) -> List[Dict[str, Any]]:
assert "step applied" in _call(approach,
"env_step",
action=[0.0] *
env.action_space.shape[0])
path = os.path.join(approach._tool_context.sandbox_dir, "simulator.py")
with open(path, "w", encoding="utf-8") as file:
file.write(source)
seen["output"] = _call(approach, "run_python", code="print(sim.fit())")
assert "Give-up recorded" in _call(approach, "give_up", note="done")
return _result()

monkeypatch.setattr(approach, "_query_agent_sync", fake_query)
approach.prepare_for_continual(Dataset([]))
ContinualRun(env, approach, create_level_player(env, approach)).run()
assert "ran past its 1200 s limit" in seen["output"]
assert "nothing was fitted or applied" in seen["output"]
assert seen["deadline"] is not None


@pytest.mark.slow
def test_play_loop_with_a_scripted_agent(tmp_path: Any) -> None:
"""Two rounds of one conversation: act in the first, which also carries the
Expand Down
Loading
Loading