Spaces:
Sleeping
Sleeping
vishal harkal commited on
Commit ·
4138efa
1
Parent(s): 6cf42ae
Fix suite success/reward metrics and baseline payload
Browse files- app/routes.py +8 -3
- app/static/results.html +39 -9
- grader/grader.py +19 -0
app/routes.py
CHANGED
|
@@ -38,7 +38,9 @@ class BaselineResponse(BaseModel):
|
|
| 38 |
seed: int | None
|
| 39 |
agent: str
|
| 40 |
steps_executed: int = Field(..., ge=0)
|
|
|
|
| 41 |
done: bool
|
|
|
|
| 42 |
score: float = Field(..., ge=0.0, le=1.0)
|
| 43 |
report: GradeReport
|
| 44 |
|
|
@@ -79,7 +81,7 @@ TASKS_ACTION_SCHEMA: Dict[str, Any] = {
|
|
| 79 |
|
| 80 |
|
| 81 |
class EnvironmentService:
|
| 82 |
-
BASELINE_MAX_STEPS =
|
| 83 |
|
| 84 |
def __init__(self) -> None:
|
| 85 |
self._lock = RLock()
|
|
@@ -183,7 +185,7 @@ class EnvironmentService:
|
|
| 183 |
resolved_task = task or self._task or "easy"
|
| 184 |
config = get_task_config(task_name=resolved_task, seed=seed)
|
| 185 |
rollout_env = LastMileDeliveryEnvironment(config=config)
|
| 186 |
-
observation = rollout_env.reset()
|
| 187 |
|
| 188 |
step_cap = config.max_steps
|
| 189 |
if max_steps is not None:
|
|
@@ -206,12 +208,15 @@ class EnvironmentService:
|
|
| 206 |
rollout_env.current_observation(),
|
| 207 |
success_condition=success_condition,
|
| 208 |
)
|
|
|
|
| 209 |
return BaselineResponse(
|
| 210 |
task=resolved_task,
|
| 211 |
seed=seed,
|
| 212 |
agent="BaselineGreedyAgent",
|
| 213 |
steps_executed=len(trajectory),
|
| 214 |
-
|
|
|
|
|
|
|
| 215 |
score=report.score,
|
| 216 |
report=report,
|
| 217 |
)
|
|
|
|
| 38 |
seed: int | None
|
| 39 |
agent: str
|
| 40 |
steps_executed: int = Field(..., ge=0)
|
| 41 |
+
total_reward: float
|
| 42 |
done: bool
|
| 43 |
+
success: bool
|
| 44 |
score: float = Field(..., ge=0.0, le=1.0)
|
| 45 |
report: GradeReport
|
| 46 |
|
|
|
|
| 81 |
|
| 82 |
|
| 83 |
class EnvironmentService:
|
| 84 |
+
BASELINE_MAX_STEPS = 256
|
| 85 |
|
| 86 |
def __init__(self) -> None:
|
| 87 |
self._lock = RLock()
|
|
|
|
| 185 |
resolved_task = task or self._task or "easy"
|
| 186 |
config = get_task_config(task_name=resolved_task, seed=seed)
|
| 187 |
rollout_env = LastMileDeliveryEnvironment(config=config)
|
| 188 |
+
observation = rollout_env.reset(seed=seed)
|
| 189 |
|
| 190 |
step_cap = config.max_steps
|
| 191 |
if max_steps is not None:
|
|
|
|
| 208 |
rollout_env.current_observation(),
|
| 209 |
success_condition=success_condition,
|
| 210 |
)
|
| 211 |
+
final_state = rollout_env.state()
|
| 212 |
return BaselineResponse(
|
| 213 |
task=resolved_task,
|
| 214 |
seed=seed,
|
| 215 |
agent="BaselineGreedyAgent",
|
| 216 |
steps_executed=len(trajectory),
|
| 217 |
+
total_reward=final_state.total_reward,
|
| 218 |
+
done=final_state.done,
|
| 219 |
+
success=report.success,
|
| 220 |
score=report.score,
|
| 221 |
report=report,
|
| 222 |
)
|
app/static/results.html
CHANGED
|
@@ -691,7 +691,8 @@
|
|
| 691 |
tr.appendChild(successTd);
|
| 692 |
|
| 693 |
const scoreTd = document.createElement("td");
|
| 694 |
-
scoreTd.textContent =
|
|
|
|
| 695 |
tr.appendChild(scoreTd);
|
| 696 |
|
| 697 |
const stepsTd = document.createElement("td");
|
|
@@ -699,7 +700,10 @@
|
|
| 699 |
tr.appendChild(stepsTd);
|
| 700 |
|
| 701 |
const rewardTd = document.createElement("td");
|
| 702 |
-
rewardTd.textContent =
|
|
|
|
|
|
|
|
|
|
| 703 |
tr.appendChild(rewardTd);
|
| 704 |
|
| 705 |
tbody.appendChild(tr);
|
|
@@ -717,7 +721,17 @@
|
|
| 717 |
return;
|
| 718 |
}
|
| 719 |
|
| 720 |
-
const
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 721 |
nodes.suiteAverage.textContent = `average_score: ${average.toFixed(4)}`;
|
| 722 |
}
|
| 723 |
|
|
@@ -992,12 +1006,21 @@
|
|
| 992 |
setStatus(`Running suite ${i + 1}/${tasks.length}: ${task}`, "warn");
|
| 993 |
const query = new URLSearchParams({ task, seed: String(state.seed) });
|
| 994 |
const payload = await apiCall(`/baseline?${query.toString()}`, { method: "GET" });
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 995 |
suiteRows.push({
|
| 996 |
task,
|
| 997 |
-
success
|
| 998 |
-
score:
|
| 999 |
steps_executed: Number(payload.steps_executed || 0),
|
| 1000 |
-
total_reward:
|
|
|
|
|
|
|
|
|
|
| 1001 |
});
|
| 1002 |
}
|
| 1003 |
|
|
@@ -1007,9 +1030,16 @@
|
|
| 1007 |
return {
|
| 1008 |
seed: state.seed,
|
| 1009 |
tasks: suiteRows,
|
| 1010 |
-
average_score:
|
| 1011 |
-
|
| 1012 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1013 |
};
|
| 1014 |
});
|
| 1015 |
}
|
|
|
|
| 691 |
tr.appendChild(successTd);
|
| 692 |
|
| 693 |
const scoreTd = document.createElement("td");
|
| 694 |
+
scoreTd.textContent =
|
| 695 |
+
item.score === null || item.score === undefined ? "n/a" : Number(item.score).toFixed(4);
|
| 696 |
tr.appendChild(scoreTd);
|
| 697 |
|
| 698 |
const stepsTd = document.createElement("td");
|
|
|
|
| 700 |
tr.appendChild(stepsTd);
|
| 701 |
|
| 702 |
const rewardTd = document.createElement("td");
|
| 703 |
+
rewardTd.textContent =
|
| 704 |
+
item.total_reward === null || item.total_reward === undefined
|
| 705 |
+
? "n/a"
|
| 706 |
+
: Number(item.total_reward).toFixed(2);
|
| 707 |
tr.appendChild(rewardTd);
|
| 708 |
|
| 709 |
tbody.appendChild(tr);
|
|
|
|
| 721 |
return;
|
| 722 |
}
|
| 723 |
|
| 724 |
+
const numericScores = rows
|
| 725 |
+
.map((item) => item.score)
|
| 726 |
+
.filter((value) => value !== null && value !== undefined)
|
| 727 |
+
.map((value) => Number(value));
|
| 728 |
+
|
| 729 |
+
if (!numericScores.length) {
|
| 730 |
+
nodes.suiteAverage.textContent = "average_score: n/a";
|
| 731 |
+
return;
|
| 732 |
+
}
|
| 733 |
+
|
| 734 |
+
const average = numericScores.reduce((acc, value) => acc + value, 0) / numericScores.length;
|
| 735 |
nodes.suiteAverage.textContent = `average_score: ${average.toFixed(4)}`;
|
| 736 |
}
|
| 737 |
|
|
|
|
| 1006 |
setStatus(`Running suite ${i + 1}/${tasks.length}: ${task}`, "warn");
|
| 1007 |
const query = new URLSearchParams({ task, seed: String(state.seed) });
|
| 1008 |
const payload = await apiCall(`/baseline?${query.toString()}`, { method: "GET" });
|
| 1009 |
+
const success =
|
| 1010 |
+
payload.success !== undefined
|
| 1011 |
+
? Boolean(payload.success)
|
| 1012 |
+
: payload.report?.success !== undefined
|
| 1013 |
+
? Boolean(payload.report.success)
|
| 1014 |
+
: Boolean(payload.done);
|
| 1015 |
suiteRows.push({
|
| 1016 |
task,
|
| 1017 |
+
success,
|
| 1018 |
+
score: payload.score === undefined || payload.score === null ? null : Number(payload.score),
|
| 1019 |
steps_executed: Number(payload.steps_executed || 0),
|
| 1020 |
+
total_reward:
|
| 1021 |
+
payload.total_reward === undefined || payload.total_reward === null
|
| 1022 |
+
? null
|
| 1023 |
+
: Number(payload.total_reward),
|
| 1024 |
});
|
| 1025 |
}
|
| 1026 |
|
|
|
|
| 1030 |
return {
|
| 1031 |
seed: state.seed,
|
| 1032 |
tasks: suiteRows,
|
| 1033 |
+
average_score: (() => {
|
| 1034 |
+
const numericScores = suiteRows
|
| 1035 |
+
.map((item) => item.score)
|
| 1036 |
+
.filter((value) => value !== null && value !== undefined)
|
| 1037 |
+
.map((value) => Number(value));
|
| 1038 |
+
if (!numericScores.length) {
|
| 1039 |
+
return null;
|
| 1040 |
+
}
|
| 1041 |
+
return Number((numericScores.reduce((acc, value) => acc + value, 0) / numericScores.length).toFixed(4));
|
| 1042 |
+
})(),
|
| 1043 |
};
|
| 1044 |
});
|
| 1045 |
}
|
grader/grader.py
CHANGED
|
@@ -22,6 +22,7 @@ DEFAULT_HIGH_PRIORITY_DEADLINE = 12
|
|
| 22 |
|
| 23 |
class GradeReport(BaseModel):
|
| 24 |
score: float = Field(..., ge=0.0, le=1.0)
|
|
|
|
| 25 |
|
| 26 |
delivered_orders: int = Field(..., ge=0)
|
| 27 |
total_orders: int = Field(..., ge=0)
|
|
@@ -190,6 +191,23 @@ class DeliveryEpisodeGrader:
|
|
| 190 |
|
| 191 |
score = round_score(raw_score, decimals=4)
|
| 192 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 193 |
logic = (
|
| 194 |
"metrics follow success_condition keys: completion_rate(_min), "
|
| 195 |
"high_priority_on_time_rate_min, max_steps, invalid_action_rate_max. "
|
|
@@ -211,6 +229,7 @@ class DeliveryEpisodeGrader:
|
|
| 211 |
|
| 212 |
return GradeReport(
|
| 213 |
score=score,
|
|
|
|
| 214 |
delivered_orders=delivered_orders,
|
| 215 |
total_orders=total_orders,
|
| 216 |
high_priority_total_orders=high_priority_total,
|
|
|
|
| 22 |
|
| 23 |
class GradeReport(BaseModel):
|
| 24 |
score: float = Field(..., ge=0.0, le=1.0)
|
| 25 |
+
success: bool
|
| 26 |
|
| 27 |
delivered_orders: int = Field(..., ge=0)
|
| 28 |
total_orders: int = Field(..., ge=0)
|
|
|
|
| 191 |
|
| 192 |
score = round_score(raw_score, decimals=4)
|
| 193 |
|
| 194 |
+
meets_completion = completion_rate_value >= targets.completion_target
|
| 195 |
+
meets_priority = True
|
| 196 |
+
if targets.priority_target is not None:
|
| 197 |
+
meets_priority = high_priority_on_time_rate >= targets.priority_target
|
| 198 |
+
meets_steps = episode_stats.steps_taken <= targets.max_steps_target
|
| 199 |
+
meets_invalid_actions = invalid_action_rate <= targets.invalid_action_rate_max
|
| 200 |
+
meets_battery_rule = (not targets.require_no_battery_depletion) or (episode_stats.battery_depletion_events == 0)
|
| 201 |
+
success = all(
|
| 202 |
+
[
|
| 203 |
+
meets_completion,
|
| 204 |
+
meets_priority,
|
| 205 |
+
meets_steps,
|
| 206 |
+
meets_invalid_actions,
|
| 207 |
+
meets_battery_rule,
|
| 208 |
+
]
|
| 209 |
+
)
|
| 210 |
+
|
| 211 |
logic = (
|
| 212 |
"metrics follow success_condition keys: completion_rate(_min), "
|
| 213 |
"high_priority_on_time_rate_min, max_steps, invalid_action_rate_max. "
|
|
|
|
| 229 |
|
| 230 |
return GradeReport(
|
| 231 |
score=score,
|
| 232 |
+
success=success,
|
| 233 |
delivered_orders=delivered_orders,
|
| 234 |
total_orders=total_orders,
|
| 235 |
high_priority_total_orders=high_priority_total,
|