synapse-x / agents /baseline.py
Nithin1026's picture
Initial submission Synapse-X
cb330aa
Raw
History Blame Contribute Delete
8.38 kB
"""
SYNAPSE-X baseline agent.
SYNAPSE-CORE-X is a deterministic decision layer that combines pressure-aware
reasoning, uncertainty gating, lookahead simulation, and cascade-aware control.
"""
from statistics import mean
from typing import Protocol
from env.models import Action, ActionPayload, Observation, StepResult, Task
class SupportsEpisodeEnv(Protocol):
def reset(self) -> Observation:
...
def step(self, action: Action) -> StepResult:
...
class SynapseAgent:
CRISIS_THRESHOLD = 0.7
BALANCED_THRESHOLD = 0.4
HIGH_UNCERTAINTY_THRESHOLD = 0.6
HIGH_CASCADE_RISK = 0.8
MEDIUM_RISK_AVERSION = 0.8
def compute_pressure(self, tasks: list[Task]) -> float:
if not tasks:
return 0.0
return mean(task.risk for task in tasks)
def score_task(self, task: Task) -> float:
return task.priority - task.risk + (1.0 - task.uncertainty)
def _is_cascade_mode(self, obs: Observation) -> bool:
return any(task.dependencies for task in obs.tasks)
def simulate_future_cost(self, task: Task) -> float:
return task.risk * task.uncertainty
def cascade_penalty(self, task: Task, active_tasks: list[Task]) -> float:
risk_spike = 1.0 if task.risk > self.HIGH_CASCADE_RISK else 0.0
future_spike = 0.6 if task.future_risk > self.HIGH_CASCADE_RISK else 0.0
dependency_bonus = 0.2 * sum(1 for candidate in active_tasks if task.id in candidate.dependencies)
return risk_spike + future_spike - dependency_bonus
def _active_tasks(self, obs: Observation) -> list[Task]:
return [task for task in obs.tasks if task.is_active]
def _deps_satisfied(self, task: Task, all_tasks: list[Task]) -> bool:
task_map = {item.id: item for item in all_tasks}
return all(task_map.get(dep_id) is not None and task_map[dep_id].completed for dep_id in task.dependencies)
def _mode(self, pressure: float) -> str:
if pressure > self.CRISIS_THRESHOLD:
return "crisis"
if pressure > self.BALANCED_THRESHOLD:
return "balanced"
return "normal"
def _task_score(self, task: Task, active_tasks: list[Task], mode: str) -> float:
base_score = self.score_task(task)
future_cost = self.simulate_future_cost(task)
score = base_score - future_cost
score -= self.cascade_penalty(task, active_tasks)
score += 1.2 * task.priority
score -= 1.1 * task.future_risk
score -= 0.9 * task.deadline_pressure
score += 0.7 if task.deadline <= 3.0 else 0.0
score -= self.MEDIUM_RISK_AVERSION * task.risk * task.uncertainty
if mode == "crisis":
score += 0.8 * (1.0 - task.risk) + 0.5 * (1.0 - task.uncertainty)
elif mode == "balanced":
score += 0.4 * task.priority + 0.3 * (1.0 - future_cost)
else:
score += 0.7 * task.priority - 0.2 * task.risk
return score
def _blocked_parent(self, active_tasks: list[Task], all_tasks: list[Task]) -> Task:
task_map = {task.id: task for task in all_tasks}
blocked = min(
active_tasks,
key=lambda task: (
sum(1 for dep_id in task.dependencies if not task_map.get(dep_id, task).completed),
task.deadline,
),
)
remaining_parents = [
task_map[dep_id]
for dep_id in blocked.dependencies
if dep_id in task_map and not task_map[dep_id].completed and not task_map[dep_id].failed
]
if remaining_parents:
return min(remaining_parents, key=lambda task: (task.risk + task.uncertainty, task.deadline))
return blocked
def _should_stabilize(self, obs: Observation, task: Task) -> bool:
if task.deadline <= 4.0:
return False
if obs.resources < task.resources_required:
return True
return (
task.risk + task.future_risk > 1.45
and task.uncertainty > self.HIGH_UNCERTAINTY_THRESHOLD
and obs.resources < 0.9
)
def _select_feasible_task(self, obs: Observation, feasible: list[Task], cascade_mode: bool) -> Task:
if not cascade_mode and obs.time > 0:
return min(
feasible,
key=lambda task: (
task.risk + 0.5 * task.uncertainty + 0.3 * task.future_risk,
-task.priority,
task.deadline,
),
)
return feasible[0]
def act(self, state: Observation | dict) -> ActionPayload:
obs = state if isinstance(state, Observation) else Observation(**state)
active_tasks = self._active_tasks(obs)
if not active_tasks:
return {"action_type": "delay", "task_id": 0}
pressure = self.compute_pressure(active_tasks)
mode = self._mode(pressure)
cascade_mode = self._is_cascade_mode(obs)
ranked = sorted(
active_tasks,
key=lambda task: self._task_score(task, active_tasks, mode),
reverse=True,
)
feasible = [
task
for task in ranked
if obs.resources >= task.resources_required and self._deps_satisfied(task, obs.tasks)
]
ready = [task for task in ranked if self._deps_satisfied(task, obs.tasks)]
if feasible:
best_task = self._select_feasible_task(obs, feasible, cascade_mode)
deadline_urgent = best_task.deadline <= 4.0
if not deadline_urgent:
if cascade_mode and self._should_stabilize(obs, best_task):
return {"action_type": "reallocate", "task_id": best_task.id}
if (
best_task.uncertainty > self.HIGH_UNCERTAINTY_THRESHOLD
and obs.resources < best_task.resources_required + 0.15
and best_task.deadline > 5.0
):
return {"action_type": "reallocate", "task_id": best_task.id}
if mode == "crisis" and best_task.risk > 0.7 and best_task.deadline > 4.0:
return {"action_type": "delay", "task_id": best_task.id}
needs_buffer = obs.resources < min(1.0, best_task.resources_required + 0.05)
if needs_buffer and len(feasible) == 1 and best_task.deadline > 4.0 and obs.resources < 1.0:
return {"action_type": "reallocate", "task_id": best_task.id}
return {"action_type": "execute", "task_id": best_task.id}
if ready:
best_task = ready[0]
if obs.resources < 1.0 and not best_task.is_terminal:
return {"action_type": "reallocate", "task_id": best_task.id}
delayable = next((task for task in ready if not task.is_terminal), None)
if delayable is not None:
return {"action_type": "delay", "task_id": delayable.id}
return {"action_type": "reallocate", "task_id": best_task.id}
parent = self._blocked_parent(active_tasks, obs.tasks)
if not parent.is_terminal and obs.resources >= parent.resources_required and self._deps_satisfied(parent, obs.tasks):
return {"action_type": "execute", "task_id": parent.id}
if not parent.is_terminal and obs.resources < 1.0:
return {"action_type": "reallocate", "task_id": parent.id}
fallback = next((task for task in active_tasks if not task.is_terminal), None)
if fallback is not None:
return {"action_type": "delay", "task_id": fallback.id}
return {"action_type": "reallocate", "task_id": 0}
def select_action(obs: Observation) -> Action:
agent = SynapseAgent()
return Action(**agent.act(obs))
def run_episode(env: SupportsEpisodeEnv, max_steps: int = 30) -> list[ActionPayload]:
obs = env.reset()
actions_taken: list[ActionPayload] = []
for _ in range(max_steps):
if obs.episode_done:
break
action = select_action(obs)
actions_taken.append(action.model_dump())
result = env.step(action)
obs = result.observation
return actions_taken
def collect_actions_for_grader(_task_name: str, _obs: Observation, env: SupportsEpisodeEnv) -> list[ActionPayload]:
return run_episode(env)