Spaces:
Sleeping
Sleeping
vishal harkal commited on
Commit ·
f104717
1
Parent(s): 67e2e92
Deploy full OpenEnv API instead of starter app
Browse files- .dockerignore +10 -0
- .gitignore +9 -0
- Dockerfile +19 -11
- README.md +393 -1
- app.py +0 -8
- app/__init__.py +0 -0
- app/main.py +37 -0
- app/routes.py +396 -0
- app/static/results.html +1063 -0
- baseline/__init__.py +0 -0
- baseline/baseline_agent.py +102 -0
- baseline/trained_q_agent.py +176 -0
- env/__init__.py +0 -0
- env/environment.py +175 -0
- env/models.py +335 -0
- env/simulator.py +622 -0
- env/utils.py +58 -0
- grader/__init__.py +0 -0
- grader/grader.py +530 -0
- grader/metrics.py +178 -0
- inference.py +786 -0
- models/q_agent_easy.json +0 -0
- openenv.yaml +127 -0
- pyproject.toml +20 -0
- q_agent_easy.json +0 -0
- requirements.txt +6 -2
- scripts/train_q_agent.py +230 -0
- scripts/validate.sh +71 -0
- server/__init__.py +1 -0
- server/app.py +20 -0
- tasks/__init__.py +0 -0
- tasks/easy.py +71 -0
- tasks/hard.py +79 -0
- tasks/medium.py +74 -0
- tasks/registry.py +73 -0
- tests/test_api.py +263 -0
- tests/test_env.py +177 -0
- tests/test_grader.py +270 -0
- tests/test_inference_format.py +117 -0
- tests/test_real_world_constraints.py +148 -0
- tests/test_reward_shaping.py +146 -0
- tests/test_tasks.py +92 -0
- uv.lock +0 -0
.dockerignore
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
__pycache__/
|
| 2 |
+
*.pyc
|
| 3 |
+
*.pyo
|
| 4 |
+
*.pyd
|
| 5 |
+
.pytest_cache/
|
| 6 |
+
.git/
|
| 7 |
+
.gitignore
|
| 8 |
+
tests/
|
| 9 |
+
.env
|
| 10 |
+
.env.*
|
.gitignore
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
.venv/
|
| 2 |
+
__pycache__/
|
| 3 |
+
.pytest_cache/
|
| 4 |
+
*.py[cod]
|
| 5 |
+
*.pyo
|
| 6 |
+
*.pyd
|
| 7 |
+
*.egg-info/
|
| 8 |
+
.coverage
|
| 9 |
+
.DS_Store
|
Dockerfile
CHANGED
|
@@ -1,16 +1,24 @@
|
|
| 1 |
-
|
| 2 |
-
# you will also find guides on how best to write your Dockerfile
|
| 3 |
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
|
| 10 |
WORKDIR /app
|
| 11 |
|
| 12 |
-
|
| 13 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 14 |
|
| 15 |
-
|
| 16 |
-
CMD ["uvicorn", "app:app", "--host", "0.0.0.0", "--port", "7860"]
|
|
|
|
| 1 |
+
FROM python:3.10-slim
|
|
|
|
| 2 |
|
| 3 |
+
ENV PYTHONDONTWRITEBYTECODE=1 \
|
| 4 |
+
PYTHONUNBUFFERED=1 \
|
| 5 |
+
PIP_NO_CACHE_DIR=1 \
|
| 6 |
+
PIP_DISABLE_PIP_VERSION_CHECK=1 \
|
| 7 |
+
PORT=7860
|
| 8 |
|
| 9 |
WORKDIR /app
|
| 10 |
|
| 11 |
+
RUN adduser --disabled-password --gecos "" appuser
|
| 12 |
+
|
| 13 |
+
COPY requirements.txt /app/requirements.txt
|
| 14 |
+
RUN python -m pip install --upgrade pip && \
|
| 15 |
+
pip install -r /app/requirements.txt
|
| 16 |
+
|
| 17 |
+
COPY . /app
|
| 18 |
+
RUN chown -R appuser:appuser /app
|
| 19 |
+
|
| 20 |
+
USER appuser
|
| 21 |
+
|
| 22 |
+
EXPOSE 7860
|
| 23 |
|
| 24 |
+
CMD ["sh", "-c", "uvicorn app.main:app --host 0.0.0.0 --port ${PORT:-7860}"]
|
|
|
README.md
CHANGED
|
@@ -8,4 +8,396 @@ app_port: 7860
|
|
| 8 |
pinned: false
|
| 9 |
---
|
| 10 |
|
| 11 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 8 |
pinned: false
|
| 9 |
---
|
| 10 |
|
| 11 |
+
# Last-Mile Delivery Optimization (OpenEnv)
|
| 12 |
+
|
| 13 |
+
OpenEnv-compatible benchmark environment for city-scale last-mile dispatch under realistic operating constraints.
|
| 14 |
+
|
| 15 |
+
## Motivation
|
| 16 |
+
|
| 17 |
+
Urban delivery systems face a difficult planning problem: dispatchers must complete as many deliveries as possible while navigating congestion, road blockages, and battery limits. Poor routing or action sequencing creates delays, SLA violations for high-priority orders, and unnecessary energy cost.
|
| 18 |
+
|
| 19 |
+
This environment models that real-world challenge in a deterministic grid world so policies can be compared fairly across difficulty levels.
|
| 20 |
+
|
| 21 |
+
## What This Environment Simulates
|
| 22 |
+
|
| 23 |
+
- Multi-order pickup and delivery workflow.
|
| 24 |
+
- Static obstacles, dynamic obstacles, and traffic penalties.
|
| 25 |
+
- Priority-aware delivery pressure (especially in hard mode).
|
| 26 |
+
- Battery constraints and charging behavior in hard mode.
|
| 27 |
+
- Deterministic seeded task generation for reproducible evaluation.
|
| 28 |
+
|
| 29 |
+
## Architecture Diagram
|
| 30 |
+
|
| 31 |
+
```mermaid
|
| 32 |
+
flowchart LR
|
| 33 |
+
U[OpenEnv Evaluator or User] -->|HTTP| API[FastAPI API Layer\napp/main.py + app/routes.py]
|
| 34 |
+
|
| 35 |
+
API -->|reset or step or state| ENV[LastMileDeliveryEnvironment\nenv/environment.py]
|
| 36 |
+
ENV --> SIM[DeliverySimulator\nenv/simulator.py]
|
| 37 |
+
SIM -->|observation, reward, done, info| API
|
| 38 |
+
|
| 39 |
+
API -->|baseline endpoint rollout| BASE[BaselineGreedyAgent\nbaseline/baseline_agent.py]
|
| 40 |
+
BASE --> ENV
|
| 41 |
+
|
| 42 |
+
API -->|grader endpoint| GRADER[DeliveryEpisodeGrader\ngrader/grader.py]
|
| 43 |
+
GRADER --> REPORT[GradeReport\nscore + metrics + safeguards]
|
| 44 |
+
|
| 45 |
+
TASKS[tasks/easy.py + tasks/medium.py + tasks/hard.py\n+ tasks/registry.py] -->|task config| ENV
|
| 46 |
+
TASKS -->|success_condition| GRADER
|
| 47 |
+
TASKS -->|metadata| API
|
| 48 |
+
```
|
| 49 |
+
|
| 50 |
+
Flow summary:
|
| 51 |
+
|
| 52 |
+
- The API orchestrates environment state transitions and stores trajectory steps.
|
| 53 |
+
- The baseline endpoint runs a fresh rollout with the built-in greedy agent.
|
| 54 |
+
- The grader endpoint scores the recorded trajectory against task success conditions.
|
| 55 |
+
- Task definitions provide deterministic configuration and grading targets.
|
| 56 |
+
|
| 57 |
+
## Action Space
|
| 58 |
+
|
| 59 |
+
The action payload is a strict JSON object with exactly four fields.
|
| 60 |
+
|
| 61 |
+
| Field | Type | Required | Allowed Values | Notes |
|
| 62 |
+
|---|---|---|---|---|
|
| 63 |
+
| move | string or null | Yes | up, down, left, right, stay, null | Movement action; stay maps to no movement |
|
| 64 |
+
| accept_order | integer or null | Yes | order index or null | When set to n, maps to internal order_n |
|
| 65 |
+
| deliver_order | boolean | Yes | true or false | Deliver current order when true |
|
| 66 |
+
| wait | boolean | Yes | true or false | Explicit wait action when true |
|
| 67 |
+
|
| 68 |
+
Validation is strict. Exactly one action intent must be active per step.
|
| 69 |
+
|
| 70 |
+
## Observation Space
|
| 71 |
+
|
| 72 |
+
Each step returns an observation object with agent, order, and map state.
|
| 73 |
+
|
| 74 |
+
| Field | Type | Description |
|
| 75 |
+
|---|---|---|
|
| 76 |
+
| grid_width, grid_height | integer | Map dimensions |
|
| 77 |
+
| agent_location | object {x, y} | Current agent position |
|
| 78 |
+
| pending_orders | array of Order | Orders not yet delivered |
|
| 79 |
+
| current_order | Order or null | Active accepted/picked-up order |
|
| 80 |
+
| obstacles | array of Position | Static + dynamic blocked cells |
|
| 81 |
+
| dynamic_obstacles | array of Position | Current dynamic blocked cells |
|
| 82 |
+
| traffic_zones | array of TrafficZone | Cells with extra movement cost |
|
| 83 |
+
| charging_stations | array of Position | Recharge cells |
|
| 84 |
+
| battery_level | integer or null | Active in battery-enabled tasks |
|
| 85 |
+
| step_count | integer | Elapsed step count |
|
| 86 |
+
| total_reward | float | Cumulative episode reward |
|
| 87 |
+
|
| 88 |
+
## Tasks
|
| 89 |
+
|
| 90 |
+
All tasks are deterministic and expose explicit success conditions.
|
| 91 |
+
|
| 92 |
+
### Easy
|
| 93 |
+
|
| 94 |
+
- Description: simple single-order delivery on a compact map.
|
| 95 |
+
- Configuration: 6x6 grid, 1 order, max_steps=60.
|
| 96 |
+
- Constraints: no obstacles, no traffic, no battery, no priority mix.
|
| 97 |
+
- Success condition: completion_rate=1.0, max_steps=60, invalid_action_rate_max=0.12.
|
| 98 |
+
|
| 99 |
+
### Medium
|
| 100 |
+
|
| 101 |
+
- Description: multi-order delivery with static obstacles and traffic friction.
|
| 102 |
+
- Configuration: 10x10 grid, 3 to 4 orders, max_steps=140.
|
| 103 |
+
- Constraints: static obstacles + traffic enabled.
|
| 104 |
+
- Success condition: completion_rate_min=0.85, max_steps=140, invalid_action_rate_max=0.08.
|
| 105 |
+
|
| 106 |
+
### Hard
|
| 107 |
+
|
| 108 |
+
- Description: high-pressure dispatch with tight battery and SLA constraints.
|
| 109 |
+
- Configuration: 16x16 grid, 6 to 8 orders, max_steps=165.
|
| 110 |
+
- Constraints: static + dynamic obstacles, traffic, high-priority orders, multi-stop delivery, battery + recharge, delay pressure.
|
| 111 |
+
- Success condition: completion_rate_min=0.90, high_priority_on_time_rate_min=0.88, battery_depletion=false, max_steps=165, invalid_action_rate_max=0.04.
|
| 112 |
+
|
| 113 |
+
## Reward Design
|
| 114 |
+
|
| 115 |
+
The reward is dense and combines task completion signals with operational efficiency:
|
| 116 |
+
|
| 117 |
+
- Base step penalty: -1.0 each step.
|
| 118 |
+
- Delivery reward: +50.0 on successful delivery.
|
| 119 |
+
- Destination milestone: +10.0 when pickup/drop destination is reached.
|
| 120 |
+
- Invalid action penalty: -20.0.
|
| 121 |
+
- Traffic movement penalty: -2.0 when entering traffic cells.
|
| 122 |
+
- Progress shaping: signed distance-based reward clipped by configured bounds.
|
| 123 |
+
- Priority delivery bonus: +15.0 (high) or +5.0 (low).
|
| 124 |
+
- Recharge reward: +2.0 when battery is restored at charging stations.
|
| 125 |
+
- Delay penalty: scaled negative penalty for overdue active orders.
|
| 126 |
+
- Battery failure penalty: -30.0 when battery depletes to zero.
|
| 127 |
+
|
| 128 |
+
Design goal: favor correct and efficient order handling while discouraging invalid and wasteful behavior.
|
| 129 |
+
|
| 130 |
+
## Grader Logic
|
| 131 |
+
|
| 132 |
+
The grader is deterministic and aligned to each task's success_condition fields.
|
| 133 |
+
|
| 134 |
+
Main metrics:
|
| 135 |
+
|
| 136 |
+
- completion_rate: delivered_orders / total_orders.
|
| 137 |
+
- high_priority_on_time_rate: on-time high-priority deliveries.
|
| 138 |
+
- efficiency_ratio: normalized by max_steps target.
|
| 139 |
+
- invalid_action_rate: invalid_actions / steps_taken.
|
| 140 |
+
|
| 141 |
+
Scoring components:
|
| 142 |
+
|
| 143 |
+
- Completion component: highest weight.
|
| 144 |
+
- Priority SLA component: enabled when high_priority_on_time_rate_min exists.
|
| 145 |
+
- Efficiency component: step-budget performance.
|
| 146 |
+
- Penalty component: reduces score when invalid action rate exceeds target.
|
| 147 |
+
|
| 148 |
+
Edge-case safeguards include:
|
| 149 |
+
|
| 150 |
+
- Zero-step and zero-order episode handling.
|
| 151 |
+
- Invalid success-condition value fallback and clamping.
|
| 152 |
+
- Division-safe invalid-action and efficiency calculations.
|
| 153 |
+
- Optional battery depletion penalty when the task requires no depletion.
|
| 154 |
+
|
| 155 |
+
## Setup
|
| 156 |
+
|
| 157 |
+
### Local
|
| 158 |
+
|
| 159 |
+
1. Create and activate a virtual environment.
|
| 160 |
+
2. Install dependencies.
|
| 161 |
+
3. Run the API server on port 7860.
|
| 162 |
+
|
| 163 |
+
```bash
|
| 164 |
+
python -m venv .venv
|
| 165 |
+
source .venv/bin/activate
|
| 166 |
+
pip install -r requirements.txt
|
| 167 |
+
uvicorn app.main:app --host 0.0.0.0 --port 7860
|
| 168 |
+
```
|
| 169 |
+
|
| 170 |
+
### Docker
|
| 171 |
+
|
| 172 |
+
```bash
|
| 173 |
+
docker build -t delivery-openenv:latest .
|
| 174 |
+
docker run --rm -p 7860:7860 \
|
| 175 |
+
-e HF_TOKEN="example-provider-key" \
|
| 176 |
+
-e MODEL_NAME="gpt-4o-mini" \
|
| 177 |
+
-e API_BASE_URL="https://api.openai.com/v1" \
|
| 178 |
+
-e OPENAI_USE_MODEL="1" \
|
| 179 |
+
delivery-openenv:latest
|
| 180 |
+
```
|
| 181 |
+
|
| 182 |
+
## API Usage
|
| 183 |
+
|
| 184 |
+
Base URL:
|
| 185 |
+
|
| 186 |
+
```text
|
| 187 |
+
http://0.0.0.0:7860
|
| 188 |
+
```
|
| 189 |
+
|
| 190 |
+
### Health
|
| 191 |
+
|
| 192 |
+
```bash
|
| 193 |
+
curl -s http://0.0.0.0:7860/health
|
| 194 |
+
```
|
| 195 |
+
|
| 196 |
+
### Frontend Results Dashboard
|
| 197 |
+
|
| 198 |
+
Open the interactive dashboard in your browser:
|
| 199 |
+
|
| 200 |
+
```bash
|
| 201 |
+
open http://0.0.0.0:7860/ui
|
| 202 |
+
```
|
| 203 |
+
|
| 204 |
+
The dashboard lets you run `reset`, `step`, `state`, `grader`, and `baseline` calls and inspect live payloads and metrics.
|
| 205 |
+
|
| 206 |
+
### List Tasks and Action Schema
|
| 207 |
+
|
| 208 |
+
```bash
|
| 209 |
+
curl -s http://0.0.0.0:7860/tasks
|
| 210 |
+
```
|
| 211 |
+
|
| 212 |
+
### Reset Environment
|
| 213 |
+
|
| 214 |
+
```bash
|
| 215 |
+
curl -s -X POST http://0.0.0.0:7860/reset \
|
| 216 |
+
-H "Content-Type: application/json" \
|
| 217 |
+
-d '{"task":"easy","seed":101}'
|
| 218 |
+
```
|
| 219 |
+
|
| 220 |
+
### Reset From Real-World Scenario
|
| 221 |
+
|
| 222 |
+
Use this endpoint when you want to inject externally curated map and order state instead of task-generated layouts.
|
| 223 |
+
|
| 224 |
+
```bash
|
| 225 |
+
curl -s -X POST http://0.0.0.0:7860/reset_from_scenario \
|
| 226 |
+
-H "Content-Type: application/json" \
|
| 227 |
+
-d '{
|
| 228 |
+
"seed": 99,
|
| 229 |
+
"success_condition": {
|
| 230 |
+
"completion_rate_min": 1.0,
|
| 231 |
+
"max_steps": 20,
|
| 232 |
+
"invalid_action_rate_max": 0.20,
|
| 233 |
+
"battery_depletion": false
|
| 234 |
+
},
|
| 235 |
+
"scenario": {
|
| 236 |
+
"width": 6,
|
| 237 |
+
"height": 6,
|
| 238 |
+
"max_steps": 20,
|
| 239 |
+
"agent_start": {"x": 0, "y": 0},
|
| 240 |
+
"orders": [
|
| 241 |
+
{
|
| 242 |
+
"order_id": "r1",
|
| 243 |
+
"pickup": {"x": 1, "y": 0},
|
| 244 |
+
"dropoff": {"x": 2, "y": 0},
|
| 245 |
+
"delivery_locations": [{"x": 2, "y": 0}],
|
| 246 |
+
"priority": "high",
|
| 247 |
+
"created_step": 0
|
| 248 |
+
}
|
| 249 |
+
],
|
| 250 |
+
"obstacles": [{"x": 4, "y": 4}],
|
| 251 |
+
"dynamic_obstacles": [{"x": 4, "y": 5}],
|
| 252 |
+
"traffic_zones": [{"location": {"x": 3, "y": 0}, "extra_cost": 3}],
|
| 253 |
+
"charging_stations": [{"x": 0, "y": 0}],
|
| 254 |
+
"battery_profile": {
|
| 255 |
+
"enabled": true,
|
| 256 |
+
"capacity": 10,
|
| 257 |
+
"recharge_rate": 4,
|
| 258 |
+
"initial_level": 7
|
| 259 |
+
}
|
| 260 |
+
}
|
| 261 |
+
}'
|
| 262 |
+
```
|
| 263 |
+
|
| 264 |
+
Contract summary for `scenario`:
|
| 265 |
+
|
| 266 |
+
- `width`, `height`, `max_steps`: grid and episode budget.
|
| 267 |
+
- `agent_start`: starting position.
|
| 268 |
+
- `orders`: list of orders with pickup/dropoff/delivery locations and priority.
|
| 269 |
+
- `obstacles`, `dynamic_obstacles`: blocked cells.
|
| 270 |
+
- `traffic_zones`: movement-cost cells (`extra_cost >= 1`).
|
| 271 |
+
- `charging_stations`: recharge cells.
|
| 272 |
+
- `battery_profile`: battery enablement, capacity, recharge rate, and optional initial level.
|
| 273 |
+
- `success_condition` (top-level): optional custom grader thresholds for this scenario.
|
| 274 |
+
|
| 275 |
+
### Step
|
| 276 |
+
|
| 277 |
+
```bash
|
| 278 |
+
curl -s -X POST http://0.0.0.0:7860/step \
|
| 279 |
+
-H "Content-Type: application/json" \
|
| 280 |
+
-d '{"move":null,"accept_order":null,"deliver_order":false,"wait":true}'
|
| 281 |
+
```
|
| 282 |
+
|
| 283 |
+
### State
|
| 284 |
+
|
| 285 |
+
```bash
|
| 286 |
+
curl -s http://0.0.0.0:7860/state
|
| 287 |
+
```
|
| 288 |
+
|
| 289 |
+
### Grade Current Episode
|
| 290 |
+
|
| 291 |
+
```bash
|
| 292 |
+
curl -s http://0.0.0.0:7860/grader
|
| 293 |
+
```
|
| 294 |
+
|
| 295 |
+
### Baseline Rollout Evaluation
|
| 296 |
+
|
| 297 |
+
```bash
|
| 298 |
+
curl -s "http://0.0.0.0:7860/baseline?task=easy"
|
| 299 |
+
```
|
| 300 |
+
|
| 301 |
+
## Baseline Scores (Reproducible Inference Run)
|
| 302 |
+
|
| 303 |
+
Measured with `inference.py --all-tasks --seed 42 --json-summary` using deterministic fallback mode (`OPENAI_USE_MODEL=0`).
|
| 304 |
+
|
| 305 |
+
| Task | Steps | Success | Score |
|
| 306 |
+
|---|---:|---:|---:|
|
| 307 |
+
| easy | 6 | true | 0.9700 |
|
| 308 |
+
| medium | 140 | false | 0.0001 |
|
| 309 |
+
| hard | 55 | false | 0.0667 |
|
| 310 |
+
|
| 311 |
+
Average score (seed=42): **0.3456**
|
| 312 |
+
|
| 313 |
+
## Real-World Validation
|
| 314 |
+
|
| 315 |
+
We tested the environment under realistic scenarios:
|
| 316 |
+
|
| 317 |
+
- Priority-based delivery selection.
|
| 318 |
+
- Traffic-aware routing decisions.
|
| 319 |
+
- Battery-constrained navigation.
|
| 320 |
+
|
| 321 |
+
Results indicate that evaluated agents can exhibit behavior aligned with real-world logistics systems under these constraints.
|
| 322 |
+
|
| 323 |
+
## OpenEnv Manifest
|
| 324 |
+
|
| 325 |
+
Project metadata and runtime wiring are defined in openenv.yaml, including:
|
| 326 |
+
|
| 327 |
+
- runtime entrypoints,
|
| 328 |
+
- environment class,
|
| 329 |
+
- task registry loader,
|
| 330 |
+
- grader class,
|
| 331 |
+
- action and observation schema declarations.
|
| 332 |
+
|
| 333 |
+
## Baseline Inference Script
|
| 334 |
+
|
| 335 |
+
The baseline runner uses the OpenAI Python client and supports the required submission variables:
|
| 336 |
+
|
| 337 |
+
- `HF_TOKEN`: API key/token used as OpenAI client `api_key`.
|
| 338 |
+
- `MODEL_NAME`: model identifier used for inference requests.
|
| 339 |
+
- `API_BASE_URL`: optional OpenAI-compatible endpoint.
|
| 340 |
+
|
| 341 |
+
Backward-compatible aliases are also supported: `OPENAI_API_KEY`, `OPENAI_MODEL`, and `OPENAI_BASE_URL`.
|
| 342 |
+
|
| 343 |
+
```bash
|
| 344 |
+
export HF_TOKEN="your-api-key"
|
| 345 |
+
export MODEL_NAME="gpt-4o-mini"
|
| 346 |
+
export API_BASE_URL="https://api.openai.com/v1"
|
| 347 |
+
export OPENAI_USE_MODEL="1"
|
| 348 |
+
|
| 349 |
+
# Reproducible baseline score across easy/medium/hard tasks.
|
| 350 |
+
python inference.py --all-tasks --seed 42 --json-summary
|
| 351 |
+
```
|
| 352 |
+
|
| 353 |
+
Notes:
|
| 354 |
+
|
| 355 |
+
- `OPENAI_USE_MODEL=1` enables remote model calls.
|
| 356 |
+
- `OPENAI_USE_MODEL=0` keeps deterministic heuristic fallback while still validating OpenAI client wiring.
|
| 357 |
+
- `--seed` controls reproducibility for baseline comparisons.
|
| 358 |
+
|
| 359 |
+
## Trainable Agent (Q-Learning)
|
| 360 |
+
|
| 361 |
+
This repository now includes a trainable tabular Q-learning agent for interactive learning on the OpenEnv loop.
|
| 362 |
+
|
| 363 |
+
Training script:
|
| 364 |
+
|
| 365 |
+
```bash
|
| 366 |
+
python scripts/train_q_agent.py --task easy --episodes 2000 --eval-episodes 200 --seed 42 --output models/q_agent_easy.json
|
| 367 |
+
```
|
| 368 |
+
|
| 369 |
+
What it does:
|
| 370 |
+
|
| 371 |
+
- Trains through repeated `reset()` / `step()` episodes.
|
| 372 |
+
- Learns action values for a compact state representation.
|
| 373 |
+
- Evaluates the learned policy and prints success/completion metrics.
|
| 374 |
+
- Saves a reusable model artifact (`models/q_agent_easy.json`).
|
| 375 |
+
|
| 376 |
+
Learner runtime components:
|
| 377 |
+
|
| 378 |
+
- Agent implementation: `baseline/trained_q_agent.py`
|
| 379 |
+
- Trainer entrypoint: `scripts/train_q_agent.py`
|
| 380 |
+
|
| 381 |
+
## Validation Commands
|
| 382 |
+
|
| 383 |
+
```bash
|
| 384 |
+
./scripts/validate.sh
|
| 385 |
+
.venv/bin/openenv validate
|
| 386 |
+
python -m pytest -q
|
| 387 |
+
```
|
| 388 |
+
|
| 389 |
+
## Test Command
|
| 390 |
+
|
| 391 |
+
```bash
|
| 392 |
+
python -m pytest -q
|
| 393 |
+
```
|
| 394 |
+
|
| 395 |
+
## Hugging Face Spaces (Docker)
|
| 396 |
+
|
| 397 |
+
1. Create a Docker Space.
|
| 398 |
+
2. Push this repository.
|
| 399 |
+
3. Set variables/secrets: HF_TOKEN, MODEL_NAME, API_BASE_URL, OPENAI_USE_MODEL.
|
| 400 |
+
4. Ensure exposed port is 7860.
|
| 401 |
+
5. Rebuild and verify /health.
|
| 402 |
+
|
| 403 |
+
Additional deployment notes are available in HF_SPACES_DEPLOYMENT.md.
|
app.py
DELETED
|
@@ -1,8 +0,0 @@
|
|
| 1 |
-
from fastapi import FastAPI
|
| 2 |
-
|
| 3 |
-
app = FastAPI()
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
@app.get("/")
|
| 7 |
-
def greet_json():
|
| 8 |
-
return {"Hello": "World!"}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
app/__init__.py
ADDED
|
File without changes
|
app/main.py
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
|
| 5 |
+
from fastapi import FastAPI
|
| 6 |
+
from fastapi.middleware.cors import CORSMiddleware
|
| 7 |
+
from fastapi.responses import RedirectResponse
|
| 8 |
+
from fastapi.staticfiles import StaticFiles
|
| 9 |
+
|
| 10 |
+
from app.routes import router
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
app = FastAPI(
|
| 14 |
+
title="Last-Mile Delivery Optimization Environment",
|
| 15 |
+
version="1.0.0",
|
| 16 |
+
description="OpenEnv-compatible environment for grid-based last-mile delivery optimization.",
|
| 17 |
+
)
|
| 18 |
+
|
| 19 |
+
app.add_middleware(
|
| 20 |
+
CORSMiddleware,
|
| 21 |
+
allow_origins=["*"],
|
| 22 |
+
allow_credentials=False,
|
| 23 |
+
allow_methods=["*"],
|
| 24 |
+
allow_headers=["*"],
|
| 25 |
+
)
|
| 26 |
+
|
| 27 |
+
STATIC_DIR = Path(__file__).resolve().parent / "static"
|
| 28 |
+
|
| 29 |
+
app.mount("/static", StaticFiles(directory=str(STATIC_DIR)), name="static")
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
@app.get("/ui", include_in_schema=False)
|
| 33 |
+
def results_ui() -> RedirectResponse:
|
| 34 |
+
return RedirectResponse(url="/static/results.html")
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
app.include_router(router)
|
app/routes.py
ADDED
|
@@ -0,0 +1,396 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from threading import RLock
|
| 4 |
+
from typing import Any, Dict, List, Optional
|
| 5 |
+
|
| 6 |
+
from fastapi import APIRouter, Body, HTTPException
|
| 7 |
+
from pydantic import BaseModel, Field, ValidationError
|
| 8 |
+
|
| 9 |
+
from baseline.baseline_agent import BaselineGreedyAgent
|
| 10 |
+
from env.environment import LastMileDeliveryEnvironment
|
| 11 |
+
from env.models import Action, EnvironmentConfig, EnvironmentState, Observation, ScenarioDefinition, StepResult
|
| 12 |
+
from grader.grader import DeliveryEpisodeGrader, GradeReport
|
| 13 |
+
from tasks.registry import get_all_task_metadata, get_task_config, get_task_definition
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
router = APIRouter()
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class ResetRequest(BaseModel):
|
| 20 |
+
task: str = "easy"
|
| 21 |
+
seed: int | None = None
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class ScenarioResetRequest(BaseModel):
|
| 25 |
+
seed: int | None = None
|
| 26 |
+
success_condition: Dict[str, Any] = Field(default_factory=dict)
|
| 27 |
+
scenario: ScenarioDefinition
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
class GraderResponse(BaseModel):
|
| 31 |
+
task: str | None
|
| 32 |
+
steps_recorded: int
|
| 33 |
+
report: GradeReport
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class BaselineResponse(BaseModel):
|
| 37 |
+
task: str
|
| 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 |
+
|
| 45 |
+
|
| 46 |
+
class TaskDescriptor(BaseModel):
|
| 47 |
+
name: str
|
| 48 |
+
description: str
|
| 49 |
+
difficulty: int
|
| 50 |
+
parameters: Dict[str, Any]
|
| 51 |
+
success_condition: Dict[str, Any]
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class TasksResponse(BaseModel):
|
| 55 |
+
tasks: List[TaskDescriptor]
|
| 56 |
+
action_schema: Dict[str, Any]
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
TASKS_ACTION_SCHEMA: Dict[str, Any] = {
|
| 60 |
+
"type": "object",
|
| 61 |
+
"required": ["move", "accept_order", "deliver_order", "wait"],
|
| 62 |
+
"additionalProperties": False,
|
| 63 |
+
"properties": {
|
| 64 |
+
"move": {
|
| 65 |
+
"type": ["string", "null"],
|
| 66 |
+
"enum": ["up", "down", "left", "right", "stay", None],
|
| 67 |
+
},
|
| 68 |
+
"accept_order": {
|
| 69 |
+
"type": ["integer", "null"],
|
| 70 |
+
},
|
| 71 |
+
"deliver_order": {
|
| 72 |
+
"type": "boolean",
|
| 73 |
+
},
|
| 74 |
+
"wait": {
|
| 75 |
+
"type": "boolean",
|
| 76 |
+
},
|
| 77 |
+
},
|
| 78 |
+
}
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
class EnvironmentService:
|
| 82 |
+
BASELINE_MAX_STEPS = 128
|
| 83 |
+
|
| 84 |
+
def __init__(self) -> None:
|
| 85 |
+
self._lock = RLock()
|
| 86 |
+
self._env: Optional[LastMileDeliveryEnvironment] = None
|
| 87 |
+
self._task: Optional[str] = None
|
| 88 |
+
self._success_condition: Dict[str, Any] = {}
|
| 89 |
+
self._trajectory: list[StepResult] = []
|
| 90 |
+
self._grader = DeliveryEpisodeGrader()
|
| 91 |
+
self._baseline_agent = BaselineGreedyAgent()
|
| 92 |
+
|
| 93 |
+
def reset(self, task: str, seed: int | None) -> Observation:
|
| 94 |
+
with self._lock:
|
| 95 |
+
config = get_task_config(task_name=task, seed=seed)
|
| 96 |
+
task_definition = get_task_definition(task)
|
| 97 |
+
self._env = LastMileDeliveryEnvironment(config=config)
|
| 98 |
+
self._task = task
|
| 99 |
+
self._success_condition = dict(task_definition.get("success_condition", {}))
|
| 100 |
+
self._trajectory = []
|
| 101 |
+
return self._env.reset(seed=seed)
|
| 102 |
+
|
| 103 |
+
def reset_from_scenario(
|
| 104 |
+
self,
|
| 105 |
+
scenario: ScenarioDefinition,
|
| 106 |
+
seed: int | None,
|
| 107 |
+
success_condition: Dict[str, Any] | None,
|
| 108 |
+
) -> Observation:
|
| 109 |
+
with self._lock:
|
| 110 |
+
config = EnvironmentConfig(
|
| 111 |
+
width=scenario.width,
|
| 112 |
+
height=scenario.height,
|
| 113 |
+
max_steps=scenario.max_steps,
|
| 114 |
+
max_orders=max(1, len(scenario.orders)),
|
| 115 |
+
obstacle_density=0.0,
|
| 116 |
+
traffic_density=0.0,
|
| 117 |
+
dynamic_obstacles_enabled=bool(scenario.dynamic_obstacles),
|
| 118 |
+
dynamic_obstacle_ratio=1.0 if scenario.dynamic_obstacles else 0.0,
|
| 119 |
+
battery_enabled=scenario.battery_profile.enabled,
|
| 120 |
+
battery_capacity=scenario.battery_profile.capacity,
|
| 121 |
+
battery_recharge_rate=scenario.battery_profile.recharge_rate,
|
| 122 |
+
charging_stations=[item.model_copy(deep=True) for item in scenario.charging_stations],
|
| 123 |
+
seed=seed,
|
| 124 |
+
)
|
| 125 |
+
self._env = LastMileDeliveryEnvironment(config=config)
|
| 126 |
+
self._task = "scenario"
|
| 127 |
+
self._success_condition = dict(success_condition or {})
|
| 128 |
+
self._trajectory = []
|
| 129 |
+
return self._env.reset_with_scenario(scenario=scenario, seed=seed)
|
| 130 |
+
|
| 131 |
+
def step(self, action: Action) -> StepResult:
|
| 132 |
+
with self._lock:
|
| 133 |
+
if self._env is None:
|
| 134 |
+
raise RuntimeError("Environment not initialized. Call /reset first.")
|
| 135 |
+
result = self._env.step_result(action)
|
| 136 |
+
self._trajectory.append(result)
|
| 137 |
+
return result
|
| 138 |
+
|
| 139 |
+
def state(self) -> EnvironmentState:
|
| 140 |
+
with self._lock:
|
| 141 |
+
if self._env is None:
|
| 142 |
+
return EnvironmentState(
|
| 143 |
+
initialized=False,
|
| 144 |
+
done=False,
|
| 145 |
+
total_reward=0.0,
|
| 146 |
+
step_count=0,
|
| 147 |
+
delivered_orders=0,
|
| 148 |
+
current_order_id=None,
|
| 149 |
+
observation=None,
|
| 150 |
+
)
|
| 151 |
+
return self._env.state()
|
| 152 |
+
|
| 153 |
+
def current_task(self) -> Optional[str]:
|
| 154 |
+
with self._lock:
|
| 155 |
+
return self._task
|
| 156 |
+
|
| 157 |
+
def grade(self) -> GradeReport:
|
| 158 |
+
with self._lock:
|
| 159 |
+
if self._env is None:
|
| 160 |
+
raise RuntimeError("Environment not initialized. Call /reset first.")
|
| 161 |
+
final_observation = self._env.current_observation()
|
| 162 |
+
success_condition = dict(self._success_condition)
|
| 163 |
+
if not success_condition:
|
| 164 |
+
resolved_task = self._task or "easy"
|
| 165 |
+
try:
|
| 166 |
+
task_definition = get_task_definition(resolved_task)
|
| 167 |
+
success_condition = dict(task_definition.get("success_condition", {}))
|
| 168 |
+
except ValueError:
|
| 169 |
+
success_condition = {}
|
| 170 |
+
return self._grader.grade_episode(
|
| 171 |
+
self._trajectory,
|
| 172 |
+
final_observation,
|
| 173 |
+
success_condition=success_condition,
|
| 174 |
+
)
|
| 175 |
+
|
| 176 |
+
def baseline_rollout(
|
| 177 |
+
self,
|
| 178 |
+
task: str | None = None,
|
| 179 |
+
seed: int | None = None,
|
| 180 |
+
max_steps: int | None = None,
|
| 181 |
+
) -> BaselineResponse:
|
| 182 |
+
with self._lock:
|
| 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:
|
| 190 |
+
step_cap = min(step_cap, max_steps)
|
| 191 |
+
step_cap = max(1, min(step_cap, self.BASELINE_MAX_STEPS))
|
| 192 |
+
|
| 193 |
+
trajectory: list[StepResult] = []
|
| 194 |
+
for _ in range(step_cap):
|
| 195 |
+
action = self._baseline_agent.act(observation)
|
| 196 |
+
result = rollout_env.step_result(action)
|
| 197 |
+
trajectory.append(result)
|
| 198 |
+
observation = result.observation
|
| 199 |
+
if result.done:
|
| 200 |
+
break
|
| 201 |
+
|
| 202 |
+
task_definition = get_task_definition(resolved_task)
|
| 203 |
+
success_condition = task_definition.get("success_condition", {})
|
| 204 |
+
report = self._grader.grade_episode(
|
| 205 |
+
trajectory,
|
| 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 |
+
done=rollout_env.state().done,
|
| 215 |
+
score=report.score,
|
| 216 |
+
report=report,
|
| 217 |
+
)
|
| 218 |
+
|
| 219 |
+
def steps_recorded(self) -> int:
|
| 220 |
+
with self._lock:
|
| 221 |
+
return len(self._trajectory)
|
| 222 |
+
|
| 223 |
+
|
| 224 |
+
def _http_error(status_code: int, message: str) -> HTTPException:
|
| 225 |
+
return HTTPException(
|
| 226 |
+
status_code=status_code,
|
| 227 |
+
detail={
|
| 228 |
+
"error": "bad_request" if status_code == 400 else "internal_server_error",
|
| 229 |
+
"message": message,
|
| 230 |
+
},
|
| 231 |
+
)
|
| 232 |
+
|
| 233 |
+
|
| 234 |
+
def _parse_step_action(payload: Any) -> Action:
|
| 235 |
+
if payload is None:
|
| 236 |
+
raise ValueError("Missing step payload")
|
| 237 |
+
|
| 238 |
+
if not isinstance(payload, dict):
|
| 239 |
+
raise ValueError("Step payload must be a JSON object")
|
| 240 |
+
|
| 241 |
+
allowed_keys = {"move", "accept_order", "deliver_order", "wait"}
|
| 242 |
+
unexpected_keys = sorted(set(payload.keys()) - allowed_keys)
|
| 243 |
+
if unexpected_keys:
|
| 244 |
+
raise ValueError(f"Unexpected action fields: {', '.join(unexpected_keys)}")
|
| 245 |
+
|
| 246 |
+
missing_keys = sorted(allowed_keys - set(payload.keys()))
|
| 247 |
+
if missing_keys:
|
| 248 |
+
raise ValueError(f"Missing action fields: {', '.join(missing_keys)}")
|
| 249 |
+
|
| 250 |
+
try:
|
| 251 |
+
return Action.model_validate(payload)
|
| 252 |
+
except ValidationError as exc:
|
| 253 |
+
raise ValueError("Invalid action payload") from exc
|
| 254 |
+
|
| 255 |
+
|
| 256 |
+
def _parse_scenario_reset_request(payload: Any) -> ScenarioResetRequest:
|
| 257 |
+
if payload is None:
|
| 258 |
+
raise ValueError("Missing scenario reset payload")
|
| 259 |
+
|
| 260 |
+
if not isinstance(payload, dict):
|
| 261 |
+
raise ValueError("Scenario reset payload must be a JSON object")
|
| 262 |
+
|
| 263 |
+
try:
|
| 264 |
+
return ScenarioResetRequest.model_validate(payload)
|
| 265 |
+
except ValidationError as exc:
|
| 266 |
+
raise ValueError("Invalid scenario payload") from exc
|
| 267 |
+
|
| 268 |
+
|
| 269 |
+
def _normalize_task_metadata(raw_tasks: List[Dict[str, Any]]) -> List[TaskDescriptor]:
|
| 270 |
+
normalized: List[TaskDescriptor] = []
|
| 271 |
+
for item in raw_tasks:
|
| 272 |
+
missing = [
|
| 273 |
+
field
|
| 274 |
+
for field in ("name", "description", "difficulty", "parameters", "success_condition")
|
| 275 |
+
if field not in item
|
| 276 |
+
]
|
| 277 |
+
if missing:
|
| 278 |
+
raise ValueError(f"Task metadata missing required fields: {', '.join(missing)}")
|
| 279 |
+
|
| 280 |
+
normalized.append(
|
| 281 |
+
TaskDescriptor(
|
| 282 |
+
name=str(item["name"]),
|
| 283 |
+
description=str(item["description"]),
|
| 284 |
+
difficulty=int(item["difficulty"]),
|
| 285 |
+
parameters=dict(item["parameters"]),
|
| 286 |
+
success_condition=dict(item["success_condition"]),
|
| 287 |
+
)
|
| 288 |
+
)
|
| 289 |
+
|
| 290 |
+
return normalized
|
| 291 |
+
|
| 292 |
+
|
| 293 |
+
env_service = EnvironmentService()
|
| 294 |
+
|
| 295 |
+
|
| 296 |
+
@router.get("/")
|
| 297 |
+
def root() -> dict:
|
| 298 |
+
return {
|
| 299 |
+
"name": "last-mile-delivery-optimization",
|
| 300 |
+
"status": "ok",
|
| 301 |
+
"health": "/health",
|
| 302 |
+
"docs": "/docs",
|
| 303 |
+
}
|
| 304 |
+
|
| 305 |
+
|
| 306 |
+
@router.get("/health")
|
| 307 |
+
def health() -> dict:
|
| 308 |
+
return {"status": "ok"}
|
| 309 |
+
|
| 310 |
+
|
| 311 |
+
@router.get("/tasks", response_model=TasksResponse)
|
| 312 |
+
def list_tasks() -> TasksResponse:
|
| 313 |
+
try:
|
| 314 |
+
raw_tasks = get_all_task_metadata()
|
| 315 |
+
tasks = _normalize_task_metadata(raw_tasks)
|
| 316 |
+
return TasksResponse(tasks=tasks, action_schema=TASKS_ACTION_SCHEMA)
|
| 317 |
+
except ValueError as exc:
|
| 318 |
+
raise _http_error(500, str(exc)) from exc
|
| 319 |
+
except Exception as exc:
|
| 320 |
+
raise _http_error(500, "Failed to retrieve task metadata") from exc
|
| 321 |
+
|
| 322 |
+
|
| 323 |
+
@router.post("/reset", response_model=Observation)
|
| 324 |
+
def reset_environment(request: ResetRequest | None = Body(default=None)) -> Observation:
|
| 325 |
+
try:
|
| 326 |
+
payload = request or ResetRequest()
|
| 327 |
+
return env_service.reset(task=payload.task, seed=payload.seed)
|
| 328 |
+
except ValueError as exc:
|
| 329 |
+
raise _http_error(400, str(exc)) from exc
|
| 330 |
+
except Exception as exc:
|
| 331 |
+
raise _http_error(500, "Failed to reset environment") from exc
|
| 332 |
+
|
| 333 |
+
|
| 334 |
+
@router.post("/reset_from_scenario", response_model=Observation)
|
| 335 |
+
def reset_environment_from_scenario(payload: Any = Body(default=None)) -> Observation:
|
| 336 |
+
try:
|
| 337 |
+
request = _parse_scenario_reset_request(payload)
|
| 338 |
+
return env_service.reset_from_scenario(
|
| 339 |
+
scenario=request.scenario,
|
| 340 |
+
seed=request.seed,
|
| 341 |
+
success_condition=request.success_condition,
|
| 342 |
+
)
|
| 343 |
+
except ValueError as exc:
|
| 344 |
+
raise _http_error(400, str(exc)) from exc
|
| 345 |
+
except Exception as exc:
|
| 346 |
+
raise _http_error(500, "Failed to reset environment from scenario") from exc
|
| 347 |
+
|
| 348 |
+
|
| 349 |
+
@router.post("/step", response_model=StepResult)
|
| 350 |
+
def environment_step(payload: Any = Body(default=None)) -> StepResult:
|
| 351 |
+
try:
|
| 352 |
+
action = _parse_step_action(payload)
|
| 353 |
+
return env_service.step(action=action)
|
| 354 |
+
except ValueError as exc:
|
| 355 |
+
raise _http_error(400, str(exc)) from exc
|
| 356 |
+
except RuntimeError as exc:
|
| 357 |
+
raise _http_error(400, str(exc)) from exc
|
| 358 |
+
except Exception as exc:
|
| 359 |
+
raise _http_error(500, "Failed to execute step") from exc
|
| 360 |
+
|
| 361 |
+
|
| 362 |
+
@router.get("/state", response_model=EnvironmentState)
|
| 363 |
+
def environment_state() -> EnvironmentState:
|
| 364 |
+
try:
|
| 365 |
+
return env_service.state()
|
| 366 |
+
except RuntimeError as exc:
|
| 367 |
+
raise _http_error(400, str(exc)) from exc
|
| 368 |
+
except Exception as exc:
|
| 369 |
+
raise _http_error(500, "Failed to retrieve state") from exc
|
| 370 |
+
|
| 371 |
+
|
| 372 |
+
@router.get("/grader", response_model=GraderResponse)
|
| 373 |
+
def environment_grader() -> GraderResponse:
|
| 374 |
+
try:
|
| 375 |
+
report = env_service.grade()
|
| 376 |
+
return GraderResponse(
|
| 377 |
+
task=env_service.current_task(),
|
| 378 |
+
steps_recorded=env_service.steps_recorded(),
|
| 379 |
+
report=report,
|
| 380 |
+
)
|
| 381 |
+
except RuntimeError as exc:
|
| 382 |
+
raise _http_error(400, str(exc)) from exc
|
| 383 |
+
except Exception as exc:
|
| 384 |
+
raise _http_error(500, "Failed to compute grade") from exc
|
| 385 |
+
|
| 386 |
+
|
| 387 |
+
@router.get("/baseline", response_model=BaselineResponse)
|
| 388 |
+
def environment_baseline(task: str | None = None, seed: int | None = None, max_steps: int | None = None) -> BaselineResponse:
|
| 389 |
+
try:
|
| 390 |
+
return env_service.baseline_rollout(task=task, seed=seed, max_steps=max_steps)
|
| 391 |
+
except ValueError as exc:
|
| 392 |
+
raise _http_error(400, str(exc)) from exc
|
| 393 |
+
except RuntimeError as exc:
|
| 394 |
+
raise _http_error(400, str(exc)) from exc
|
| 395 |
+
except Exception as exc:
|
| 396 |
+
raise _http_error(500, "Failed to compute baseline rollout") from exc
|
app/static/results.html
ADDED
|
@@ -0,0 +1,1063 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
<!doctype html>
|
| 2 |
+
<html lang="en">
|
| 3 |
+
<head>
|
| 4 |
+
<meta charset="UTF-8" />
|
| 5 |
+
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
| 6 |
+
<title>OpenEnv Result Studio</title>
|
| 7 |
+
<link rel="preconnect" href="https://fonts.googleapis.com" />
|
| 8 |
+
<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin />
|
| 9 |
+
<link href="https://fonts.googleapis.com/css2?family=Space+Grotesk:wght@400;500;700&family=JetBrains+Mono:wght@400;600&display=swap" rel="stylesheet" />
|
| 10 |
+
<style>
|
| 11 |
+
:root {
|
| 12 |
+
--ink: #11212a;
|
| 13 |
+
--muted: #4a6675;
|
| 14 |
+
--bg-a: #e7f1f5;
|
| 15 |
+
--bg-b: #f6efe4;
|
| 16 |
+
--card: rgba(255, 255, 255, 0.82);
|
| 17 |
+
--card-border: #c7d9df;
|
| 18 |
+
--accent: #13786d;
|
| 19 |
+
--accent-hover: #0f5f56;
|
| 20 |
+
--accent-soft: #cdece7;
|
| 21 |
+
--warn: #f18f3b;
|
| 22 |
+
--error: #b63d2f;
|
| 23 |
+
--ok: #1f7a40;
|
| 24 |
+
--mono: "JetBrains Mono", "Consolas", monospace;
|
| 25 |
+
}
|
| 26 |
+
|
| 27 |
+
* {
|
| 28 |
+
box-sizing: border-box;
|
| 29 |
+
}
|
| 30 |
+
|
| 31 |
+
body {
|
| 32 |
+
margin: 0;
|
| 33 |
+
min-height: 100vh;
|
| 34 |
+
color: var(--ink);
|
| 35 |
+
font-family: "Space Grotesk", "Segoe UI", sans-serif;
|
| 36 |
+
background:
|
| 37 |
+
radial-gradient(1000px 420px at 8% -5%, #b8dfe8 0%, transparent 72%),
|
| 38 |
+
radial-gradient(920px 420px at 96% -6%, #f6d9bb 0%, transparent 70%),
|
| 39 |
+
linear-gradient(140deg, var(--bg-a), var(--bg-b));
|
| 40 |
+
}
|
| 41 |
+
|
| 42 |
+
body::before {
|
| 43 |
+
content: "";
|
| 44 |
+
position: fixed;
|
| 45 |
+
inset: 0;
|
| 46 |
+
pointer-events: none;
|
| 47 |
+
opacity: 0.25;
|
| 48 |
+
background-image:
|
| 49 |
+
linear-gradient(rgba(17, 33, 42, 0.06) 1px, transparent 1px),
|
| 50 |
+
linear-gradient(90deg, rgba(17, 33, 42, 0.06) 1px, transparent 1px);
|
| 51 |
+
background-size: 32px 32px;
|
| 52 |
+
mask-image: radial-gradient(circle at center, black 38%, transparent 85%);
|
| 53 |
+
}
|
| 54 |
+
|
| 55 |
+
main {
|
| 56 |
+
position: relative;
|
| 57 |
+
max-width: 1160px;
|
| 58 |
+
margin: 0 auto;
|
| 59 |
+
padding: 30px 16px 44px;
|
| 60 |
+
display: grid;
|
| 61 |
+
gap: 14px;
|
| 62 |
+
}
|
| 63 |
+
|
| 64 |
+
.card {
|
| 65 |
+
background: var(--card);
|
| 66 |
+
border: 1px solid var(--card-border);
|
| 67 |
+
border-radius: 18px;
|
| 68 |
+
box-shadow: 0 16px 38px rgba(17, 33, 42, 0.12);
|
| 69 |
+
backdrop-filter: blur(8px);
|
| 70 |
+
animation: rise 0.5s ease both;
|
| 71 |
+
}
|
| 72 |
+
|
| 73 |
+
.card:nth-child(2) { animation-delay: 0.05s; }
|
| 74 |
+
.card:nth-child(3) { animation-delay: 0.09s; }
|
| 75 |
+
.card:nth-child(4) { animation-delay: 0.13s; }
|
| 76 |
+
|
| 77 |
+
.hero {
|
| 78 |
+
padding: 22px;
|
| 79 |
+
display: grid;
|
| 80 |
+
gap: 10px;
|
| 81 |
+
}
|
| 82 |
+
|
| 83 |
+
.hero h1 {
|
| 84 |
+
margin: 0;
|
| 85 |
+
font-size: clamp(1.45rem, 2.5vw, 2.25rem);
|
| 86 |
+
letter-spacing: 0.2px;
|
| 87 |
+
}
|
| 88 |
+
|
| 89 |
+
.hero p {
|
| 90 |
+
margin: 0;
|
| 91 |
+
color: var(--muted);
|
| 92 |
+
line-height: 1.4;
|
| 93 |
+
max-width: 840px;
|
| 94 |
+
}
|
| 95 |
+
|
| 96 |
+
.status-row {
|
| 97 |
+
display: flex;
|
| 98 |
+
align-items: center;
|
| 99 |
+
gap: 10px;
|
| 100 |
+
flex-wrap: wrap;
|
| 101 |
+
}
|
| 102 |
+
|
| 103 |
+
.pill {
|
| 104 |
+
display: inline-flex;
|
| 105 |
+
align-items: center;
|
| 106 |
+
gap: 6px;
|
| 107 |
+
border-radius: 999px;
|
| 108 |
+
padding: 5px 10px;
|
| 109 |
+
font-size: 0.84rem;
|
| 110 |
+
font-weight: 600;
|
| 111 |
+
border: 1px solid transparent;
|
| 112 |
+
}
|
| 113 |
+
|
| 114 |
+
.pill.ok {
|
| 115 |
+
color: #134b23;
|
| 116 |
+
background: #d9f4df;
|
| 117 |
+
border-color: #9ed2a9;
|
| 118 |
+
}
|
| 119 |
+
|
| 120 |
+
.pill.warn {
|
| 121 |
+
color: #78450d;
|
| 122 |
+
background: #ffe6c9;
|
| 123 |
+
border-color: #efc48e;
|
| 124 |
+
}
|
| 125 |
+
|
| 126 |
+
.pill.error {
|
| 127 |
+
color: #6b251d;
|
| 128 |
+
background: #f6d3cf;
|
| 129 |
+
border-color: #e2a19a;
|
| 130 |
+
}
|
| 131 |
+
|
| 132 |
+
.controls {
|
| 133 |
+
padding: 18px;
|
| 134 |
+
display: grid;
|
| 135 |
+
gap: 14px;
|
| 136 |
+
}
|
| 137 |
+
|
| 138 |
+
.form-row {
|
| 139 |
+
display: grid;
|
| 140 |
+
grid-template-columns: repeat(2, minmax(0, 1fr));
|
| 141 |
+
gap: 10px;
|
| 142 |
+
}
|
| 143 |
+
|
| 144 |
+
.field {
|
| 145 |
+
display: grid;
|
| 146 |
+
gap: 5px;
|
| 147 |
+
}
|
| 148 |
+
|
| 149 |
+
.field label {
|
| 150 |
+
font-size: 0.72rem;
|
| 151 |
+
letter-spacing: 0.08em;
|
| 152 |
+
color: var(--muted);
|
| 153 |
+
text-transform: uppercase;
|
| 154 |
+
font-weight: 700;
|
| 155 |
+
}
|
| 156 |
+
|
| 157 |
+
.field input,
|
| 158 |
+
.field select {
|
| 159 |
+
width: 100%;
|
| 160 |
+
border-radius: 10px;
|
| 161 |
+
border: 1px solid #aac2ca;
|
| 162 |
+
background: #f7fbfc;
|
| 163 |
+
color: var(--ink);
|
| 164 |
+
padding: 10px 11px;
|
| 165 |
+
font-size: 0.95rem;
|
| 166 |
+
font-family: inherit;
|
| 167 |
+
outline: none;
|
| 168 |
+
transition: border-color 0.18s ease, box-shadow 0.18s ease;
|
| 169 |
+
}
|
| 170 |
+
|
| 171 |
+
.field input:focus,
|
| 172 |
+
.field select:focus {
|
| 173 |
+
border-color: var(--accent);
|
| 174 |
+
box-shadow: 0 0 0 3px rgba(19, 120, 109, 0.15);
|
| 175 |
+
}
|
| 176 |
+
|
| 177 |
+
.button-row {
|
| 178 |
+
display: flex;
|
| 179 |
+
flex-wrap: wrap;
|
| 180 |
+
gap: 8px;
|
| 181 |
+
}
|
| 182 |
+
|
| 183 |
+
button {
|
| 184 |
+
border: 1px solid #9fb8c0;
|
| 185 |
+
border-radius: 10px;
|
| 186 |
+
padding: 10px 13px;
|
| 187 |
+
font-family: inherit;
|
| 188 |
+
font-size: 0.9rem;
|
| 189 |
+
font-weight: 600;
|
| 190 |
+
color: var(--ink);
|
| 191 |
+
background: #f4fbfd;
|
| 192 |
+
cursor: pointer;
|
| 193 |
+
transition: transform 0.12s ease, box-shadow 0.12s ease, background 0.16s ease;
|
| 194 |
+
}
|
| 195 |
+
|
| 196 |
+
button:hover {
|
| 197 |
+
transform: translateY(-1px);
|
| 198 |
+
box-shadow: 0 6px 16px rgba(17, 33, 42, 0.12);
|
| 199 |
+
background: #ebf7fa;
|
| 200 |
+
}
|
| 201 |
+
|
| 202 |
+
button.primary {
|
| 203 |
+
background: var(--accent);
|
| 204 |
+
color: #ffffff;
|
| 205 |
+
border-color: #0f5f56;
|
| 206 |
+
}
|
| 207 |
+
|
| 208 |
+
button.primary:hover {
|
| 209 |
+
background: var(--accent-hover);
|
| 210 |
+
}
|
| 211 |
+
|
| 212 |
+
button.warn {
|
| 213 |
+
background: #fff4e8;
|
| 214 |
+
border-color: #f3bc86;
|
| 215 |
+
color: #7e4a0f;
|
| 216 |
+
}
|
| 217 |
+
|
| 218 |
+
.metrics {
|
| 219 |
+
padding: 18px;
|
| 220 |
+
display: grid;
|
| 221 |
+
grid-template-columns: repeat(6, minmax(0, 1fr));
|
| 222 |
+
gap: 10px;
|
| 223 |
+
}
|
| 224 |
+
|
| 225 |
+
.metric {
|
| 226 |
+
background: rgba(255, 255, 255, 0.82);
|
| 227 |
+
border: 1px solid #cddbe0;
|
| 228 |
+
border-radius: 12px;
|
| 229 |
+
padding: 12px;
|
| 230 |
+
display: grid;
|
| 231 |
+
gap: 5px;
|
| 232 |
+
}
|
| 233 |
+
|
| 234 |
+
.metric label {
|
| 235 |
+
font-size: 0.72rem;
|
| 236 |
+
letter-spacing: 0.08em;
|
| 237 |
+
text-transform: uppercase;
|
| 238 |
+
color: var(--muted);
|
| 239 |
+
font-weight: 700;
|
| 240 |
+
}
|
| 241 |
+
|
| 242 |
+
.metric strong {
|
| 243 |
+
font-family: var(--mono);
|
| 244 |
+
font-size: 1rem;
|
| 245 |
+
color: #0c3640;
|
| 246 |
+
word-break: break-word;
|
| 247 |
+
}
|
| 248 |
+
|
| 249 |
+
.results-grid {
|
| 250 |
+
display: grid;
|
| 251 |
+
grid-template-columns: 1.1fr 0.9fr;
|
| 252 |
+
gap: 14px;
|
| 253 |
+
}
|
| 254 |
+
|
| 255 |
+
.panel {
|
| 256 |
+
padding: 16px;
|
| 257 |
+
display: grid;
|
| 258 |
+
gap: 10px;
|
| 259 |
+
}
|
| 260 |
+
|
| 261 |
+
.panel h2 {
|
| 262 |
+
margin: 0;
|
| 263 |
+
font-size: 1.05rem;
|
| 264 |
+
letter-spacing: 0.02em;
|
| 265 |
+
}
|
| 266 |
+
|
| 267 |
+
.mono {
|
| 268 |
+
font-family: var(--mono);
|
| 269 |
+
font-size: 0.82rem;
|
| 270 |
+
line-height: 1.45;
|
| 271 |
+
}
|
| 272 |
+
|
| 273 |
+
.log-list {
|
| 274 |
+
list-style: none;
|
| 275 |
+
margin: 0;
|
| 276 |
+
padding: 0;
|
| 277 |
+
display: grid;
|
| 278 |
+
gap: 8px;
|
| 279 |
+
max-height: 390px;
|
| 280 |
+
overflow: auto;
|
| 281 |
+
}
|
| 282 |
+
|
| 283 |
+
.log-item {
|
| 284 |
+
border: 1px solid #c8d7dc;
|
| 285 |
+
background: rgba(248, 252, 253, 0.95);
|
| 286 |
+
border-radius: 12px;
|
| 287 |
+
padding: 10px;
|
| 288 |
+
display: grid;
|
| 289 |
+
gap: 6px;
|
| 290 |
+
}
|
| 291 |
+
|
| 292 |
+
.log-head {
|
| 293 |
+
display: flex;
|
| 294 |
+
justify-content: space-between;
|
| 295 |
+
gap: 8px;
|
| 296 |
+
align-items: center;
|
| 297 |
+
font-size: 0.86rem;
|
| 298 |
+
font-weight: 700;
|
| 299 |
+
}
|
| 300 |
+
|
| 301 |
+
.tag {
|
| 302 |
+
font-size: 0.72rem;
|
| 303 |
+
padding: 2px 7px;
|
| 304 |
+
border-radius: 999px;
|
| 305 |
+
border: 1px solid;
|
| 306 |
+
font-weight: 700;
|
| 307 |
+
letter-spacing: 0.04em;
|
| 308 |
+
}
|
| 309 |
+
|
| 310 |
+
.tag.ok {
|
| 311 |
+
color: #205f32;
|
| 312 |
+
border-color: #8fc39d;
|
| 313 |
+
background: #e1f5e6;
|
| 314 |
+
}
|
| 315 |
+
|
| 316 |
+
.tag.error {
|
| 317 |
+
color: #75281f;
|
| 318 |
+
border-color: #d2948b;
|
| 319 |
+
background: #f9dcd8;
|
| 320 |
+
}
|
| 321 |
+
|
| 322 |
+
pre {
|
| 323 |
+
margin: 0;
|
| 324 |
+
padding: 12px;
|
| 325 |
+
border: 1px solid #bfd3d9;
|
| 326 |
+
border-radius: 12px;
|
| 327 |
+
background: #f4fbfd;
|
| 328 |
+
max-height: 390px;
|
| 329 |
+
overflow: auto;
|
| 330 |
+
white-space: pre-wrap;
|
| 331 |
+
word-break: break-word;
|
| 332 |
+
}
|
| 333 |
+
|
| 334 |
+
.grader-table {
|
| 335 |
+
width: 100%;
|
| 336 |
+
border-collapse: collapse;
|
| 337 |
+
font-family: var(--mono);
|
| 338 |
+
font-size: 0.8rem;
|
| 339 |
+
border: 1px solid #bfd3d9;
|
| 340 |
+
border-radius: 12px;
|
| 341 |
+
overflow: hidden;
|
| 342 |
+
}
|
| 343 |
+
|
| 344 |
+
.grader-table th,
|
| 345 |
+
.grader-table td {
|
| 346 |
+
text-align: left;
|
| 347 |
+
padding: 8px 10px;
|
| 348 |
+
border-bottom: 1px solid #d7e4e8;
|
| 349 |
+
}
|
| 350 |
+
|
| 351 |
+
.grader-table th {
|
| 352 |
+
background: #e4f2f5;
|
| 353 |
+
letter-spacing: 0.06em;
|
| 354 |
+
text-transform: uppercase;
|
| 355 |
+
font-size: 0.72rem;
|
| 356 |
+
color: #365765;
|
| 357 |
+
}
|
| 358 |
+
|
| 359 |
+
.grader-table tr:last-child td {
|
| 360 |
+
border-bottom: none;
|
| 361 |
+
}
|
| 362 |
+
|
| 363 |
+
.suite-wrap {
|
| 364 |
+
padding: 16px;
|
| 365 |
+
display: grid;
|
| 366 |
+
gap: 10px;
|
| 367 |
+
}
|
| 368 |
+
|
| 369 |
+
.suite-toolbar {
|
| 370 |
+
display: flex;
|
| 371 |
+
justify-content: space-between;
|
| 372 |
+
align-items: center;
|
| 373 |
+
gap: 8px;
|
| 374 |
+
flex-wrap: wrap;
|
| 375 |
+
}
|
| 376 |
+
|
| 377 |
+
.suite-note {
|
| 378 |
+
margin: 0;
|
| 379 |
+
color: var(--muted);
|
| 380 |
+
font-size: 0.85rem;
|
| 381 |
+
}
|
| 382 |
+
|
| 383 |
+
.suite-average {
|
| 384 |
+
font-family: var(--mono);
|
| 385 |
+
font-weight: 700;
|
| 386 |
+
color: #0c3640;
|
| 387 |
+
font-size: 0.88rem;
|
| 388 |
+
}
|
| 389 |
+
|
| 390 |
+
.suite-table {
|
| 391 |
+
width: 100%;
|
| 392 |
+
border-collapse: collapse;
|
| 393 |
+
border: 1px solid #bfd3d9;
|
| 394 |
+
border-radius: 12px;
|
| 395 |
+
overflow: hidden;
|
| 396 |
+
font-family: var(--mono);
|
| 397 |
+
font-size: 0.8rem;
|
| 398 |
+
}
|
| 399 |
+
|
| 400 |
+
.suite-table th,
|
| 401 |
+
.suite-table td {
|
| 402 |
+
text-align: left;
|
| 403 |
+
padding: 8px 10px;
|
| 404 |
+
border-bottom: 1px solid #d7e4e8;
|
| 405 |
+
}
|
| 406 |
+
|
| 407 |
+
.suite-table th {
|
| 408 |
+
background: #e4f2f5;
|
| 409 |
+
letter-spacing: 0.06em;
|
| 410 |
+
text-transform: uppercase;
|
| 411 |
+
font-size: 0.72rem;
|
| 412 |
+
color: #365765;
|
| 413 |
+
}
|
| 414 |
+
|
| 415 |
+
.suite-table tr:last-child td {
|
| 416 |
+
border-bottom: none;
|
| 417 |
+
}
|
| 418 |
+
|
| 419 |
+
.flag {
|
| 420 |
+
display: inline-flex;
|
| 421 |
+
align-items: center;
|
| 422 |
+
border-radius: 999px;
|
| 423 |
+
border: 1px solid;
|
| 424 |
+
padding: 2px 7px;
|
| 425 |
+
font-size: 0.72rem;
|
| 426 |
+
letter-spacing: 0.04em;
|
| 427 |
+
font-weight: 700;
|
| 428 |
+
}
|
| 429 |
+
|
| 430 |
+
.flag.ok {
|
| 431 |
+
color: #205f32;
|
| 432 |
+
border-color: #8fc39d;
|
| 433 |
+
background: #e1f5e6;
|
| 434 |
+
}
|
| 435 |
+
|
| 436 |
+
.flag.error {
|
| 437 |
+
color: #75281f;
|
| 438 |
+
border-color: #d2948b;
|
| 439 |
+
background: #f9dcd8;
|
| 440 |
+
}
|
| 441 |
+
|
| 442 |
+
@media (max-width: 980px) {
|
| 443 |
+
.metrics {
|
| 444 |
+
grid-template-columns: repeat(3, minmax(0, 1fr));
|
| 445 |
+
}
|
| 446 |
+
|
| 447 |
+
.results-grid {
|
| 448 |
+
grid-template-columns: 1fr;
|
| 449 |
+
}
|
| 450 |
+
}
|
| 451 |
+
|
| 452 |
+
@media (max-width: 640px) {
|
| 453 |
+
.form-row {
|
| 454 |
+
grid-template-columns: 1fr;
|
| 455 |
+
}
|
| 456 |
+
|
| 457 |
+
.metrics {
|
| 458 |
+
grid-template-columns: repeat(2, minmax(0, 1fr));
|
| 459 |
+
}
|
| 460 |
+
|
| 461 |
+
.button-row {
|
| 462 |
+
display: grid;
|
| 463 |
+
grid-template-columns: 1fr 1fr;
|
| 464 |
+
}
|
| 465 |
+
}
|
| 466 |
+
|
| 467 |
+
@keyframes rise {
|
| 468 |
+
from {
|
| 469 |
+
opacity: 0;
|
| 470 |
+
transform: translateY(8px);
|
| 471 |
+
}
|
| 472 |
+
to {
|
| 473 |
+
opacity: 1;
|
| 474 |
+
transform: translateY(0);
|
| 475 |
+
}
|
| 476 |
+
}
|
| 477 |
+
</style>
|
| 478 |
+
</head>
|
| 479 |
+
<body>
|
| 480 |
+
<main>
|
| 481 |
+
<section class="card hero">
|
| 482 |
+
<h1>OpenEnv Result Studio</h1>
|
| 483 |
+
<p>Interactive frontend for the last-mile delivery benchmark. Run reset/step/state/grader/baseline directly from the browser and inspect live deterministic results.</p>
|
| 484 |
+
<div class="status-row">
|
| 485 |
+
<span id="statusPill" class="pill warn">Idle</span>
|
| 486 |
+
<span id="statusText" class="mono">Ready to run checks.</span>
|
| 487 |
+
</div>
|
| 488 |
+
</section>
|
| 489 |
+
|
| 490 |
+
<section class="card controls">
|
| 491 |
+
<div class="form-row">
|
| 492 |
+
<div class="field">
|
| 493 |
+
<label for="taskInput">Task</label>
|
| 494 |
+
<select id="taskInput">
|
| 495 |
+
<option value="easy">easy</option>
|
| 496 |
+
<option value="medium">medium</option>
|
| 497 |
+
<option value="hard">hard</option>
|
| 498 |
+
</select>
|
| 499 |
+
</div>
|
| 500 |
+
<div class="field">
|
| 501 |
+
<label for="seedInput">Seed</label>
|
| 502 |
+
<input id="seedInput" type="number" min="0" value="101" />
|
| 503 |
+
</div>
|
| 504 |
+
</div>
|
| 505 |
+
|
| 506 |
+
<div class="form-row">
|
| 507 |
+
<div class="field">
|
| 508 |
+
<label for="apiBaseInput">API Base URL (optional)</label>
|
| 509 |
+
<input id="apiBaseInput" type="text" placeholder="http://127.0.0.1:8000" />
|
| 510 |
+
</div>
|
| 511 |
+
</div>
|
| 512 |
+
|
| 513 |
+
<div class="button-row">
|
| 514 |
+
<button class="primary" id="btnReset">Reset</button>
|
| 515 |
+
<button id="btnStep">Step (wait)</button>
|
| 516 |
+
<button id="btnState">State</button>
|
| 517 |
+
<button id="btnGrader">Grader</button>
|
| 518 |
+
<button id="btnBaseline">Baseline</button>
|
| 519 |
+
<button class="primary" id="btnSuite">Run All Tasks</button>
|
| 520 |
+
<button class="warn" id="btnDemo">Run Demo Sequence</button>
|
| 521 |
+
</div>
|
| 522 |
+
</section>
|
| 523 |
+
|
| 524 |
+
<section class="card metrics">
|
| 525 |
+
<div class="metric"><label>Task</label><strong id="mTask">easy</strong></div>
|
| 526 |
+
<div class="metric"><label>Done</label><strong id="mDone">false</strong></div>
|
| 527 |
+
<div class="metric"><label>Step Count</label><strong id="mStep">0</strong></div>
|
| 528 |
+
<div class="metric"><label>Total Reward</label><strong id="mTotalReward">0.00</strong></div>
|
| 529 |
+
<div class="metric"><label>Pending Orders</label><strong id="mPending">0</strong></div>
|
| 530 |
+
<div class="metric"><label>Score</label><strong id="mScore">n/a</strong></div>
|
| 531 |
+
</section>
|
| 532 |
+
|
| 533 |
+
<section class="results-grid">
|
| 534 |
+
<article class="card panel">
|
| 535 |
+
<h2>Run Log</h2>
|
| 536 |
+
<ul id="logList" class="log-list"></ul>
|
| 537 |
+
</article>
|
| 538 |
+
|
| 539 |
+
<article class="card panel">
|
| 540 |
+
<h2>Latest Payload</h2>
|
| 541 |
+
<pre id="payloadView" class="mono">{}</pre>
|
| 542 |
+
</article>
|
| 543 |
+
</section>
|
| 544 |
+
|
| 545 |
+
<section class="card panel">
|
| 546 |
+
<h2>Grader Metrics</h2>
|
| 547 |
+
<table class="grader-table" id="graderTable">
|
| 548 |
+
<thead>
|
| 549 |
+
<tr><th>Metric</th><th>Value</th></tr>
|
| 550 |
+
</thead>
|
| 551 |
+
<tbody>
|
| 552 |
+
<tr><td>score</td><td>n/a</td></tr>
|
| 553 |
+
</tbody>
|
| 554 |
+
</table>
|
| 555 |
+
</section>
|
| 556 |
+
|
| 557 |
+
<section class="card suite-wrap">
|
| 558 |
+
<div class="suite-toolbar">
|
| 559 |
+
<p class="suite-note">One-click baseline rollout across easy, medium, and hard using the current seed.</p>
|
| 560 |
+
<span id="suiteAverage" class="suite-average">average_score: n/a</span>
|
| 561 |
+
</div>
|
| 562 |
+
<table class="suite-table" id="suiteTable">
|
| 563 |
+
<thead>
|
| 564 |
+
<tr><th>Task</th><th>Success</th><th>Score</th><th>Steps</th><th>Total Reward</th></tr>
|
| 565 |
+
</thead>
|
| 566 |
+
<tbody>
|
| 567 |
+
<tr><td colspan="5">No suite run yet.</td></tr>
|
| 568 |
+
</tbody>
|
| 569 |
+
</table>
|
| 570 |
+
</section>
|
| 571 |
+
</main>
|
| 572 |
+
|
| 573 |
+
<script>
|
| 574 |
+
const state = {
|
| 575 |
+
task: "easy",
|
| 576 |
+
seed: 101,
|
| 577 |
+
apiBaseUrl: "",
|
| 578 |
+
done: false,
|
| 579 |
+
stepCount: 0,
|
| 580 |
+
totalReward: 0,
|
| 581 |
+
pendingOrders: 0,
|
| 582 |
+
score: null,
|
| 583 |
+
lastPayload: {},
|
| 584 |
+
logCount: 0,
|
| 585 |
+
suiteRows: [],
|
| 586 |
+
};
|
| 587 |
+
|
| 588 |
+
const nodes = {
|
| 589 |
+
statusPill: document.getElementById("statusPill"),
|
| 590 |
+
statusText: document.getElementById("statusText"),
|
| 591 |
+
taskInput: document.getElementById("taskInput"),
|
| 592 |
+
seedInput: document.getElementById("seedInput"),
|
| 593 |
+
apiBaseInput: document.getElementById("apiBaseInput"),
|
| 594 |
+
btnReset: document.getElementById("btnReset"),
|
| 595 |
+
btnStep: document.getElementById("btnStep"),
|
| 596 |
+
btnState: document.getElementById("btnState"),
|
| 597 |
+
btnGrader: document.getElementById("btnGrader"),
|
| 598 |
+
btnBaseline: document.getElementById("btnBaseline"),
|
| 599 |
+
btnSuite: document.getElementById("btnSuite"),
|
| 600 |
+
btnDemo: document.getElementById("btnDemo"),
|
| 601 |
+
mTask: document.getElementById("mTask"),
|
| 602 |
+
mDone: document.getElementById("mDone"),
|
| 603 |
+
mStep: document.getElementById("mStep"),
|
| 604 |
+
mTotalReward: document.getElementById("mTotalReward"),
|
| 605 |
+
mPending: document.getElementById("mPending"),
|
| 606 |
+
mScore: document.getElementById("mScore"),
|
| 607 |
+
logList: document.getElementById("logList"),
|
| 608 |
+
payloadView: document.getElementById("payloadView"),
|
| 609 |
+
graderTable: document.getElementById("graderTable"),
|
| 610 |
+
suiteTable: document.getElementById("suiteTable"),
|
| 611 |
+
suiteAverage: document.getElementById("suiteAverage"),
|
| 612 |
+
};
|
| 613 |
+
|
| 614 |
+
function setStatus(message, mode) {
|
| 615 |
+
nodes.statusText.textContent = message;
|
| 616 |
+
nodes.statusPill.classList.remove("ok", "warn", "error");
|
| 617 |
+
if (mode === "ok") {
|
| 618 |
+
nodes.statusPill.classList.add("ok");
|
| 619 |
+
nodes.statusPill.textContent = "OK";
|
| 620 |
+
} else if (mode === "error") {
|
| 621 |
+
nodes.statusPill.classList.add("error");
|
| 622 |
+
nodes.statusPill.textContent = "Error";
|
| 623 |
+
} else {
|
| 624 |
+
nodes.statusPill.classList.add("warn");
|
| 625 |
+
nodes.statusPill.textContent = "Running";
|
| 626 |
+
}
|
| 627 |
+
}
|
| 628 |
+
|
| 629 |
+
function renderMetrics() {
|
| 630 |
+
nodes.mTask.textContent = state.task;
|
| 631 |
+
nodes.mDone.textContent = String(state.done);
|
| 632 |
+
nodes.mStep.textContent = String(state.stepCount);
|
| 633 |
+
nodes.mTotalReward.textContent = Number(state.totalReward).toFixed(2);
|
| 634 |
+
nodes.mPending.textContent = String(state.pendingOrders);
|
| 635 |
+
nodes.mScore.textContent = state.score === null ? "n/a" : Number(state.score).toFixed(4);
|
| 636 |
+
nodes.payloadView.textContent = JSON.stringify(state.lastPayload, null, 2);
|
| 637 |
+
}
|
| 638 |
+
|
| 639 |
+
function renderGraderTable(report) {
|
| 640 |
+
const tbody = document.createElement("tbody");
|
| 641 |
+
const pairs = [
|
| 642 |
+
["score", report?.score],
|
| 643 |
+
["completion_rate", report?.completion_rate],
|
| 644 |
+
["efficiency_ratio", report?.efficiency_ratio],
|
| 645 |
+
["invalid_action_rate", report?.invalid_action_rate],
|
| 646 |
+
["steps_taken", report?.steps_taken],
|
| 647 |
+
["delivered_orders", report?.delivered_orders],
|
| 648 |
+
["total_orders", report?.total_orders],
|
| 649 |
+
];
|
| 650 |
+
|
| 651 |
+
pairs.forEach(([key, value]) => {
|
| 652 |
+
const tr = document.createElement("tr");
|
| 653 |
+
const tdKey = document.createElement("td");
|
| 654 |
+
tdKey.textContent = String(key);
|
| 655 |
+
const tdValue = document.createElement("td");
|
| 656 |
+
tdValue.textContent = value === undefined || value === null ? "n/a" : String(value);
|
| 657 |
+
tr.appendChild(tdKey);
|
| 658 |
+
tr.appendChild(tdValue);
|
| 659 |
+
tbody.appendChild(tr);
|
| 660 |
+
});
|
| 661 |
+
|
| 662 |
+
const oldTbody = nodes.graderTable.querySelector("tbody");
|
| 663 |
+
if (oldTbody) {
|
| 664 |
+
nodes.graderTable.removeChild(oldTbody);
|
| 665 |
+
}
|
| 666 |
+
nodes.graderTable.appendChild(tbody);
|
| 667 |
+
}
|
| 668 |
+
|
| 669 |
+
function renderSuiteTable(rows) {
|
| 670 |
+
const tbody = document.createElement("tbody");
|
| 671 |
+
if (!Array.isArray(rows) || rows.length === 0) {
|
| 672 |
+
const tr = document.createElement("tr");
|
| 673 |
+
const td = document.createElement("td");
|
| 674 |
+
td.colSpan = 5;
|
| 675 |
+
td.textContent = "No suite run yet.";
|
| 676 |
+
tr.appendChild(td);
|
| 677 |
+
tbody.appendChild(tr);
|
| 678 |
+
} else {
|
| 679 |
+
rows.forEach((item) => {
|
| 680 |
+
const tr = document.createElement("tr");
|
| 681 |
+
|
| 682 |
+
const taskTd = document.createElement("td");
|
| 683 |
+
taskTd.textContent = String(item.task);
|
| 684 |
+
tr.appendChild(taskTd);
|
| 685 |
+
|
| 686 |
+
const successTd = document.createElement("td");
|
| 687 |
+
const flag = document.createElement("span");
|
| 688 |
+
flag.className = `flag ${item.success ? "ok" : "error"}`;
|
| 689 |
+
flag.textContent = item.success ? "TRUE" : "FALSE";
|
| 690 |
+
successTd.appendChild(flag);
|
| 691 |
+
tr.appendChild(successTd);
|
| 692 |
+
|
| 693 |
+
const scoreTd = document.createElement("td");
|
| 694 |
+
scoreTd.textContent = Number(item.score || 0).toFixed(4);
|
| 695 |
+
tr.appendChild(scoreTd);
|
| 696 |
+
|
| 697 |
+
const stepsTd = document.createElement("td");
|
| 698 |
+
stepsTd.textContent = String(item.steps_executed || 0);
|
| 699 |
+
tr.appendChild(stepsTd);
|
| 700 |
+
|
| 701 |
+
const rewardTd = document.createElement("td");
|
| 702 |
+
rewardTd.textContent = Number(item.total_reward || 0).toFixed(2);
|
| 703 |
+
tr.appendChild(rewardTd);
|
| 704 |
+
|
| 705 |
+
tbody.appendChild(tr);
|
| 706 |
+
});
|
| 707 |
+
}
|
| 708 |
+
|
| 709 |
+
const oldTbody = nodes.suiteTable.querySelector("tbody");
|
| 710 |
+
if (oldTbody) {
|
| 711 |
+
nodes.suiteTable.removeChild(oldTbody);
|
| 712 |
+
}
|
| 713 |
+
nodes.suiteTable.appendChild(tbody);
|
| 714 |
+
|
| 715 |
+
if (!rows.length) {
|
| 716 |
+
nodes.suiteAverage.textContent = "average_score: n/a";
|
| 717 |
+
return;
|
| 718 |
+
}
|
| 719 |
+
|
| 720 |
+
const average = rows.reduce((acc, item) => acc + Number(item.score || 0), 0) / rows.length;
|
| 721 |
+
nodes.suiteAverage.textContent = `average_score: ${average.toFixed(4)}`;
|
| 722 |
+
}
|
| 723 |
+
|
| 724 |
+
function addLog(label, payload, ok) {
|
| 725 |
+
state.logCount += 1;
|
| 726 |
+
const li = document.createElement("li");
|
| 727 |
+
li.className = "log-item";
|
| 728 |
+
|
| 729 |
+
const header = document.createElement("div");
|
| 730 |
+
header.className = "log-head";
|
| 731 |
+
const title = document.createElement("span");
|
| 732 |
+
title.textContent = `${state.logCount}. ${label}`;
|
| 733 |
+
|
| 734 |
+
const tag = document.createElement("span");
|
| 735 |
+
tag.className = `tag ${ok ? "ok" : "error"}`;
|
| 736 |
+
tag.textContent = ok ? "SUCCESS" : "FAILED";
|
| 737 |
+
|
| 738 |
+
header.appendChild(title);
|
| 739 |
+
header.appendChild(tag);
|
| 740 |
+
li.appendChild(header);
|
| 741 |
+
|
| 742 |
+
const pre = document.createElement("pre");
|
| 743 |
+
pre.className = "mono";
|
| 744 |
+
pre.textContent = JSON.stringify(payload, null, 2);
|
| 745 |
+
li.appendChild(pre);
|
| 746 |
+
|
| 747 |
+
nodes.logList.prepend(li);
|
| 748 |
+
}
|
| 749 |
+
|
| 750 |
+
function updateFromObservation(obs, doneFlag) {
|
| 751 |
+
if (!obs) {
|
| 752 |
+
return;
|
| 753 |
+
}
|
| 754 |
+
state.done = Boolean(doneFlag);
|
| 755 |
+
state.stepCount = Number(obs.step_count || 0);
|
| 756 |
+
state.totalReward = Number(obs.total_reward || 0);
|
| 757 |
+
state.pendingOrders = Array.isArray(obs.pending_orders) ? obs.pending_orders.length : 0;
|
| 758 |
+
}
|
| 759 |
+
|
| 760 |
+
function updateFromState(payload) {
|
| 761 |
+
state.done = Boolean(payload.done);
|
| 762 |
+
state.stepCount = Number(payload.step_count || 0);
|
| 763 |
+
state.totalReward = Number(payload.total_reward || 0);
|
| 764 |
+
if (payload.observation && Array.isArray(payload.observation.pending_orders)) {
|
| 765 |
+
state.pendingOrders = payload.observation.pending_orders.length;
|
| 766 |
+
}
|
| 767 |
+
}
|
| 768 |
+
|
| 769 |
+
function normalizeApiBaseUrl(value) {
|
| 770 |
+
return String(value || "").trim().replace(/\/+$/, "");
|
| 771 |
+
}
|
| 772 |
+
|
| 773 |
+
function applyApiBaseUrl(value, persist) {
|
| 774 |
+
state.apiBaseUrl = normalizeApiBaseUrl(value);
|
| 775 |
+
if (nodes.apiBaseInput) {
|
| 776 |
+
nodes.apiBaseInput.value = state.apiBaseUrl;
|
| 777 |
+
}
|
| 778 |
+
|
| 779 |
+
if (!persist) {
|
| 780 |
+
return;
|
| 781 |
+
}
|
| 782 |
+
|
| 783 |
+
try {
|
| 784 |
+
window.localStorage.setItem("openenv_api_base_url", state.apiBaseUrl);
|
| 785 |
+
} catch (_) {
|
| 786 |
+
// Ignore storage errors in restricted browser contexts.
|
| 787 |
+
}
|
| 788 |
+
}
|
| 789 |
+
|
| 790 |
+
function currentApiBaseUrl() {
|
| 791 |
+
return normalizeApiBaseUrl(nodes.apiBaseInput?.value || state.apiBaseUrl);
|
| 792 |
+
}
|
| 793 |
+
|
| 794 |
+
function buildApiUrl(path) {
|
| 795 |
+
const normalizedPath = path.startsWith("/") ? path : `/${path}`;
|
| 796 |
+
const baseUrl = currentApiBaseUrl();
|
| 797 |
+
if (!baseUrl) {
|
| 798 |
+
return normalizedPath;
|
| 799 |
+
}
|
| 800 |
+
return `${baseUrl}${normalizedPath}`;
|
| 801 |
+
}
|
| 802 |
+
|
| 803 |
+
function saveApiBaseUrl() {
|
| 804 |
+
applyApiBaseUrl(currentApiBaseUrl(), true);
|
| 805 |
+
}
|
| 806 |
+
|
| 807 |
+
function loadApiBaseUrl() {
|
| 808 |
+
let saved = "";
|
| 809 |
+
try {
|
| 810 |
+
saved = window.localStorage.getItem("openenv_api_base_url") || "";
|
| 811 |
+
} catch (_) {
|
| 812 |
+
saved = "";
|
| 813 |
+
}
|
| 814 |
+
|
| 815 |
+
applyApiBaseUrl(saved, false);
|
| 816 |
+
}
|
| 817 |
+
|
| 818 |
+
function candidateApiBaseUrls() {
|
| 819 |
+
const candidates = [];
|
| 820 |
+
const add = (value) => {
|
| 821 |
+
const normalized = normalizeApiBaseUrl(value);
|
| 822 |
+
if (normalized && !candidates.includes(normalized)) {
|
| 823 |
+
candidates.push(normalized);
|
| 824 |
+
}
|
| 825 |
+
};
|
| 826 |
+
|
| 827 |
+
add(state.apiBaseUrl);
|
| 828 |
+
if (window.location.protocol === "http:" || window.location.protocol === "https:") {
|
| 829 |
+
add(window.location.origin);
|
| 830 |
+
}
|
| 831 |
+
|
| 832 |
+
add("http://127.0.0.1:8000");
|
| 833 |
+
add("http://localhost:8000");
|
| 834 |
+
add("http://127.0.0.1:7860");
|
| 835 |
+
add("http://localhost:7860");
|
| 836 |
+
return candidates;
|
| 837 |
+
}
|
| 838 |
+
|
| 839 |
+
async function probeHealth(baseUrl) {
|
| 840 |
+
const url = `${baseUrl}/health`;
|
| 841 |
+
try {
|
| 842 |
+
const response = await fetch(url, { method: "GET" });
|
| 843 |
+
if (!response.ok) {
|
| 844 |
+
return false;
|
| 845 |
+
}
|
| 846 |
+
|
| 847 |
+
const contentType = (response.headers.get("content-type") || "").toLowerCase();
|
| 848 |
+
if (!contentType.includes("application/json")) {
|
| 849 |
+
return false;
|
| 850 |
+
}
|
| 851 |
+
|
| 852 |
+
const payload = await response.json();
|
| 853 |
+
return payload && payload.status === "ok";
|
| 854 |
+
} catch (_) {
|
| 855 |
+
return false;
|
| 856 |
+
}
|
| 857 |
+
}
|
| 858 |
+
|
| 859 |
+
async function autoDetectApiBaseUrl() {
|
| 860 |
+
const candidates = candidateApiBaseUrls();
|
| 861 |
+
for (let i = 0; i < candidates.length; i += 1) {
|
| 862 |
+
const candidate = candidates[i];
|
| 863 |
+
if (await probeHealth(candidate)) {
|
| 864 |
+
applyApiBaseUrl(candidate, true);
|
| 865 |
+
return true;
|
| 866 |
+
}
|
| 867 |
+
}
|
| 868 |
+
return false;
|
| 869 |
+
}
|
| 870 |
+
|
| 871 |
+
async function apiCall(path, options = {}) {
|
| 872 |
+
const url = buildApiUrl(path);
|
| 873 |
+
const response = await fetch(url, {
|
| 874 |
+
headers: { "Content-Type": "application/json" },
|
| 875 |
+
...options,
|
| 876 |
+
});
|
| 877 |
+
|
| 878 |
+
const rawText = await response.text();
|
| 879 |
+
let data = {};
|
| 880 |
+
if (rawText) {
|
| 881 |
+
try {
|
| 882 |
+
data = JSON.parse(rawText);
|
| 883 |
+
} catch (_) {
|
| 884 |
+
data = {
|
| 885 |
+
message: "Non-JSON response from endpoint",
|
| 886 |
+
endpoint: url,
|
| 887 |
+
status: response.status,
|
| 888 |
+
content_type: response.headers.get("content-type") || "unknown",
|
| 889 |
+
body_preview: rawText.slice(0, 240),
|
| 890 |
+
hint: "Set API Base URL to your FastAPI server, for example http://127.0.0.1:8000",
|
| 891 |
+
};
|
| 892 |
+
}
|
| 893 |
+
}
|
| 894 |
+
|
| 895 |
+
if (!response.ok) {
|
| 896 |
+
const err = new Error(`HTTP ${response.status}`);
|
| 897 |
+
err.payload = data;
|
| 898 |
+
throw err;
|
| 899 |
+
}
|
| 900 |
+
|
| 901 |
+
if (data && data.message === "Non-JSON response from endpoint") {
|
| 902 |
+
const err = new Error("API returned non-JSON content");
|
| 903 |
+
err.payload = data;
|
| 904 |
+
throw err;
|
| 905 |
+
}
|
| 906 |
+
|
| 907 |
+
return data;
|
| 908 |
+
}
|
| 909 |
+
|
| 910 |
+
async function withAction(label, fn) {
|
| 911 |
+
setStatus(`${label} in progress...`, "warn");
|
| 912 |
+
try {
|
| 913 |
+
const payload = await fn();
|
| 914 |
+
state.lastPayload = payload;
|
| 915 |
+
addLog(label, payload, true);
|
| 916 |
+
renderMetrics();
|
| 917 |
+
setStatus(`${label} complete.`, "ok");
|
| 918 |
+
} catch (error) {
|
| 919 |
+
const payload = error.payload || { message: String(error) };
|
| 920 |
+
state.lastPayload = payload;
|
| 921 |
+
addLog(label, payload, false);
|
| 922 |
+
renderMetrics();
|
| 923 |
+
setStatus(`${label} failed.`, "error");
|
| 924 |
+
}
|
| 925 |
+
}
|
| 926 |
+
|
| 927 |
+
async function doReset() {
|
| 928 |
+
state.task = nodes.taskInput.value;
|
| 929 |
+
state.seed = Number(nodes.seedInput.value || 0);
|
| 930 |
+
state.score = null;
|
| 931 |
+
return withAction("reset", async () => {
|
| 932 |
+
const payload = await apiCall("/reset", {
|
| 933 |
+
method: "POST",
|
| 934 |
+
body: JSON.stringify({ task: state.task, seed: state.seed }),
|
| 935 |
+
});
|
| 936 |
+
updateFromObservation(payload, false);
|
| 937 |
+
renderGraderTable(null);
|
| 938 |
+
return payload;
|
| 939 |
+
});
|
| 940 |
+
}
|
| 941 |
+
|
| 942 |
+
async function doStepWait() {
|
| 943 |
+
return withAction("step(wait)", async () => {
|
| 944 |
+
const payload = await apiCall("/step", {
|
| 945 |
+
method: "POST",
|
| 946 |
+
body: JSON.stringify({ move: null, accept_order: null, deliver_order: false, wait: true }),
|
| 947 |
+
});
|
| 948 |
+
updateFromObservation(payload.observation, payload.done);
|
| 949 |
+
return payload;
|
| 950 |
+
});
|
| 951 |
+
}
|
| 952 |
+
|
| 953 |
+
async function doState() {
|
| 954 |
+
return withAction("state", async () => {
|
| 955 |
+
const payload = await apiCall("/state", { method: "GET" });
|
| 956 |
+
updateFromState(payload);
|
| 957 |
+
return payload;
|
| 958 |
+
});
|
| 959 |
+
}
|
| 960 |
+
|
| 961 |
+
async function doGrader() {
|
| 962 |
+
return withAction("grader", async () => {
|
| 963 |
+
const payload = await apiCall("/grader", { method: "GET" });
|
| 964 |
+
if (payload.report && payload.report.score !== undefined) {
|
| 965 |
+
state.score = Number(payload.report.score);
|
| 966 |
+
}
|
| 967 |
+
renderGraderTable(payload.report || null);
|
| 968 |
+
return payload;
|
| 969 |
+
});
|
| 970 |
+
}
|
| 971 |
+
|
| 972 |
+
async function doBaseline() {
|
| 973 |
+
state.task = nodes.taskInput.value;
|
| 974 |
+
state.seed = Number(nodes.seedInput.value || 0);
|
| 975 |
+
return withAction("baseline", async () => {
|
| 976 |
+
const query = new URLSearchParams({ task: state.task, seed: String(state.seed) });
|
| 977 |
+
const payload = await apiCall(`/baseline?${query.toString()}`, { method: "GET" });
|
| 978 |
+
state.score = payload.score === undefined ? state.score : Number(payload.score);
|
| 979 |
+
renderGraderTable(payload.report || null);
|
| 980 |
+
return payload;
|
| 981 |
+
});
|
| 982 |
+
}
|
| 983 |
+
|
| 984 |
+
async function doSuite() {
|
| 985 |
+
state.seed = Number(nodes.seedInput.value || 0);
|
| 986 |
+
return withAction("baseline_suite", async () => {
|
| 987 |
+
const tasks = ["easy", "medium", "hard"];
|
| 988 |
+
const suiteRows = [];
|
| 989 |
+
|
| 990 |
+
for (let i = 0; i < tasks.length; i += 1) {
|
| 991 |
+
const task = tasks[i];
|
| 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: Boolean(payload.done),
|
| 998 |
+
score: Number(payload.score || 0),
|
| 999 |
+
steps_executed: Number(payload.steps_executed || 0),
|
| 1000 |
+
total_reward: Number(payload.report?.total_reward || 0),
|
| 1001 |
+
});
|
| 1002 |
+
}
|
| 1003 |
+
|
| 1004 |
+
state.suiteRows = suiteRows;
|
| 1005 |
+
renderSuiteTable(state.suiteRows);
|
| 1006 |
+
|
| 1007 |
+
return {
|
| 1008 |
+
seed: state.seed,
|
| 1009 |
+
tasks: suiteRows,
|
| 1010 |
+
average_score: suiteRows.length
|
| 1011 |
+
? Number((suiteRows.reduce((acc, item) => acc + Number(item.score || 0), 0) / suiteRows.length).toFixed(4))
|
| 1012 |
+
: 0,
|
| 1013 |
+
};
|
| 1014 |
+
});
|
| 1015 |
+
}
|
| 1016 |
+
|
| 1017 |
+
async function doDemoSequence() {
|
| 1018 |
+
setStatus("Running demo sequence reset -> step -> step -> state -> grader", "warn");
|
| 1019 |
+
await doReset();
|
| 1020 |
+
await doStepWait();
|
| 1021 |
+
await doStepWait();
|
| 1022 |
+
await doState();
|
| 1023 |
+
await doGrader();
|
| 1024 |
+
setStatus("Demo sequence finished.", "ok");
|
| 1025 |
+
}
|
| 1026 |
+
|
| 1027 |
+
nodes.btnReset.addEventListener("click", doReset);
|
| 1028 |
+
nodes.btnStep.addEventListener("click", doStepWait);
|
| 1029 |
+
nodes.btnState.addEventListener("click", doState);
|
| 1030 |
+
nodes.btnGrader.addEventListener("click", doGrader);
|
| 1031 |
+
nodes.btnBaseline.addEventListener("click", doBaseline);
|
| 1032 |
+
nodes.btnSuite.addEventListener("click", doSuite);
|
| 1033 |
+
nodes.btnDemo.addEventListener("click", doDemoSequence);
|
| 1034 |
+
nodes.apiBaseInput.addEventListener("change", saveApiBaseUrl);
|
| 1035 |
+
nodes.apiBaseInput.addEventListener("blur", saveApiBaseUrl);
|
| 1036 |
+
|
| 1037 |
+
async function initializeUi() {
|
| 1038 |
+
loadApiBaseUrl();
|
| 1039 |
+
renderMetrics();
|
| 1040 |
+
renderGraderTable(null);
|
| 1041 |
+
renderSuiteTable([]);
|
| 1042 |
+
|
| 1043 |
+
if (state.apiBaseUrl) {
|
| 1044 |
+
const savedHealthy = await probeHealth(state.apiBaseUrl);
|
| 1045 |
+
if (savedHealthy) {
|
| 1046 |
+
setStatus(`Ready. Endpoint: ${state.apiBaseUrl}`, "ok");
|
| 1047 |
+
return;
|
| 1048 |
+
}
|
| 1049 |
+
}
|
| 1050 |
+
|
| 1051 |
+
const detected = await autoDetectApiBaseUrl();
|
| 1052 |
+
if (detected) {
|
| 1053 |
+
setStatus(`Ready. Endpoint: ${currentApiBaseUrl()}`, "ok");
|
| 1054 |
+
return;
|
| 1055 |
+
}
|
| 1056 |
+
|
| 1057 |
+
setStatus("No API detected. Set API Base URL and ensure /health is reachable.", "error");
|
| 1058 |
+
}
|
| 1059 |
+
|
| 1060 |
+
initializeUi();
|
| 1061 |
+
</script>
|
| 1062 |
+
</body>
|
| 1063 |
+
</html>
|
baseline/__init__.py
ADDED
|
File without changes
|
baseline/baseline_agent.py
ADDED
|
@@ -0,0 +1,102 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from typing import List, Optional, Tuple
|
| 4 |
+
|
| 5 |
+
from env.models import ActionType, Direction, Observation, Order, OrderStatus, Position, SimulatorAction
|
| 6 |
+
from env.utils import manhattan
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
class BaselineGreedyAgent:
|
| 10 |
+
"""Greedy baseline that accepts nearest pending order and follows shortest-axis moves."""
|
| 11 |
+
|
| 12 |
+
def act(self, observation: Observation) -> SimulatorAction:
|
| 13 |
+
agent = observation.agent_location
|
| 14 |
+
current_order = observation.current_order
|
| 15 |
+
|
| 16 |
+
if current_order is None:
|
| 17 |
+
pending_orders = [
|
| 18 |
+
order for order in observation.pending_orders if order.status == OrderStatus.PENDING
|
| 19 |
+
]
|
| 20 |
+
if not pending_orders:
|
| 21 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 22 |
+
|
| 23 |
+
selected = self._nearest_pending(agent, pending_orders)
|
| 24 |
+
return SimulatorAction(action_type=ActionType.ACCEPT_ORDER, order_id=selected.order_id)
|
| 25 |
+
|
| 26 |
+
if current_order.status == OrderStatus.ACCEPTED:
|
| 27 |
+
if self._same(agent, current_order.pickup):
|
| 28 |
+
return SimulatorAction(action_type=ActionType.DELIVER_ORDER)
|
| 29 |
+
return self._move_towards(agent, current_order.pickup, observation)
|
| 30 |
+
|
| 31 |
+
if current_order.status == OrderStatus.PICKED_UP:
|
| 32 |
+
if self._same(agent, current_order.dropoff):
|
| 33 |
+
return SimulatorAction(action_type=ActionType.DELIVER_ORDER)
|
| 34 |
+
return self._move_towards(agent, current_order.dropoff, observation)
|
| 35 |
+
|
| 36 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 37 |
+
|
| 38 |
+
def _nearest_pending(self, agent: Position, orders: List[Order]) -> Order:
|
| 39 |
+
origin = (agent.x, agent.y)
|
| 40 |
+
return min(orders, key=lambda order: manhattan(origin, (order.pickup.x, order.pickup.y)))
|
| 41 |
+
|
| 42 |
+
def _move_towards(self, source: Position, target: Position, observation: Observation) -> SimulatorAction:
|
| 43 |
+
candidates = self._direction_candidates(source, target)
|
| 44 |
+
for direction in candidates:
|
| 45 |
+
if self._is_safe_move(source, direction, observation):
|
| 46 |
+
return SimulatorAction(action_type=ActionType.MOVE, direction=direction)
|
| 47 |
+
|
| 48 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 49 |
+
|
| 50 |
+
def _direction_candidates(self, source: Position, target: Position) -> List[Direction]:
|
| 51 |
+
directions: List[Direction] = []
|
| 52 |
+
|
| 53 |
+
dx = target.x - source.x
|
| 54 |
+
dy = target.y - source.y
|
| 55 |
+
|
| 56 |
+
if abs(dx) >= abs(dy):
|
| 57 |
+
if dx > 0:
|
| 58 |
+
directions.append(Direction.RIGHT)
|
| 59 |
+
elif dx < 0:
|
| 60 |
+
directions.append(Direction.LEFT)
|
| 61 |
+
|
| 62 |
+
if dy > 0:
|
| 63 |
+
directions.append(Direction.DOWN)
|
| 64 |
+
elif dy < 0:
|
| 65 |
+
directions.append(Direction.UP)
|
| 66 |
+
else:
|
| 67 |
+
if dy > 0:
|
| 68 |
+
directions.append(Direction.DOWN)
|
| 69 |
+
elif dy < 0:
|
| 70 |
+
directions.append(Direction.UP)
|
| 71 |
+
|
| 72 |
+
if dx > 0:
|
| 73 |
+
directions.append(Direction.RIGHT)
|
| 74 |
+
elif dx < 0:
|
| 75 |
+
directions.append(Direction.LEFT)
|
| 76 |
+
|
| 77 |
+
# Fall back to any movement if preferred route is blocked.
|
| 78 |
+
for d in [Direction.UP, Direction.DOWN, Direction.LEFT, Direction.RIGHT]:
|
| 79 |
+
if d not in directions:
|
| 80 |
+
directions.append(d)
|
| 81 |
+
return directions
|
| 82 |
+
|
| 83 |
+
def _is_safe_move(self, source: Position, direction: Direction, observation: Observation) -> bool:
|
| 84 |
+
x, y = source.x, source.y
|
| 85 |
+
if direction == Direction.UP:
|
| 86 |
+
y -= 1
|
| 87 |
+
elif direction == Direction.DOWN:
|
| 88 |
+
y += 1
|
| 89 |
+
elif direction == Direction.LEFT:
|
| 90 |
+
x -= 1
|
| 91 |
+
else:
|
| 92 |
+
x += 1
|
| 93 |
+
|
| 94 |
+
if x < 0 or y < 0 or x >= observation.grid_width or y >= observation.grid_height:
|
| 95 |
+
return False
|
| 96 |
+
|
| 97 |
+
obstacle_set = {(o.x, o.y) for o in observation.obstacles}
|
| 98 |
+
return (x, y) not in obstacle_set
|
| 99 |
+
|
| 100 |
+
@staticmethod
|
| 101 |
+
def _same(a: Position, b: Position) -> bool:
|
| 102 |
+
return a.x == b.x and a.y == b.y
|
baseline/trained_q_agent.py
ADDED
|
@@ -0,0 +1,176 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import json
|
| 4 |
+
from typing import Dict, List, Optional, Tuple
|
| 5 |
+
|
| 6 |
+
from env.models import ActionType, Direction, Observation, Order, OrderStatus, SimulatorAction
|
| 7 |
+
from env.utils import manhattan
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
StateKey = str
|
| 11 |
+
QValues = List[float]
|
| 12 |
+
QTable = Dict[StateKey, QValues]
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
ACTION_UP = 0
|
| 16 |
+
ACTION_DOWN = 1
|
| 17 |
+
ACTION_LEFT = 2
|
| 18 |
+
ACTION_RIGHT = 3
|
| 19 |
+
ACTION_ACCEPT = 4
|
| 20 |
+
ACTION_DELIVER = 5
|
| 21 |
+
ACTION_WAIT = 6
|
| 22 |
+
NUM_ACTIONS = 7
|
| 23 |
+
|
| 24 |
+
ACTION_LABELS: List[str] = [
|
| 25 |
+
"move_up",
|
| 26 |
+
"move_down",
|
| 27 |
+
"move_left",
|
| 28 |
+
"move_right",
|
| 29 |
+
"accept_nearest",
|
| 30 |
+
"deliver",
|
| 31 |
+
"wait",
|
| 32 |
+
]
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def _nearest_pending_order(observation: Observation) -> Optional[Order]:
|
| 36 |
+
pending = [order for order in observation.pending_orders if order.status == OrderStatus.PENDING]
|
| 37 |
+
if not pending:
|
| 38 |
+
return None
|
| 39 |
+
|
| 40 |
+
origin = (observation.agent_location.x, observation.agent_location.y)
|
| 41 |
+
return min(pending, key=lambda order: manhattan(origin, (order.pickup.x, order.pickup.y)))
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def _target_from_observation(observation: Observation) -> Tuple[int, int, int]:
|
| 45 |
+
current = observation.current_order
|
| 46 |
+
origin = (observation.agent_location.x, observation.agent_location.y)
|
| 47 |
+
|
| 48 |
+
# stage: 0=no active order, 1=going to pickup, 2=going to dropoff, 3=other
|
| 49 |
+
if current is None:
|
| 50 |
+
pending = _nearest_pending_order(observation)
|
| 51 |
+
if pending is None:
|
| 52 |
+
return -1, -1, 0
|
| 53 |
+
return pending.pickup.x, pending.pickup.y, 0
|
| 54 |
+
|
| 55 |
+
if current.status == OrderStatus.ACCEPTED:
|
| 56 |
+
return current.pickup.x, current.pickup.y, 1
|
| 57 |
+
|
| 58 |
+
if current.status == OrderStatus.PICKED_UP:
|
| 59 |
+
targets = current.delivery_locations or [current.dropoff]
|
| 60 |
+
selected = min(targets, key=lambda point: manhattan(origin, (point.x, point.y)))
|
| 61 |
+
return selected.x, selected.y, 2
|
| 62 |
+
|
| 63 |
+
return -1, -1, 3
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def _blocked_neighbor_count(observation: Observation) -> int:
|
| 67 |
+
blocked = {(item.x, item.y) for item in observation.obstacles}
|
| 68 |
+
blocked.update((item.x, item.y) for item in observation.dynamic_obstacles)
|
| 69 |
+
|
| 70 |
+
x = observation.agent_location.x
|
| 71 |
+
y = observation.agent_location.y
|
| 72 |
+
|
| 73 |
+
neighbors = [
|
| 74 |
+
(x, y - 1),
|
| 75 |
+
(x, y + 1),
|
| 76 |
+
(x - 1, y),
|
| 77 |
+
(x + 1, y),
|
| 78 |
+
]
|
| 79 |
+
|
| 80 |
+
count = 0
|
| 81 |
+
for nx, ny in neighbors:
|
| 82 |
+
if nx < 0 or ny < 0 or nx >= observation.grid_width or ny >= observation.grid_height:
|
| 83 |
+
count += 1
|
| 84 |
+
elif (nx, ny) in blocked:
|
| 85 |
+
count += 1
|
| 86 |
+
return count
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def _battery_bucket(level: Optional[int]) -> int:
|
| 90 |
+
if level is None:
|
| 91 |
+
return -1
|
| 92 |
+
if level <= 10:
|
| 93 |
+
return 0
|
| 94 |
+
if level <= 25:
|
| 95 |
+
return 1
|
| 96 |
+
if level <= 50:
|
| 97 |
+
return 2
|
| 98 |
+
return 3
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
def observation_key(observation: Observation) -> StateKey:
|
| 102 |
+
target_x, target_y, stage = _target_from_observation(observation)
|
| 103 |
+
|
| 104 |
+
if target_x < 0:
|
| 105 |
+
dist_bucket = 0
|
| 106 |
+
else:
|
| 107 |
+
dist = manhattan(
|
| 108 |
+
(observation.agent_location.x, observation.agent_location.y),
|
| 109 |
+
(target_x, target_y),
|
| 110 |
+
)
|
| 111 |
+
dist_bucket = min(20, dist)
|
| 112 |
+
|
| 113 |
+
pending_count = min(9, len(observation.pending_orders))
|
| 114 |
+
|
| 115 |
+
features = (
|
| 116 |
+
observation.grid_width,
|
| 117 |
+
observation.grid_height,
|
| 118 |
+
observation.agent_location.x,
|
| 119 |
+
observation.agent_location.y,
|
| 120 |
+
stage,
|
| 121 |
+
target_x,
|
| 122 |
+
target_y,
|
| 123 |
+
dist_bucket,
|
| 124 |
+
pending_count,
|
| 125 |
+
_battery_bucket(observation.battery_level),
|
| 126 |
+
_blocked_neighbor_count(observation),
|
| 127 |
+
)
|
| 128 |
+
|
| 129 |
+
return "|".join(str(value) for value in features)
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
def action_from_index(action_index: int, observation: Observation) -> SimulatorAction:
|
| 133 |
+
if action_index == ACTION_UP:
|
| 134 |
+
return SimulatorAction(action_type=ActionType.MOVE, direction=Direction.UP)
|
| 135 |
+
if action_index == ACTION_DOWN:
|
| 136 |
+
return SimulatorAction(action_type=ActionType.MOVE, direction=Direction.DOWN)
|
| 137 |
+
if action_index == ACTION_LEFT:
|
| 138 |
+
return SimulatorAction(action_type=ActionType.MOVE, direction=Direction.LEFT)
|
| 139 |
+
if action_index == ACTION_RIGHT:
|
| 140 |
+
return SimulatorAction(action_type=ActionType.MOVE, direction=Direction.RIGHT)
|
| 141 |
+
|
| 142 |
+
if action_index == ACTION_ACCEPT:
|
| 143 |
+
nearest = _nearest_pending_order(observation)
|
| 144 |
+
if nearest is None:
|
| 145 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 146 |
+
return SimulatorAction(action_type=ActionType.ACCEPT_ORDER, order_id=nearest.order_id)
|
| 147 |
+
|
| 148 |
+
if action_index == ACTION_DELIVER:
|
| 149 |
+
if observation.current_order is None:
|
| 150 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 151 |
+
return SimulatorAction(action_type=ActionType.DELIVER_ORDER)
|
| 152 |
+
|
| 153 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 154 |
+
|
| 155 |
+
|
| 156 |
+
class TrainedQAgent:
|
| 157 |
+
def __init__(self, q_table: QTable):
|
| 158 |
+
self.q_table = q_table
|
| 159 |
+
|
| 160 |
+
@classmethod
|
| 161 |
+
def load(cls, model_path: str) -> "TrainedQAgent":
|
| 162 |
+
with open(model_path, "r", encoding="utf-8") as handle:
|
| 163 |
+
payload = json.load(handle)
|
| 164 |
+
q_values = payload.get("q_values", {})
|
| 165 |
+
if not isinstance(q_values, dict):
|
| 166 |
+
raise ValueError("Invalid model file: q_values must be a dictionary")
|
| 167 |
+
return cls(q_table={str(key): [float(item) for item in value] for key, value in q_values.items()})
|
| 168 |
+
|
| 169 |
+
def act(self, observation: Observation) -> SimulatorAction:
|
| 170 |
+
key = observation_key(observation)
|
| 171 |
+
values = self.q_table.get(key)
|
| 172 |
+
if values is None or not values:
|
| 173 |
+
return action_from_index(ACTION_WAIT, observation)
|
| 174 |
+
|
| 175 |
+
best_index = max(range(min(NUM_ACTIONS, len(values))), key=lambda index: values[index])
|
| 176 |
+
return action_from_index(best_index, observation)
|
env/__init__.py
ADDED
|
File without changes
|
env/environment.py
ADDED
|
@@ -0,0 +1,175 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from typing import Any, Dict, Optional, Tuple, Union
|
| 4 |
+
|
| 5 |
+
from pydantic import ValidationError
|
| 6 |
+
|
| 7 |
+
from env.models import (
|
| 8 |
+
Action,
|
| 9 |
+
ActionType,
|
| 10 |
+
EnvironmentConfig,
|
| 11 |
+
EnvironmentState,
|
| 12 |
+
Observation,
|
| 13 |
+
RewardConfig,
|
| 14 |
+
ScenarioDefinition,
|
| 15 |
+
SimulatorAction,
|
| 16 |
+
StepInfo,
|
| 17 |
+
StepResult,
|
| 18 |
+
)
|
| 19 |
+
from env.simulator import DeliverySimulator
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
class LastMileDeliveryEnvironment:
|
| 23 |
+
"""OpenEnv-compatible environment for last-mile delivery optimization."""
|
| 24 |
+
|
| 25 |
+
def __init__(
|
| 26 |
+
self,
|
| 27 |
+
config: Optional[EnvironmentConfig] = None,
|
| 28 |
+
reward_config: Optional[RewardConfig] = None,
|
| 29 |
+
):
|
| 30 |
+
self.config = config or EnvironmentConfig()
|
| 31 |
+
self.reward_config = reward_config or RewardConfig()
|
| 32 |
+
|
| 33 |
+
self._simulator = DeliverySimulator(config=self.config, reward_config=self.reward_config)
|
| 34 |
+
self._initialized = False
|
| 35 |
+
self._total_reward = 0.0
|
| 36 |
+
self._last_observation: Optional[Observation] = None
|
| 37 |
+
|
| 38 |
+
def reset(self, seed: Optional[int] = None) -> Observation:
|
| 39 |
+
if seed is not None:
|
| 40 |
+
self.config.seed = seed
|
| 41 |
+
|
| 42 |
+
self._simulator.reset(seed=seed)
|
| 43 |
+
self._initialized = True
|
| 44 |
+
self._total_reward = 0.0
|
| 45 |
+
self._last_observation = self._simulator.build_observation(total_reward=self._total_reward)
|
| 46 |
+
return self._last_observation
|
| 47 |
+
|
| 48 |
+
def reset_with_scenario(self, scenario: ScenarioDefinition, seed: Optional[int] = None) -> Observation:
|
| 49 |
+
if seed is not None:
|
| 50 |
+
self.config.seed = seed
|
| 51 |
+
|
| 52 |
+
self._simulator.reset_with_scenario(scenario=scenario, seed=seed)
|
| 53 |
+
self._initialized = True
|
| 54 |
+
self._total_reward = 0.0
|
| 55 |
+
self._last_observation = self._simulator.build_observation(total_reward=self._total_reward)
|
| 56 |
+
return self._last_observation
|
| 57 |
+
|
| 58 |
+
def step(self, action: Union[Action, SimulatorAction, Dict[str, Any]]) -> Tuple[Observation, float, bool, Dict[str, Any]]:
|
| 59 |
+
"""OpenEnv step contract: returns (observation, reward, done, info)."""
|
| 60 |
+
if not self._initialized:
|
| 61 |
+
raise RuntimeError("Environment not initialized. Call reset() before step().")
|
| 62 |
+
|
| 63 |
+
if self._simulator.state.done:
|
| 64 |
+
observation = self._build_observation()
|
| 65 |
+
info = StepInfo(message="episode_already_done")
|
| 66 |
+
return observation, 0.0, True, info.model_dump()
|
| 67 |
+
|
| 68 |
+
validated_action, error_message = self._validate_action(action)
|
| 69 |
+
if validated_action is None:
|
| 70 |
+
reward, info = self._apply_invalid_action_penalty(error_message or "invalid_action_payload")
|
| 71 |
+
else:
|
| 72 |
+
reward, info = self._simulator.apply_action(validated_action)
|
| 73 |
+
|
| 74 |
+
self._total_reward += reward
|
| 75 |
+
observation = self._build_observation()
|
| 76 |
+
|
| 77 |
+
return observation, reward, self._simulator.state.done, info.model_dump()
|
| 78 |
+
|
| 79 |
+
def step_result(self, action: Union[Action, SimulatorAction, Dict[str, Any]]) -> StepResult:
|
| 80 |
+
"""Compatibility wrapper for existing API layers that consume StepResult."""
|
| 81 |
+
observation, reward, done, info = self.step(action)
|
| 82 |
+
return StepResult(
|
| 83 |
+
observation=observation,
|
| 84 |
+
reward=reward,
|
| 85 |
+
done=done,
|
| 86 |
+
info=StepInfo.model_validate(info),
|
| 87 |
+
)
|
| 88 |
+
|
| 89 |
+
def reset_state(self, seed: Optional[int] = None) -> EnvironmentState:
|
| 90 |
+
"""Compatibility helper for external runners expecting state immediately after reset."""
|
| 91 |
+
self.reset(seed=seed)
|
| 92 |
+
return self.state()
|
| 93 |
+
|
| 94 |
+
def current_observation(self) -> Observation:
|
| 95 |
+
if not self._initialized or self._last_observation is None:
|
| 96 |
+
raise RuntimeError("Environment not initialized. Call reset() first.")
|
| 97 |
+
return self._last_observation
|
| 98 |
+
|
| 99 |
+
def state(self) -> EnvironmentState:
|
| 100 |
+
if not self._initialized:
|
| 101 |
+
return EnvironmentState(
|
| 102 |
+
initialized=False,
|
| 103 |
+
done=False,
|
| 104 |
+
total_reward=0.0,
|
| 105 |
+
step_count=0,
|
| 106 |
+
delivered_orders=0,
|
| 107 |
+
current_order_id=None,
|
| 108 |
+
observation=None,
|
| 109 |
+
)
|
| 110 |
+
|
| 111 |
+
return EnvironmentState(
|
| 112 |
+
initialized=True,
|
| 113 |
+
done=self._simulator.state.done,
|
| 114 |
+
total_reward=self._total_reward,
|
| 115 |
+
step_count=self._simulator.state.step_count,
|
| 116 |
+
delivered_orders=self._simulator.state.delivered_orders,
|
| 117 |
+
current_order_id=self._simulator.state.current_order_id,
|
| 118 |
+
observation=self._build_observation(),
|
| 119 |
+
)
|
| 120 |
+
|
| 121 |
+
def _build_observation(self) -> Observation:
|
| 122 |
+
self._last_observation = self._simulator.build_observation(total_reward=self._total_reward)
|
| 123 |
+
return self._last_observation
|
| 124 |
+
|
| 125 |
+
def _apply_invalid_action_penalty(self, message: str) -> Tuple[float, StepInfo]:
|
| 126 |
+
# Invalid payloads still consume a step to keep episode progression consistent.
|
| 127 |
+
base_reward, base_info = self._simulator.apply_action(SimulatorAction(action_type=ActionType.WAIT))
|
| 128 |
+
reward = base_reward + self.reward_config.invalid_action_penalty
|
| 129 |
+
info = base_info.model_copy(deep=True)
|
| 130 |
+
info.invalid_action = True
|
| 131 |
+
info.message = message
|
| 132 |
+
info.destination_reached = False
|
| 133 |
+
info.delivered_order_id = None
|
| 134 |
+
return reward, info
|
| 135 |
+
|
| 136 |
+
def _validate_action(
|
| 137 |
+
self,
|
| 138 |
+
action: Union[Action, SimulatorAction, Dict[str, Any]],
|
| 139 |
+
) -> Tuple[Optional[SimulatorAction], Optional[str]]:
|
| 140 |
+
parsed: Optional[SimulatorAction] = None
|
| 141 |
+
if isinstance(action, SimulatorAction):
|
| 142 |
+
parsed = action
|
| 143 |
+
elif isinstance(action, Action):
|
| 144 |
+
try:
|
| 145 |
+
parsed = action.to_simulator_action()
|
| 146 |
+
except ValueError:
|
| 147 |
+
return None, "invalid_action_payload"
|
| 148 |
+
else:
|
| 149 |
+
try:
|
| 150 |
+
public_action = Action.model_validate(action)
|
| 151 |
+
parsed = public_action.to_simulator_action()
|
| 152 |
+
except ValidationError:
|
| 153 |
+
return None, "invalid_action_payload"
|
| 154 |
+
except ValueError:
|
| 155 |
+
return None, "invalid_action_payload"
|
| 156 |
+
|
| 157 |
+
if parsed.action_type == ActionType.MOVE:
|
| 158 |
+
if parsed.direction is None:
|
| 159 |
+
return None, "direction_required_for_move"
|
| 160 |
+
if parsed.order_id is not None:
|
| 161 |
+
return None, "order_id_not_allowed_for_move"
|
| 162 |
+
|
| 163 |
+
if parsed.action_type == ActionType.ACCEPT_ORDER and parsed.direction is not None:
|
| 164 |
+
return None, "direction_not_allowed_for_accept_order"
|
| 165 |
+
|
| 166 |
+
if parsed.action_type in {ActionType.DELIVER_ORDER, ActionType.WAIT}:
|
| 167 |
+
if parsed.direction is not None:
|
| 168 |
+
return None, "direction_not_allowed_for_action"
|
| 169 |
+
if parsed.order_id is not None:
|
| 170 |
+
return None, "order_id_not_allowed_for_action"
|
| 171 |
+
|
| 172 |
+
if parsed.action_type == ActionType.DELIVER_ORDER and self._simulator.state.current_order_id is None:
|
| 173 |
+
return None, "deliver_without_active_order"
|
| 174 |
+
|
| 175 |
+
return parsed, None
|
env/models.py
ADDED
|
@@ -0,0 +1,335 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from enum import Enum
|
| 4 |
+
from typing import Dict, List, Optional, Tuple
|
| 5 |
+
|
| 6 |
+
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
Coordinate = Tuple[int, int]
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class Direction(str, Enum):
|
| 13 |
+
UP = "up"
|
| 14 |
+
DOWN = "down"
|
| 15 |
+
LEFT = "left"
|
| 16 |
+
RIGHT = "right"
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class MoveDirection(str, Enum):
|
| 20 |
+
UP = "up"
|
| 21 |
+
DOWN = "down"
|
| 22 |
+
LEFT = "left"
|
| 23 |
+
RIGHT = "right"
|
| 24 |
+
STAY = "stay"
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
class ActionType(str, Enum):
|
| 28 |
+
MOVE = "move"
|
| 29 |
+
ACCEPT_ORDER = "accept_order"
|
| 30 |
+
DELIVER_ORDER = "deliver_order"
|
| 31 |
+
WAIT = "wait"
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
class OrderStatus(str, Enum):
|
| 35 |
+
PENDING = "pending"
|
| 36 |
+
ACCEPTED = "accepted"
|
| 37 |
+
PICKED_UP = "picked_up"
|
| 38 |
+
DELIVERED = "delivered"
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
class OrderPriority(str, Enum):
|
| 42 |
+
HIGH = "high"
|
| 43 |
+
LOW = "low"
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
class Position(BaseModel):
|
| 47 |
+
x: int = Field(..., ge=0)
|
| 48 |
+
y: int = Field(..., ge=0)
|
| 49 |
+
|
| 50 |
+
def to_tuple(self) -> Coordinate:
|
| 51 |
+
return self.x, self.y
|
| 52 |
+
|
| 53 |
+
@classmethod
|
| 54 |
+
def from_tuple(cls, value: Coordinate) -> "Position":
|
| 55 |
+
return cls(x=value[0], y=value[1])
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
class TrafficZone(BaseModel):
|
| 59 |
+
location: Position
|
| 60 |
+
extra_cost: int = Field(default=1, ge=1)
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
class SimulatorAction(BaseModel):
|
| 64 |
+
action_type: ActionType
|
| 65 |
+
direction: Optional[Direction] = None
|
| 66 |
+
order_id: Optional[str] = None
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
class Action(BaseModel):
|
| 70 |
+
model_config = ConfigDict(extra="forbid")
|
| 71 |
+
|
| 72 |
+
move: Optional[MoveDirection] = None
|
| 73 |
+
accept_order: Optional[int] = None
|
| 74 |
+
deliver_order: bool = False
|
| 75 |
+
wait: bool = False
|
| 76 |
+
|
| 77 |
+
@model_validator(mode="after")
|
| 78 |
+
def _validate_single_active_field(self) -> "Action":
|
| 79 |
+
if isinstance(self.accept_order, bool):
|
| 80 |
+
raise ValueError("accept_order must be an integer or null")
|
| 81 |
+
|
| 82 |
+
active_count = sum(
|
| 83 |
+
[
|
| 84 |
+
self.move is not None,
|
| 85 |
+
self.accept_order is not None,
|
| 86 |
+
self.deliver_order,
|
| 87 |
+
self.wait,
|
| 88 |
+
]
|
| 89 |
+
)
|
| 90 |
+
if active_count != 1:
|
| 91 |
+
raise ValueError("Exactly one action field must be active")
|
| 92 |
+
return self
|
| 93 |
+
|
| 94 |
+
def to_simulator_action(self) -> SimulatorAction:
|
| 95 |
+
if self.move is not None:
|
| 96 |
+
if self.move == MoveDirection.STAY:
|
| 97 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 98 |
+
return SimulatorAction(action_type=ActionType.MOVE, direction=Direction(self.move.value))
|
| 99 |
+
|
| 100 |
+
if self.accept_order is not None:
|
| 101 |
+
return SimulatorAction(action_type=ActionType.ACCEPT_ORDER, order_id=f"order_{self.accept_order}")
|
| 102 |
+
|
| 103 |
+
if self.deliver_order:
|
| 104 |
+
return SimulatorAction(action_type=ActionType.DELIVER_ORDER)
|
| 105 |
+
|
| 106 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 107 |
+
|
| 108 |
+
@classmethod
|
| 109 |
+
def from_simulator_action(cls, action: SimulatorAction) -> "Action":
|
| 110 |
+
if action.action_type == ActionType.MOVE:
|
| 111 |
+
if action.direction is None:
|
| 112 |
+
return cls(move=MoveDirection.STAY)
|
| 113 |
+
return cls(move=MoveDirection(action.direction.value))
|
| 114 |
+
|
| 115 |
+
if action.action_type == ActionType.ACCEPT_ORDER:
|
| 116 |
+
order_id = action.order_id or ""
|
| 117 |
+
if order_id.startswith("order_") and order_id[6:].isdigit():
|
| 118 |
+
order_number = int(order_id[6:])
|
| 119 |
+
else:
|
| 120 |
+
digits = "".join(ch for ch in order_id if ch.isdigit())
|
| 121 |
+
order_number = int(digits) if digits else 0
|
| 122 |
+
return cls(accept_order=order_number)
|
| 123 |
+
|
| 124 |
+
if action.action_type == ActionType.DELIVER_ORDER:
|
| 125 |
+
return cls(deliver_order=True)
|
| 126 |
+
|
| 127 |
+
return cls(wait=True)
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
class Order(BaseModel):
|
| 131 |
+
order_id: str
|
| 132 |
+
pickup: Position
|
| 133 |
+
dropoff: Position
|
| 134 |
+
delivery_locations: List[Position] = Field(default_factory=list)
|
| 135 |
+
priority: OrderPriority = OrderPriority.LOW
|
| 136 |
+
created_step: int = 0
|
| 137 |
+
accepted_step: Optional[int] = None
|
| 138 |
+
picked_up_step: Optional[int] = None
|
| 139 |
+
delivered_step: Optional[int] = None
|
| 140 |
+
status: OrderStatus = OrderStatus.PENDING
|
| 141 |
+
|
| 142 |
+
|
| 143 |
+
class Observation(BaseModel):
|
| 144 |
+
grid_width: int
|
| 145 |
+
grid_height: int
|
| 146 |
+
agent_location: Position
|
| 147 |
+
pending_orders: List[Order] = Field(default_factory=list)
|
| 148 |
+
current_order: Optional[Order] = None
|
| 149 |
+
obstacles: List[Position] = Field(default_factory=list)
|
| 150 |
+
dynamic_obstacles: List[Position] = Field(default_factory=list)
|
| 151 |
+
traffic_zones: List[TrafficZone] = Field(default_factory=list)
|
| 152 |
+
charging_stations: List[Position] = Field(default_factory=list)
|
| 153 |
+
battery_level: Optional[int] = None
|
| 154 |
+
step_count: int = 0
|
| 155 |
+
total_reward: float = 0.0
|
| 156 |
+
|
| 157 |
+
|
| 158 |
+
class Reward(BaseModel):
|
| 159 |
+
value: float
|
| 160 |
+
components: Dict[str, float] = Field(default_factory=dict)
|
| 161 |
+
|
| 162 |
+
|
| 163 |
+
class StepInfo(BaseModel):
|
| 164 |
+
invalid_action: bool = False
|
| 165 |
+
message: str = ""
|
| 166 |
+
destination_reached: bool = False
|
| 167 |
+
delivered_order_id: Optional[str] = None
|
| 168 |
+
delay_penalty_applied: bool = False
|
| 169 |
+
battery_depleted: bool = False
|
| 170 |
+
made_progress: Optional[bool] = None
|
| 171 |
+
|
| 172 |
+
|
| 173 |
+
class StepResult(BaseModel):
|
| 174 |
+
observation: Observation
|
| 175 |
+
reward: float
|
| 176 |
+
done: bool
|
| 177 |
+
info: StepInfo
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
class EnvironmentState(BaseModel):
|
| 181 |
+
initialized: bool
|
| 182 |
+
done: bool
|
| 183 |
+
total_reward: float
|
| 184 |
+
step_count: int
|
| 185 |
+
delivered_orders: int
|
| 186 |
+
current_order_id: Optional[str] = None
|
| 187 |
+
observation: Optional[Observation] = None
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
class EnvironmentConfig(BaseModel):
|
| 191 |
+
width: int = Field(default=10, ge=2)
|
| 192 |
+
height: int = Field(default=10, ge=2)
|
| 193 |
+
max_steps: int = Field(default=200, ge=1)
|
| 194 |
+
max_orders: int = Field(default=5, ge=1)
|
| 195 |
+
delivery_locations_per_order: int = Field(default=1, ge=1, le=3)
|
| 196 |
+
priority_high_ratio: float = Field(default=0.4, ge=0.0, le=1.0)
|
| 197 |
+
obstacle_density: float = Field(default=0.1, ge=0.0, le=0.4)
|
| 198 |
+
dynamic_obstacles_enabled: bool = False
|
| 199 |
+
dynamic_obstacle_ratio: float = Field(default=0.25, ge=0.0, le=1.0)
|
| 200 |
+
traffic_density: float = Field(default=0.1, ge=0.0, le=0.4)
|
| 201 |
+
traffic_extra_cost: int = Field(default=2, ge=1)
|
| 202 |
+
battery_enabled: bool = False
|
| 203 |
+
battery_capacity: int = Field(default=100, ge=1)
|
| 204 |
+
battery_recharge_rate: int = Field(default=12, ge=1)
|
| 205 |
+
charging_stations: List[Position] = Field(default_factory=lambda: [Position(x=0, y=0)])
|
| 206 |
+
seed: Optional[int] = None
|
| 207 |
+
|
| 208 |
+
|
| 209 |
+
class RewardConfig(BaseModel):
|
| 210 |
+
delivery_reward: float = 50.0
|
| 211 |
+
destination_reward: float = 10.0
|
| 212 |
+
step_penalty: float = -1.0
|
| 213 |
+
invalid_action_penalty: float = -20.0
|
| 214 |
+
delay_penalty: float = -5.0
|
| 215 |
+
progress_reward_scale: float = 1.25
|
| 216 |
+
progress_reward_clip: float = 2.0
|
| 217 |
+
delay_penalty_scale: float = 0.35
|
| 218 |
+
delay_penalty_cap: float = 6.0
|
| 219 |
+
efficiency_bonus_max: float = 12.0
|
| 220 |
+
traffic_penalty: float = -2.0
|
| 221 |
+
high_priority_delivery_bonus: float = 15.0
|
| 222 |
+
low_priority_delivery_bonus: float = 5.0
|
| 223 |
+
recharge_reward: float = 2.0
|
| 224 |
+
battery_failure_penalty: float = -30.0
|
| 225 |
+
|
| 226 |
+
|
| 227 |
+
class ScenarioBatteryProfile(BaseModel):
|
| 228 |
+
enabled: bool = False
|
| 229 |
+
capacity: int = Field(default=100, ge=1)
|
| 230 |
+
recharge_rate: int = Field(default=12, ge=1)
|
| 231 |
+
initial_level: Optional[int] = Field(default=None, ge=0)
|
| 232 |
+
|
| 233 |
+
@model_validator(mode="after")
|
| 234 |
+
def _validate_initial_level(self) -> "ScenarioBatteryProfile":
|
| 235 |
+
if self.initial_level is not None and self.initial_level > self.capacity:
|
| 236 |
+
raise ValueError("battery_profile.initial_level cannot exceed capacity")
|
| 237 |
+
return self
|
| 238 |
+
|
| 239 |
+
|
| 240 |
+
class ScenarioOrder(BaseModel):
|
| 241 |
+
order_id: Optional[str] = None
|
| 242 |
+
pickup: Position
|
| 243 |
+
dropoff: Position
|
| 244 |
+
delivery_locations: List[Position] = Field(default_factory=list)
|
| 245 |
+
priority: OrderPriority = OrderPriority.LOW
|
| 246 |
+
created_step: int = Field(default=0, ge=0)
|
| 247 |
+
|
| 248 |
+
@model_validator(mode="after")
|
| 249 |
+
def _ensure_delivery_locations(self) -> "ScenarioOrder":
|
| 250 |
+
if not self.delivery_locations:
|
| 251 |
+
self.delivery_locations = [self.dropoff]
|
| 252 |
+
return self
|
| 253 |
+
|
| 254 |
+
|
| 255 |
+
class ScenarioDefinition(BaseModel):
|
| 256 |
+
width: int = Field(..., ge=2)
|
| 257 |
+
height: int = Field(..., ge=2)
|
| 258 |
+
max_steps: int = Field(default=200, ge=1)
|
| 259 |
+
agent_start: Position = Field(default_factory=lambda: Position(x=0, y=0))
|
| 260 |
+
orders: List[ScenarioOrder] = Field(..., min_length=1)
|
| 261 |
+
obstacles: List[Position] = Field(default_factory=list)
|
| 262 |
+
dynamic_obstacles: List[Position] = Field(default_factory=list)
|
| 263 |
+
traffic_zones: List[TrafficZone] = Field(default_factory=list)
|
| 264 |
+
charging_stations: List[Position] = Field(default_factory=lambda: [Position(x=0, y=0)])
|
| 265 |
+
battery_profile: ScenarioBatteryProfile = Field(default_factory=ScenarioBatteryProfile)
|
| 266 |
+
|
| 267 |
+
@model_validator(mode="after")
|
| 268 |
+
def _validate_geometry(self) -> "ScenarioDefinition":
|
| 269 |
+
def _in_bounds(position: Position) -> bool:
|
| 270 |
+
return 0 <= position.x < self.width and 0 <= position.y < self.height
|
| 271 |
+
|
| 272 |
+
def _dedupe(label: str, points: List[Position]) -> set[Coordinate]:
|
| 273 |
+
coordinates = [item.to_tuple() for item in points]
|
| 274 |
+
if len(set(coordinates)) != len(coordinates):
|
| 275 |
+
raise ValueError(f"{label} contains duplicate coordinates")
|
| 276 |
+
return set(coordinates)
|
| 277 |
+
|
| 278 |
+
agent_cell = self.agent_start.to_tuple()
|
| 279 |
+
if not _in_bounds(self.agent_start):
|
| 280 |
+
raise ValueError("agent_start must be within grid bounds")
|
| 281 |
+
|
| 282 |
+
static_cells = _dedupe("obstacles", self.obstacles)
|
| 283 |
+
dynamic_cells = _dedupe("dynamic_obstacles", self.dynamic_obstacles)
|
| 284 |
+
if static_cells.intersection(dynamic_cells):
|
| 285 |
+
raise ValueError("obstacles and dynamic_obstacles cannot overlap")
|
| 286 |
+
|
| 287 |
+
obstacle_cells = static_cells.union(dynamic_cells)
|
| 288 |
+
if agent_cell in obstacle_cells:
|
| 289 |
+
raise ValueError("agent_start cannot overlap an obstacle")
|
| 290 |
+
|
| 291 |
+
charging_cells = _dedupe("charging_stations", self.charging_stations)
|
| 292 |
+
if self.battery_profile.enabled and not charging_cells:
|
| 293 |
+
raise ValueError("charging_stations is required when battery_profile.enabled is true")
|
| 294 |
+
if charging_cells.intersection(obstacle_cells):
|
| 295 |
+
raise ValueError("charging_stations cannot overlap obstacles")
|
| 296 |
+
|
| 297 |
+
traffic_points = [zone.location for zone in self.traffic_zones]
|
| 298 |
+
traffic_cells = _dedupe("traffic_zones", traffic_points)
|
| 299 |
+
if traffic_cells.intersection(obstacle_cells):
|
| 300 |
+
raise ValueError("traffic_zones cannot overlap obstacles")
|
| 301 |
+
|
| 302 |
+
for label, points in [
|
| 303 |
+
("obstacles", self.obstacles),
|
| 304 |
+
("dynamic_obstacles", self.dynamic_obstacles),
|
| 305 |
+
("charging_stations", self.charging_stations),
|
| 306 |
+
("traffic_zones", traffic_points),
|
| 307 |
+
]:
|
| 308 |
+
for point in points:
|
| 309 |
+
if not _in_bounds(point):
|
| 310 |
+
raise ValueError(f"{label} contains out-of-bounds coordinates")
|
| 311 |
+
|
| 312 |
+
seen_order_ids: set[str] = set()
|
| 313 |
+
for index, order in enumerate(self.orders):
|
| 314 |
+
if order.order_id:
|
| 315 |
+
if order.order_id in seen_order_ids:
|
| 316 |
+
raise ValueError("orders contains duplicate order_id values")
|
| 317 |
+
seen_order_ids.add(order.order_id)
|
| 318 |
+
|
| 319 |
+
for point in [order.pickup, *order.delivery_locations]:
|
| 320 |
+
if not _in_bounds(point):
|
| 321 |
+
raise ValueError(f"orders[{index}] contains out-of-bounds coordinates")
|
| 322 |
+
if point.to_tuple() in obstacle_cells:
|
| 323 |
+
raise ValueError(f"orders[{index}] cannot overlap obstacles")
|
| 324 |
+
|
| 325 |
+
return self
|
| 326 |
+
|
| 327 |
+
|
| 328 |
+
class SimulationState(BaseModel):
|
| 329 |
+
agent_location: Position = Field(default_factory=lambda: Position(x=0, y=0))
|
| 330 |
+
orders: Dict[str, Order] = Field(default_factory=dict)
|
| 331 |
+
current_order_id: Optional[str] = None
|
| 332 |
+
step_count: int = 0
|
| 333 |
+
battery_level: Optional[int] = None
|
| 334 |
+
done: bool = False
|
| 335 |
+
delivered_orders: int = 0
|
env/simulator.py
ADDED
|
@@ -0,0 +1,622 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import random
|
| 4 |
+
from typing import Dict, List, Optional, Set, Tuple
|
| 5 |
+
|
| 6 |
+
from env.models import (
|
| 7 |
+
ActionType,
|
| 8 |
+
Coordinate,
|
| 9 |
+
Direction,
|
| 10 |
+
EnvironmentConfig,
|
| 11 |
+
Observation,
|
| 12 |
+
Order,
|
| 13 |
+
OrderPriority,
|
| 14 |
+
OrderStatus,
|
| 15 |
+
Position,
|
| 16 |
+
RewardConfig,
|
| 17 |
+
SimulatorAction,
|
| 18 |
+
SimulationState,
|
| 19 |
+
ScenarioDefinition,
|
| 20 |
+
StepInfo,
|
| 21 |
+
TrafficZone,
|
| 22 |
+
)
|
| 23 |
+
from env.utils import as_tuple, in_bounds, manhattan, move_coordinate, sample_locations, to_position_list
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
class DeliverySimulator:
|
| 27 |
+
"""Core deterministic simulator for last-mile delivery dynamics."""
|
| 28 |
+
|
| 29 |
+
def __init__(self, config: EnvironmentConfig, reward_config: Optional[RewardConfig] = None):
|
| 30 |
+
self.config = config
|
| 31 |
+
self.reward_config = reward_config or RewardConfig()
|
| 32 |
+
self._active_seed: Optional[int] = config.seed
|
| 33 |
+
self.rng = random.Random(self._active_seed)
|
| 34 |
+
|
| 35 |
+
self._static_obstacles: Set[Coordinate] = set()
|
| 36 |
+
self._dynamic_obstacles: Set[Coordinate] = set()
|
| 37 |
+
self._obstacles: Set[Coordinate] = set()
|
| 38 |
+
self._traffic: Dict[Coordinate, int] = {}
|
| 39 |
+
self._charging_stations: Set[Coordinate] = set()
|
| 40 |
+
self._state = SimulationState()
|
| 41 |
+
|
| 42 |
+
@property
|
| 43 |
+
def state(self) -> SimulationState:
|
| 44 |
+
return self._state
|
| 45 |
+
|
| 46 |
+
@property
|
| 47 |
+
def obstacles(self) -> Set[Coordinate]:
|
| 48 |
+
return self._obstacles
|
| 49 |
+
|
| 50 |
+
@property
|
| 51 |
+
def dynamic_obstacles(self) -> Set[Coordinate]:
|
| 52 |
+
return self._dynamic_obstacles
|
| 53 |
+
|
| 54 |
+
@property
|
| 55 |
+
def charging_stations(self) -> Set[Coordinate]:
|
| 56 |
+
return self._charging_stations
|
| 57 |
+
|
| 58 |
+
@property
|
| 59 |
+
def traffic(self) -> Dict[Coordinate, int]:
|
| 60 |
+
return self._traffic
|
| 61 |
+
|
| 62 |
+
def reset(self, seed: Optional[int] = None) -> SimulationState:
|
| 63 |
+
if seed is not None:
|
| 64 |
+
self._active_seed = seed
|
| 65 |
+
self.config.seed = seed
|
| 66 |
+
elif self.config.seed is not None and self.config.seed != self._active_seed:
|
| 67 |
+
self._active_seed = self.config.seed
|
| 68 |
+
|
| 69 |
+
self.rng.seed(self._active_seed)
|
| 70 |
+
|
| 71 |
+
configured_stations = {as_tuple(p) for p in self.config.charging_stations}
|
| 72 |
+
self._charging_stations = configured_stations or {(0, 0)}
|
| 73 |
+
self._generate_map()
|
| 74 |
+
self._state = SimulationState(
|
| 75 |
+
agent_location=Position(x=0, y=0),
|
| 76 |
+
orders=self._generate_orders(),
|
| 77 |
+
current_order_id=None,
|
| 78 |
+
step_count=0,
|
| 79 |
+
battery_level=self.config.battery_capacity if self.config.battery_enabled else None,
|
| 80 |
+
done=False,
|
| 81 |
+
delivered_orders=0,
|
| 82 |
+
)
|
| 83 |
+
return self._state
|
| 84 |
+
|
| 85 |
+
def reset_with_scenario(self, scenario: ScenarioDefinition, seed: Optional[int] = None) -> SimulationState:
|
| 86 |
+
if seed is not None:
|
| 87 |
+
self._active_seed = seed
|
| 88 |
+
self.config.seed = seed
|
| 89 |
+
elif self.config.seed is not None and self.config.seed != self._active_seed:
|
| 90 |
+
self._active_seed = self.config.seed
|
| 91 |
+
|
| 92 |
+
self.rng.seed(self._active_seed)
|
| 93 |
+
|
| 94 |
+
self.config.width = scenario.width
|
| 95 |
+
self.config.height = scenario.height
|
| 96 |
+
self.config.max_steps = scenario.max_steps
|
| 97 |
+
self.config.max_orders = max(1, len(scenario.orders))
|
| 98 |
+
self.config.dynamic_obstacles_enabled = bool(scenario.dynamic_obstacles)
|
| 99 |
+
self.config.dynamic_obstacle_ratio = 1.0 if scenario.dynamic_obstacles else 0.0
|
| 100 |
+
self.config.obstacle_density = 0.0
|
| 101 |
+
self.config.traffic_density = 0.0
|
| 102 |
+
self.config.battery_enabled = scenario.battery_profile.enabled
|
| 103 |
+
self.config.battery_capacity = scenario.battery_profile.capacity
|
| 104 |
+
self.config.battery_recharge_rate = scenario.battery_profile.recharge_rate
|
| 105 |
+
self.config.charging_stations = [item.model_copy(deep=True) for item in scenario.charging_stations]
|
| 106 |
+
|
| 107 |
+
self._static_obstacles = {as_tuple(item) for item in scenario.obstacles}
|
| 108 |
+
self._dynamic_obstacles = {as_tuple(item) for item in scenario.dynamic_obstacles}
|
| 109 |
+
self._obstacles = self._static_obstacles.union(self._dynamic_obstacles)
|
| 110 |
+
self._traffic = {
|
| 111 |
+
as_tuple(zone.location): zone.extra_cost
|
| 112 |
+
for zone in scenario.traffic_zones
|
| 113 |
+
}
|
| 114 |
+
self._charging_stations = {as_tuple(item) for item in scenario.charging_stations} or {(0, 0)}
|
| 115 |
+
|
| 116 |
+
initial_battery_level: Optional[int] = None
|
| 117 |
+
if scenario.battery_profile.enabled:
|
| 118 |
+
initial_level = scenario.battery_profile.initial_level
|
| 119 |
+
initial_battery_level = (
|
| 120 |
+
scenario.battery_profile.capacity
|
| 121 |
+
if initial_level is None
|
| 122 |
+
else initial_level
|
| 123 |
+
)
|
| 124 |
+
|
| 125 |
+
self._state = SimulationState(
|
| 126 |
+
agent_location=scenario.agent_start.model_copy(deep=True),
|
| 127 |
+
orders=self._orders_from_scenario(scenario),
|
| 128 |
+
current_order_id=None,
|
| 129 |
+
step_count=0,
|
| 130 |
+
battery_level=initial_battery_level,
|
| 131 |
+
done=False,
|
| 132 |
+
delivered_orders=0,
|
| 133 |
+
)
|
| 134 |
+
return self._state
|
| 135 |
+
|
| 136 |
+
def build_observation(self, total_reward: float = 0.0) -> Observation:
|
| 137 |
+
pending_orders = [
|
| 138 |
+
order for order in self._state.orders.values() if order.status != OrderStatus.DELIVERED
|
| 139 |
+
]
|
| 140 |
+
current_order = None
|
| 141 |
+
if self._state.current_order_id is not None:
|
| 142 |
+
current_order = self._state.orders.get(self._state.current_order_id)
|
| 143 |
+
|
| 144 |
+
traffic_zones = [
|
| 145 |
+
TrafficZone(location=Position.from_tuple(loc), extra_cost=cost)
|
| 146 |
+
for loc, cost in self._traffic.items()
|
| 147 |
+
]
|
| 148 |
+
|
| 149 |
+
return Observation(
|
| 150 |
+
grid_width=self.config.width,
|
| 151 |
+
grid_height=self.config.height,
|
| 152 |
+
agent_location=self._state.agent_location,
|
| 153 |
+
pending_orders=pending_orders,
|
| 154 |
+
current_order=current_order,
|
| 155 |
+
obstacles=to_position_list(self._obstacles),
|
| 156 |
+
dynamic_obstacles=to_position_list(self._dynamic_obstacles),
|
| 157 |
+
traffic_zones=traffic_zones,
|
| 158 |
+
charging_stations=to_position_list(self._charging_stations),
|
| 159 |
+
battery_level=self._state.battery_level,
|
| 160 |
+
step_count=self._state.step_count,
|
| 161 |
+
total_reward=total_reward,
|
| 162 |
+
)
|
| 163 |
+
|
| 164 |
+
def apply_action(self, action: SimulatorAction) -> Tuple[float, StepInfo]:
|
| 165 |
+
if self._state.done:
|
| 166 |
+
return 0.0, StepInfo(invalid_action=True, message="episode_already_done")
|
| 167 |
+
|
| 168 |
+
reward = self.reward_config.step_penalty
|
| 169 |
+
info = StepInfo()
|
| 170 |
+
self._state.step_count += 1
|
| 171 |
+
|
| 172 |
+
if action.action_type == ActionType.MOVE:
|
| 173 |
+
delta_reward, step_info = self._handle_move(action.direction)
|
| 174 |
+
reward += delta_reward
|
| 175 |
+
info = step_info
|
| 176 |
+
elif action.action_type == ActionType.ACCEPT_ORDER:
|
| 177 |
+
delta_reward, step_info = self._handle_accept(action.order_id)
|
| 178 |
+
reward += delta_reward
|
| 179 |
+
info = step_info
|
| 180 |
+
elif action.action_type == ActionType.DELIVER_ORDER:
|
| 181 |
+
delta_reward, step_info = self._handle_deliver()
|
| 182 |
+
reward += delta_reward
|
| 183 |
+
info = step_info
|
| 184 |
+
elif action.action_type == ActionType.WAIT:
|
| 185 |
+
delta_reward, step_info = self._handle_wait()
|
| 186 |
+
reward += delta_reward
|
| 187 |
+
info = step_info
|
| 188 |
+
else:
|
| 189 |
+
reward += self.reward_config.invalid_action_penalty
|
| 190 |
+
info = StepInfo(invalid_action=True, message="unknown_action")
|
| 191 |
+
|
| 192 |
+
if self.config.dynamic_obstacles_enabled and self._dynamic_obstacles:
|
| 193 |
+
self._update_dynamic_obstacles()
|
| 194 |
+
|
| 195 |
+
delay_penalty = self._compute_delay_penalty()
|
| 196 |
+
if delay_penalty < 0:
|
| 197 |
+
reward += delay_penalty
|
| 198 |
+
info.delay_penalty_applied = True
|
| 199 |
+
|
| 200 |
+
if self._state.battery_level is not None and self._state.battery_level <= 0:
|
| 201 |
+
self._state.done = True
|
| 202 |
+
reward += self.reward_config.battery_failure_penalty
|
| 203 |
+
info.battery_depleted = True
|
| 204 |
+
|
| 205 |
+
if self._state.step_count >= self.config.max_steps:
|
| 206 |
+
self._state.done = True
|
| 207 |
+
|
| 208 |
+
if self._all_delivered():
|
| 209 |
+
self._state.done = True
|
| 210 |
+
|
| 211 |
+
return reward, info
|
| 212 |
+
|
| 213 |
+
def _all_delivered(self) -> bool:
|
| 214 |
+
return all(order.status == OrderStatus.DELIVERED for order in self._state.orders.values())
|
| 215 |
+
|
| 216 |
+
def _generate_map(self) -> None:
|
| 217 |
+
blocked: Set[Coordinate] = {(0, 0)}.union(self._charging_stations)
|
| 218 |
+
|
| 219 |
+
obstacle_count = int(self.config.width * self.config.height * self.config.obstacle_density)
|
| 220 |
+
traffic_count = int(self.config.width * self.config.height * self.config.traffic_density)
|
| 221 |
+
|
| 222 |
+
obstacle_locations = set(
|
| 223 |
+
sample_locations(
|
| 224 |
+
self.rng,
|
| 225 |
+
self.config.width,
|
| 226 |
+
self.config.height,
|
| 227 |
+
obstacle_count,
|
| 228 |
+
blocked=blocked,
|
| 229 |
+
)
|
| 230 |
+
)
|
| 231 |
+
|
| 232 |
+
if self.config.dynamic_obstacles_enabled and obstacle_locations:
|
| 233 |
+
dynamic_count = int(len(obstacle_locations) * self.config.dynamic_obstacle_ratio)
|
| 234 |
+
dynamic_count = min(dynamic_count, len(obstacle_locations))
|
| 235 |
+
dynamic_items = self.rng.sample(sorted(obstacle_locations), dynamic_count) if dynamic_count > 0 else []
|
| 236 |
+
self._dynamic_obstacles = set(dynamic_items)
|
| 237 |
+
else:
|
| 238 |
+
self._dynamic_obstacles = set()
|
| 239 |
+
|
| 240 |
+
self._static_obstacles = obstacle_locations - self._dynamic_obstacles
|
| 241 |
+
self._obstacles = self._static_obstacles.union(self._dynamic_obstacles)
|
| 242 |
+
blocked = blocked.union(self._obstacles)
|
| 243 |
+
|
| 244 |
+
traffic_locations = sample_locations(
|
| 245 |
+
self.rng,
|
| 246 |
+
self.config.width,
|
| 247 |
+
self.config.height,
|
| 248 |
+
traffic_count,
|
| 249 |
+
blocked=blocked,
|
| 250 |
+
)
|
| 251 |
+
self._traffic = {loc: self.config.traffic_extra_cost for loc in traffic_locations}
|
| 252 |
+
|
| 253 |
+
def _generate_orders(self) -> Dict[str, Order]:
|
| 254 |
+
blocked = {(0, 0)}.union(self._charging_stations).union(self._obstacles)
|
| 255 |
+
points_per_order = 1 + self.config.delivery_locations_per_order
|
| 256 |
+
required_locations = self.config.max_orders * points_per_order
|
| 257 |
+
|
| 258 |
+
order_points = sample_locations(
|
| 259 |
+
self.rng,
|
| 260 |
+
self.config.width,
|
| 261 |
+
self.config.height,
|
| 262 |
+
required_locations,
|
| 263 |
+
blocked=blocked,
|
| 264 |
+
)
|
| 265 |
+
|
| 266 |
+
orders: Dict[str, Order] = {}
|
| 267 |
+
index = 0
|
| 268 |
+
order_number = 1
|
| 269 |
+
while index < len(order_points):
|
| 270 |
+
if index + points_per_order > len(order_points):
|
| 271 |
+
break
|
| 272 |
+
|
| 273 |
+
pickup_point = order_points[index]
|
| 274 |
+
delivery_points = order_points[index + 1 : index + points_per_order]
|
| 275 |
+
if not delivery_points:
|
| 276 |
+
break
|
| 277 |
+
|
| 278 |
+
priority = (
|
| 279 |
+
OrderPriority.HIGH
|
| 280 |
+
if self.rng.random() < self.config.priority_high_ratio
|
| 281 |
+
else OrderPriority.LOW
|
| 282 |
+
)
|
| 283 |
+
|
| 284 |
+
order_id = f"order_{order_number}"
|
| 285 |
+
orders[order_id] = Order(
|
| 286 |
+
order_id=order_id,
|
| 287 |
+
pickup=Position.from_tuple(pickup_point),
|
| 288 |
+
dropoff=Position.from_tuple(delivery_points[0]),
|
| 289 |
+
delivery_locations=[Position.from_tuple(loc) for loc in delivery_points],
|
| 290 |
+
priority=priority,
|
| 291 |
+
created_step=0,
|
| 292 |
+
status=OrderStatus.PENDING,
|
| 293 |
+
)
|
| 294 |
+
|
| 295 |
+
index += points_per_order
|
| 296 |
+
order_number += 1
|
| 297 |
+
|
| 298 |
+
return orders
|
| 299 |
+
|
| 300 |
+
def _orders_from_scenario(self, scenario: ScenarioDefinition) -> Dict[str, Order]:
|
| 301 |
+
orders: Dict[str, Order] = {}
|
| 302 |
+
for index, scenario_order in enumerate(scenario.orders, start=1):
|
| 303 |
+
order_id = scenario_order.order_id or f"order_{index}"
|
| 304 |
+
if order_id in orders:
|
| 305 |
+
raise ValueError(f"Duplicate scenario order_id '{order_id}'")
|
| 306 |
+
|
| 307 |
+
delivery_locations = scenario_order.delivery_locations or [scenario_order.dropoff]
|
| 308 |
+
orders[order_id] = Order(
|
| 309 |
+
order_id=order_id,
|
| 310 |
+
pickup=scenario_order.pickup.model_copy(deep=True),
|
| 311 |
+
dropoff=scenario_order.dropoff.model_copy(deep=True),
|
| 312 |
+
delivery_locations=[item.model_copy(deep=True) for item in delivery_locations],
|
| 313 |
+
priority=scenario_order.priority,
|
| 314 |
+
created_step=scenario_order.created_step,
|
| 315 |
+
status=OrderStatus.PENDING,
|
| 316 |
+
)
|
| 317 |
+
|
| 318 |
+
return orders
|
| 319 |
+
|
| 320 |
+
def _update_dynamic_obstacles(self) -> None:
|
| 321 |
+
reserved = {as_tuple(self._state.agent_location)}.union(self._charging_stations)
|
| 322 |
+
for order in self._state.orders.values():
|
| 323 |
+
if order.status == OrderStatus.DELIVERED:
|
| 324 |
+
continue
|
| 325 |
+
reserved.add(as_tuple(order.pickup))
|
| 326 |
+
targets = order.delivery_locations or [order.dropoff]
|
| 327 |
+
reserved.update(as_tuple(target) for target in targets)
|
| 328 |
+
|
| 329 |
+
updated_dynamic: Set[Coordinate] = set()
|
| 330 |
+
|
| 331 |
+
for obstacle in sorted(self._dynamic_obstacles):
|
| 332 |
+
candidates = [
|
| 333 |
+
obstacle,
|
| 334 |
+
move_coordinate(obstacle, Direction.UP),
|
| 335 |
+
move_coordinate(obstacle, Direction.DOWN),
|
| 336 |
+
move_coordinate(obstacle, Direction.LEFT),
|
| 337 |
+
move_coordinate(obstacle, Direction.RIGHT),
|
| 338 |
+
]
|
| 339 |
+
|
| 340 |
+
valid_candidates: List[Coordinate] = []
|
| 341 |
+
for candidate in candidates:
|
| 342 |
+
if not in_bounds(candidate, self.config.width, self.config.height):
|
| 343 |
+
continue
|
| 344 |
+
if candidate in self._static_obstacles:
|
| 345 |
+
continue
|
| 346 |
+
if candidate in updated_dynamic:
|
| 347 |
+
continue
|
| 348 |
+
if candidate in reserved:
|
| 349 |
+
continue
|
| 350 |
+
valid_candidates.append(candidate)
|
| 351 |
+
|
| 352 |
+
if not valid_candidates:
|
| 353 |
+
updated_dynamic.add(obstacle)
|
| 354 |
+
continue
|
| 355 |
+
|
| 356 |
+
selected = valid_candidates[self.rng.randrange(len(valid_candidates))]
|
| 357 |
+
updated_dynamic.add(selected)
|
| 358 |
+
|
| 359 |
+
self._dynamic_obstacles = updated_dynamic
|
| 360 |
+
self._obstacles = self._static_obstacles.union(self._dynamic_obstacles)
|
| 361 |
+
|
| 362 |
+
def _consume_battery(self, amount: int) -> None:
|
| 363 |
+
if self._state.battery_level is None:
|
| 364 |
+
return
|
| 365 |
+
self._state.battery_level = max(0, self._state.battery_level - amount)
|
| 366 |
+
|
| 367 |
+
def _handle_move(self, direction: Optional[Direction]) -> Tuple[float, StepInfo]:
|
| 368 |
+
if direction is None:
|
| 369 |
+
return self.reward_config.invalid_action_penalty, StepInfo(
|
| 370 |
+
invalid_action=True,
|
| 371 |
+
message="direction_required",
|
| 372 |
+
)
|
| 373 |
+
|
| 374 |
+
current = as_tuple(self._state.agent_location)
|
| 375 |
+
progress_target = self._get_progress_target(current)
|
| 376 |
+
nxt = move_coordinate(current, direction)
|
| 377 |
+
if not in_bounds(nxt, self.config.width, self.config.height):
|
| 378 |
+
return self.reward_config.invalid_action_penalty, StepInfo(
|
| 379 |
+
invalid_action=True,
|
| 380 |
+
message="out_of_bounds",
|
| 381 |
+
)
|
| 382 |
+
|
| 383 |
+
if nxt in self._obstacles:
|
| 384 |
+
return self.reward_config.invalid_action_penalty, StepInfo(
|
| 385 |
+
invalid_action=True,
|
| 386 |
+
message="hit_obstacle",
|
| 387 |
+
)
|
| 388 |
+
|
| 389 |
+
self._state.agent_location = Position.from_tuple(nxt)
|
| 390 |
+
battery_cost = 1 + self._traffic.get(nxt, 0)
|
| 391 |
+
self._consume_battery(battery_cost)
|
| 392 |
+
|
| 393 |
+
progress_reward, progress_delta = self._progress_shaping_reward(
|
| 394 |
+
previous_location=current,
|
| 395 |
+
new_location=nxt,
|
| 396 |
+
target=progress_target,
|
| 397 |
+
)
|
| 398 |
+
|
| 399 |
+
info = StepInfo(message="moved", made_progress=progress_delta > 0)
|
| 400 |
+
reward = 0.0
|
| 401 |
+
|
| 402 |
+
if nxt in self._traffic:
|
| 403 |
+
reward += self.reward_config.traffic_penalty
|
| 404 |
+
info.message = "moved_in_traffic"
|
| 405 |
+
|
| 406 |
+
reward += progress_reward
|
| 407 |
+
|
| 408 |
+
destination_reached = self._check_destination_reached()
|
| 409 |
+
info.destination_reached = destination_reached
|
| 410 |
+
if destination_reached:
|
| 411 |
+
reward += self.reward_config.destination_reward
|
| 412 |
+
info.made_progress = True
|
| 413 |
+
|
| 414 |
+
return reward, info
|
| 415 |
+
|
| 416 |
+
def _check_destination_reached(self) -> bool:
|
| 417 |
+
if self._state.current_order_id is None:
|
| 418 |
+
return False
|
| 419 |
+
|
| 420 |
+
order = self._state.orders[self._state.current_order_id]
|
| 421 |
+
current = as_tuple(self._state.agent_location)
|
| 422 |
+
|
| 423 |
+
if order.status == OrderStatus.ACCEPTED and current == as_tuple(order.pickup):
|
| 424 |
+
order.status = OrderStatus.PICKED_UP
|
| 425 |
+
order.picked_up_step = self._state.step_count
|
| 426 |
+
return True
|
| 427 |
+
|
| 428 |
+
target_locations = order.delivery_locations or [order.dropoff]
|
| 429 |
+
if order.status == OrderStatus.PICKED_UP and any(current == as_tuple(target) for target in target_locations):
|
| 430 |
+
return True
|
| 431 |
+
|
| 432 |
+
return False
|
| 433 |
+
|
| 434 |
+
def _handle_wait(self) -> Tuple[float, StepInfo]:
|
| 435 |
+
info = StepInfo(message="wait", made_progress=False)
|
| 436 |
+
reward = 0.0
|
| 437 |
+
|
| 438 |
+
if self._state.battery_level is None:
|
| 439 |
+
return reward, info
|
| 440 |
+
|
| 441 |
+
location = as_tuple(self._state.agent_location)
|
| 442 |
+
if location not in self._charging_stations:
|
| 443 |
+
return reward, info
|
| 444 |
+
|
| 445 |
+
if self._state.battery_level >= self.config.battery_capacity:
|
| 446 |
+
info.message = "wait_full_battery"
|
| 447 |
+
return reward, info
|
| 448 |
+
|
| 449 |
+
before = self._state.battery_level
|
| 450 |
+
self._state.battery_level = min(
|
| 451 |
+
self.config.battery_capacity,
|
| 452 |
+
self._state.battery_level + self.config.battery_recharge_rate,
|
| 453 |
+
)
|
| 454 |
+
gained = self._state.battery_level - before
|
| 455 |
+
|
| 456 |
+
if gained > 0:
|
| 457 |
+
reward += self.reward_config.recharge_reward
|
| 458 |
+
info.message = "recharged"
|
| 459 |
+
info.made_progress = True
|
| 460 |
+
|
| 461 |
+
return reward, info
|
| 462 |
+
|
| 463 |
+
def _handle_accept(self, order_id: Optional[str]) -> Tuple[float, StepInfo]:
|
| 464 |
+
if self._state.current_order_id is not None:
|
| 465 |
+
return self.reward_config.invalid_action_penalty, StepInfo(
|
| 466 |
+
invalid_action=True,
|
| 467 |
+
message="current_order_already_assigned",
|
| 468 |
+
)
|
| 469 |
+
|
| 470 |
+
selected_id = order_id
|
| 471 |
+
if selected_id is None:
|
| 472 |
+
pending_orders = [
|
| 473 |
+
order
|
| 474 |
+
for order in self._state.orders.values()
|
| 475 |
+
if order.status == OrderStatus.PENDING
|
| 476 |
+
]
|
| 477 |
+
pending_orders.sort(
|
| 478 |
+
key=lambda order: (
|
| 479 |
+
0 if order.priority == OrderPriority.HIGH else 1,
|
| 480 |
+
order.order_id,
|
| 481 |
+
)
|
| 482 |
+
)
|
| 483 |
+
pending_ids = [order.order_id for order in pending_orders]
|
| 484 |
+
if not pending_ids:
|
| 485 |
+
return self.reward_config.invalid_action_penalty, StepInfo(
|
| 486 |
+
invalid_action=True,
|
| 487 |
+
message="no_pending_order",
|
| 488 |
+
)
|
| 489 |
+
selected_id = pending_ids[0]
|
| 490 |
+
|
| 491 |
+
order = self._state.orders.get(selected_id)
|
| 492 |
+
if order is None or order.status != OrderStatus.PENDING:
|
| 493 |
+
return self.reward_config.invalid_action_penalty, StepInfo(
|
| 494 |
+
invalid_action=True,
|
| 495 |
+
message="order_not_available",
|
| 496 |
+
)
|
| 497 |
+
|
| 498 |
+
order.status = OrderStatus.ACCEPTED
|
| 499 |
+
order.accepted_step = self._state.step_count
|
| 500 |
+
self._state.current_order_id = selected_id
|
| 501 |
+
self._consume_battery(1)
|
| 502 |
+
return 0.0, StepInfo(message="order_accepted", made_progress=True)
|
| 503 |
+
|
| 504 |
+
def _handle_deliver(self) -> Tuple[float, StepInfo]:
|
| 505 |
+
if self._state.current_order_id is None:
|
| 506 |
+
return self.reward_config.invalid_action_penalty, StepInfo(
|
| 507 |
+
invalid_action=True,
|
| 508 |
+
message="no_current_order",
|
| 509 |
+
)
|
| 510 |
+
|
| 511 |
+
order = self._state.orders[self._state.current_order_id]
|
| 512 |
+
current = as_tuple(self._state.agent_location)
|
| 513 |
+
if order.status == OrderStatus.ACCEPTED and current == as_tuple(order.pickup):
|
| 514 |
+
order.status = OrderStatus.PICKED_UP
|
| 515 |
+
order.picked_up_step = self._state.step_count
|
| 516 |
+
self._consume_battery(1)
|
| 517 |
+
return self.reward_config.destination_reward, StepInfo(
|
| 518 |
+
message="order_picked_up",
|
| 519 |
+
destination_reached=True,
|
| 520 |
+
made_progress=True,
|
| 521 |
+
)
|
| 522 |
+
|
| 523 |
+
valid_dropoffs = order.delivery_locations or [order.dropoff]
|
| 524 |
+
is_valid_dropoff = any(current == as_tuple(target) for target in valid_dropoffs)
|
| 525 |
+
if order.status != OrderStatus.PICKED_UP or not is_valid_dropoff:
|
| 526 |
+
return self.reward_config.invalid_action_penalty, StepInfo(
|
| 527 |
+
invalid_action=True,
|
| 528 |
+
message="not_at_dropoff_or_not_picked_up",
|
| 529 |
+
)
|
| 530 |
+
|
| 531 |
+
order.status = OrderStatus.DELIVERED
|
| 532 |
+
order.delivered_step = self._state.step_count
|
| 533 |
+
self._state.delivered_orders += 1
|
| 534 |
+
delivered_id = order.order_id
|
| 535 |
+
self._state.current_order_id = None
|
| 536 |
+
self._consume_battery(1)
|
| 537 |
+
|
| 538 |
+
priority_bonus = (
|
| 539 |
+
self.reward_config.high_priority_delivery_bonus
|
| 540 |
+
if order.priority == OrderPriority.HIGH
|
| 541 |
+
else self.reward_config.low_priority_delivery_bonus
|
| 542 |
+
)
|
| 543 |
+
|
| 544 |
+
efficiency_bonus = self._delivery_efficiency_bonus(order)
|
| 545 |
+
|
| 546 |
+
return self.reward_config.delivery_reward + priority_bonus + efficiency_bonus, StepInfo(
|
| 547 |
+
message=f"order_delivered_{order.priority.value}",
|
| 548 |
+
delivered_order_id=delivered_id,
|
| 549 |
+
made_progress=True,
|
| 550 |
+
)
|
| 551 |
+
|
| 552 |
+
def _get_progress_target(self, from_location: Coordinate) -> Optional[Coordinate]:
|
| 553 |
+
if self._state.current_order_id is not None:
|
| 554 |
+
current_order = self._state.orders[self._state.current_order_id]
|
| 555 |
+
if current_order.status == OrderStatus.ACCEPTED:
|
| 556 |
+
return as_tuple(current_order.pickup)
|
| 557 |
+
if current_order.status == OrderStatus.PICKED_UP:
|
| 558 |
+
targets = current_order.delivery_locations or [current_order.dropoff]
|
| 559 |
+
target_coords = [as_tuple(target) for target in targets]
|
| 560 |
+
return min(target_coords, key=lambda t: manhattan(from_location, t))
|
| 561 |
+
|
| 562 |
+
pending_orders = [
|
| 563 |
+
order
|
| 564 |
+
for order in self._state.orders.values()
|
| 565 |
+
if order.status == OrderStatus.PENDING
|
| 566 |
+
]
|
| 567 |
+
if not pending_orders:
|
| 568 |
+
return None
|
| 569 |
+
|
| 570 |
+
pickup_targets = [as_tuple(order.pickup) for order in pending_orders]
|
| 571 |
+
return min(pickup_targets, key=lambda t: manhattan(from_location, t))
|
| 572 |
+
|
| 573 |
+
def _progress_shaping_reward(
|
| 574 |
+
self,
|
| 575 |
+
previous_location: Coordinate,
|
| 576 |
+
new_location: Coordinate,
|
| 577 |
+
target: Optional[Coordinate],
|
| 578 |
+
) -> Tuple[float, int]:
|
| 579 |
+
if target is None:
|
| 580 |
+
return 0.0, 0
|
| 581 |
+
|
| 582 |
+
old_distance = manhattan(previous_location, target)
|
| 583 |
+
new_distance = manhattan(new_location, target)
|
| 584 |
+
distance_delta = old_distance - new_distance
|
| 585 |
+
|
| 586 |
+
raw_reward = self.reward_config.progress_reward_scale * distance_delta
|
| 587 |
+
clip_value = self.reward_config.progress_reward_clip
|
| 588 |
+
shaped_reward = max(-clip_value, min(clip_value, raw_reward))
|
| 589 |
+
return shaped_reward, distance_delta
|
| 590 |
+
|
| 591 |
+
def _delivery_efficiency_bonus(self, order: Order) -> float:
|
| 592 |
+
if order.accepted_step is None:
|
| 593 |
+
return 0.0
|
| 594 |
+
|
| 595 |
+
elapsed_steps = max(1, self._state.step_count - order.accepted_step)
|
| 596 |
+
delivery_targets = order.delivery_locations or [order.dropoff]
|
| 597 |
+
travel_lower_bound = min(
|
| 598 |
+
manhattan(as_tuple(order.pickup), as_tuple(target))
|
| 599 |
+
for target in delivery_targets
|
| 600 |
+
)
|
| 601 |
+
|
| 602 |
+
# Two control actions are minimally required: pickup and final deliver.
|
| 603 |
+
expected_steps = max(2, travel_lower_bound + 2)
|
| 604 |
+
efficiency_ratio = min(1.0, expected_steps / elapsed_steps)
|
| 605 |
+
return self.reward_config.efficiency_bonus_max * efficiency_ratio
|
| 606 |
+
|
| 607 |
+
def _compute_delay_penalty(self) -> float:
|
| 608 |
+
if self._state.current_order_id is None:
|
| 609 |
+
return 0.0
|
| 610 |
+
order = self._state.orders[self._state.current_order_id]
|
| 611 |
+
if order.accepted_step is None:
|
| 612 |
+
return 0.0
|
| 613 |
+
|
| 614 |
+
elapsed = self._state.step_count - order.accepted_step
|
| 615 |
+
threshold = 12 if order.priority == OrderPriority.HIGH else 20
|
| 616 |
+
overdue_steps = elapsed - threshold
|
| 617 |
+
if overdue_steps <= 0:
|
| 618 |
+
return 0.0
|
| 619 |
+
|
| 620 |
+
penalty = overdue_steps * self.reward_config.delay_penalty_scale
|
| 621 |
+
penalty = min(self.reward_config.delay_penalty_cap, penalty)
|
| 622 |
+
return -penalty
|
env/utils.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import random
|
| 4 |
+
from typing import Iterable, List, Optional, Sequence, Set, Tuple
|
| 5 |
+
|
| 6 |
+
from env.models import Coordinate, Direction, Position
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
def move_coordinate(location: Coordinate, direction: Direction) -> Coordinate:
|
| 10 |
+
x, y = location
|
| 11 |
+
if direction == Direction.UP:
|
| 12 |
+
return x, y - 1
|
| 13 |
+
if direction == Direction.DOWN:
|
| 14 |
+
return x, y + 1
|
| 15 |
+
if direction == Direction.LEFT:
|
| 16 |
+
return x - 1, y
|
| 17 |
+
return x + 1, y
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def in_bounds(location: Coordinate, width: int, height: int) -> bool:
|
| 21 |
+
return 0 <= location[0] < width and 0 <= location[1] < height
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def manhattan(a: Coordinate, b: Coordinate) -> int:
|
| 25 |
+
return abs(a[0] - b[0]) + abs(a[1] - b[1])
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def to_position_list(coordinates: Iterable[Coordinate]) -> List[Position]:
|
| 29 |
+
return [Position(x=x, y=y) for x, y in coordinates]
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def sample_locations(
|
| 33 |
+
rng: random.Random,
|
| 34 |
+
width: int,
|
| 35 |
+
height: int,
|
| 36 |
+
count: int,
|
| 37 |
+
blocked: Optional[Set[Coordinate]] = None,
|
| 38 |
+
) -> List[Coordinate]:
|
| 39 |
+
blocked = blocked or set()
|
| 40 |
+
available: List[Coordinate] = [
|
| 41 |
+
(x, y)
|
| 42 |
+
for x in range(width)
|
| 43 |
+
for y in range(height)
|
| 44 |
+
if (x, y) not in blocked
|
| 45 |
+
]
|
| 46 |
+
if count > len(available):
|
| 47 |
+
count = len(available)
|
| 48 |
+
return rng.sample(available, count)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def nearest_location(origin: Coordinate, candidates: Sequence[Coordinate]) -> Optional[Coordinate]:
|
| 52 |
+
if not candidates:
|
| 53 |
+
return None
|
| 54 |
+
return min(candidates, key=lambda c: manhattan(origin, c))
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def as_tuple(position: Position) -> Coordinate:
|
| 58 |
+
return position.x, position.y
|
grader/__init__.py
ADDED
|
File without changes
|
grader/grader.py
ADDED
|
@@ -0,0 +1,530 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from dataclasses import dataclass
|
| 4 |
+
from typing import Any, Dict, List, Sequence
|
| 5 |
+
|
| 6 |
+
from pydantic import BaseModel, Field
|
| 7 |
+
|
| 8 |
+
from env.models import Observation, Order, OrderPriority, StepResult
|
| 9 |
+
from grader.metrics import (
|
| 10 |
+
clamp,
|
| 11 |
+
round_score,
|
| 12 |
+
)
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
DEFAULT_SUCCESS_CONDITION: Dict[str, Any] = {
|
| 16 |
+
"completion_rate_min": 1.0,
|
| 17 |
+
"max_steps": 200,
|
| 18 |
+
"invalid_action_rate_max": 0.10,
|
| 19 |
+
}
|
| 20 |
+
DEFAULT_HIGH_PRIORITY_DEADLINE = 12
|
| 21 |
+
|
| 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)
|
| 28 |
+
high_priority_total_orders: int = Field(..., ge=0)
|
| 29 |
+
high_priority_delivered_orders: int = Field(..., ge=0)
|
| 30 |
+
high_priority_on_time_deliveries: int = Field(..., ge=0)
|
| 31 |
+
|
| 32 |
+
steps_taken: int = Field(..., ge=0)
|
| 33 |
+
max_steps_target: int = Field(..., ge=0)
|
| 34 |
+
optimal_steps: int = Field(..., ge=0)
|
| 35 |
+
|
| 36 |
+
completion_rate: float = Field(..., ge=0.0, le=1.0)
|
| 37 |
+
high_priority_on_time_rate: float = Field(..., ge=0.0, le=1.0)
|
| 38 |
+
efficiency_ratio: float = Field(..., ge=0.0, le=1.0)
|
| 39 |
+
invalid_action_rate: float = Field(..., ge=0.0, le=1.0)
|
| 40 |
+
|
| 41 |
+
completion_component: float = Field(..., ge=0.0, le=1.0)
|
| 42 |
+
priority_component: float = Field(..., ge=0.0, le=1.0)
|
| 43 |
+
efficiency_component: float = Field(..., ge=0.0, le=1.0)
|
| 44 |
+
penalty_component: float = Field(..., ge=0.0, le=1.0)
|
| 45 |
+
|
| 46 |
+
invalid_actions: int = Field(..., ge=0)
|
| 47 |
+
delay_events: int = Field(..., ge=0)
|
| 48 |
+
battery_depletion_events: int = Field(..., ge=0)
|
| 49 |
+
no_progress_events: int = Field(..., ge=0)
|
| 50 |
+
battery_remaining: int | None = None
|
| 51 |
+
|
| 52 |
+
success_condition_used: Dict[str, Any] = Field(default_factory=dict)
|
| 53 |
+
scoring_logic: str
|
| 54 |
+
safeguards_applied: List[str] = Field(default_factory=list)
|
| 55 |
+
edge_case_handling: List[str] = Field(default_factory=list)
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
@dataclass(frozen=True)
|
| 59 |
+
class EpisodeStats:
|
| 60 |
+
steps_taken: int
|
| 61 |
+
invalid_actions: int
|
| 62 |
+
delay_events: int
|
| 63 |
+
battery_depletion_events: int
|
| 64 |
+
no_progress_events: int
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
@dataclass(frozen=True)
|
| 68 |
+
class SuccessTargets:
|
| 69 |
+
completion_target: float
|
| 70 |
+
priority_target: float | None
|
| 71 |
+
max_steps_target: int
|
| 72 |
+
invalid_action_rate_max: float
|
| 73 |
+
require_no_battery_depletion: bool
|
| 74 |
+
high_priority_deadline_steps: int
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
class DeliveryEpisodeGrader:
|
| 78 |
+
"""Deterministic task-aware grader aligned with task success_condition metrics."""
|
| 79 |
+
|
| 80 |
+
def __init__(self) -> None:
|
| 81 |
+
self.base_weights = {
|
| 82 |
+
"completion": 0.55,
|
| 83 |
+
"priority": 0.25,
|
| 84 |
+
"efficiency": 0.20,
|
| 85 |
+
}
|
| 86 |
+
|
| 87 |
+
def grade_episode(
|
| 88 |
+
self,
|
| 89 |
+
trajectory: Sequence[StepResult],
|
| 90 |
+
final_observation: Observation,
|
| 91 |
+
success_condition: Dict[str, Any] | None = None,
|
| 92 |
+
) -> GradeReport:
|
| 93 |
+
order_catalog = self._collect_order_catalog(trajectory, final_observation)
|
| 94 |
+
delivered_steps = self._collect_delivered_steps(trajectory)
|
| 95 |
+
delivered_ids = set(delivered_steps.keys())
|
| 96 |
+
delivered_orders = len(delivered_ids)
|
| 97 |
+
|
| 98 |
+
remaining_ids = {order.order_id for order in final_observation.pending_orders}
|
| 99 |
+
if final_observation.current_order is not None:
|
| 100 |
+
remaining_ids.add(final_observation.current_order.order_id)
|
| 101 |
+
|
| 102 |
+
known_ids = set(order_catalog.keys())
|
| 103 |
+
all_ids = known_ids.union(delivered_ids).union(remaining_ids)
|
| 104 |
+
total_orders = len(all_ids)
|
| 105 |
+
|
| 106 |
+
episode_stats = self._collect_episode_stats(trajectory)
|
| 107 |
+
|
| 108 |
+
edge_case_notes: List[str] = []
|
| 109 |
+
safeguards: List[str] = []
|
| 110 |
+
|
| 111 |
+
targets = self._resolve_success_targets(
|
| 112 |
+
success_condition=success_condition,
|
| 113 |
+
fallback_max_steps=max(1, final_observation.step_count),
|
| 114 |
+
edge_case_notes=edge_case_notes,
|
| 115 |
+
)
|
| 116 |
+
|
| 117 |
+
if episode_stats.steps_taken == 0:
|
| 118 |
+
edge_case_notes.append("No steps were recorded for this episode.")
|
| 119 |
+
|
| 120 |
+
if total_orders == 0:
|
| 121 |
+
completion_rate_value = 1.0
|
| 122 |
+
edge_case_notes.append("No orders were present; completion_rate set to 1.0 by convention.")
|
| 123 |
+
else:
|
| 124 |
+
completion_rate_value = delivered_orders / total_orders
|
| 125 |
+
|
| 126 |
+
(
|
| 127 |
+
high_priority_total,
|
| 128 |
+
high_priority_delivered,
|
| 129 |
+
high_priority_on_time,
|
| 130 |
+
high_priority_on_time_rate,
|
| 131 |
+
) = self._high_priority_metrics(
|
| 132 |
+
order_catalog=order_catalog,
|
| 133 |
+
delivered_steps=delivered_steps,
|
| 134 |
+
deadline_steps=targets.high_priority_deadline_steps,
|
| 135 |
+
edge_case_notes=edge_case_notes,
|
| 136 |
+
)
|
| 137 |
+
|
| 138 |
+
efficiency_ratio = self._efficiency_ratio(
|
| 139 |
+
steps_taken=episode_stats.steps_taken,
|
| 140 |
+
max_steps_target=targets.max_steps_target,
|
| 141 |
+
)
|
| 142 |
+
|
| 143 |
+
invalid_action_rate = self._invalid_action_rate(
|
| 144 |
+
invalid_actions=episode_stats.invalid_actions,
|
| 145 |
+
steps_taken=episode_stats.steps_taken,
|
| 146 |
+
)
|
| 147 |
+
|
| 148 |
+
completion_component = self._target_progress(
|
| 149 |
+
value=completion_rate_value,
|
| 150 |
+
target=targets.completion_target,
|
| 151 |
+
edge_case_notes=edge_case_notes,
|
| 152 |
+
metric_name="completion_rate",
|
| 153 |
+
)
|
| 154 |
+
|
| 155 |
+
if targets.priority_target is None:
|
| 156 |
+
priority_component = 1.0
|
| 157 |
+
edge_case_notes.append("high_priority_on_time_rate_min not set; priority component treated as neutral.")
|
| 158 |
+
else:
|
| 159 |
+
priority_component = self._target_progress(
|
| 160 |
+
value=high_priority_on_time_rate,
|
| 161 |
+
target=targets.priority_target,
|
| 162 |
+
edge_case_notes=edge_case_notes,
|
| 163 |
+
metric_name="high_priority_on_time_rate",
|
| 164 |
+
)
|
| 165 |
+
|
| 166 |
+
efficiency_component = efficiency_ratio
|
| 167 |
+
penalty_component = self._invalid_action_penalty_component(
|
| 168 |
+
invalid_action_rate=invalid_action_rate,
|
| 169 |
+
invalid_action_rate_max=targets.invalid_action_rate_max,
|
| 170 |
+
edge_case_notes=edge_case_notes,
|
| 171 |
+
)
|
| 172 |
+
|
| 173 |
+
if targets.require_no_battery_depletion and episode_stats.battery_depletion_events > 0:
|
| 174 |
+
penalty_component *= 0.5
|
| 175 |
+
safeguards.append("battery_depletion_penalty_applied")
|
| 176 |
+
|
| 177 |
+
completion_weight = self.base_weights["completion"]
|
| 178 |
+
priority_weight = self.base_weights["priority"] if targets.priority_target is not None else 0.0
|
| 179 |
+
efficiency_weight = self.base_weights["efficiency"]
|
| 180 |
+
if targets.priority_target is None:
|
| 181 |
+
completion_weight += 0.15
|
| 182 |
+
efficiency_weight += 0.10
|
| 183 |
+
|
| 184 |
+
raw_base = (
|
| 185 |
+
completion_weight * completion_component
|
| 186 |
+
+ priority_weight * priority_component
|
| 187 |
+
+ efficiency_weight * efficiency_component
|
| 188 |
+
)
|
| 189 |
+
raw_score = raw_base * penalty_component
|
| 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. "
|
| 196 |
+
"components: completion (highest), priority SLA (medium when configured), "
|
| 197 |
+
"efficiency from steps/max_steps (lower), and penalties reduce via invalid_action_rate."
|
| 198 |
+
)
|
| 199 |
+
|
| 200 |
+
if not edge_case_notes:
|
| 201 |
+
edge_case_notes.append("No edge-case adjustments were needed.")
|
| 202 |
+
|
| 203 |
+
success_condition_used = {
|
| 204 |
+
"completion_rate_target": targets.completion_target,
|
| 205 |
+
"high_priority_on_time_rate_target": targets.priority_target,
|
| 206 |
+
"max_steps": targets.max_steps_target,
|
| 207 |
+
"invalid_action_rate_max": targets.invalid_action_rate_max,
|
| 208 |
+
"battery_depletion": not targets.require_no_battery_depletion,
|
| 209 |
+
"high_priority_deadline_steps": targets.high_priority_deadline_steps,
|
| 210 |
+
}
|
| 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,
|
| 217 |
+
high_priority_delivered_orders=high_priority_delivered,
|
| 218 |
+
high_priority_on_time_deliveries=high_priority_on_time,
|
| 219 |
+
steps_taken=episode_stats.steps_taken,
|
| 220 |
+
max_steps_target=targets.max_steps_target,
|
| 221 |
+
optimal_steps=targets.max_steps_target,
|
| 222 |
+
completion_rate=completion_rate_value,
|
| 223 |
+
high_priority_on_time_rate=high_priority_on_time_rate,
|
| 224 |
+
efficiency_ratio=efficiency_ratio,
|
| 225 |
+
invalid_action_rate=invalid_action_rate,
|
| 226 |
+
completion_component=completion_component,
|
| 227 |
+
priority_component=priority_component,
|
| 228 |
+
efficiency_component=efficiency_component,
|
| 229 |
+
penalty_component=penalty_component,
|
| 230 |
+
invalid_actions=episode_stats.invalid_actions,
|
| 231 |
+
delay_events=episode_stats.delay_events,
|
| 232 |
+
battery_depletion_events=episode_stats.battery_depletion_events,
|
| 233 |
+
no_progress_events=episode_stats.no_progress_events,
|
| 234 |
+
battery_remaining=final_observation.battery_level,
|
| 235 |
+
success_condition_used=success_condition_used,
|
| 236 |
+
scoring_logic=logic,
|
| 237 |
+
safeguards_applied=safeguards,
|
| 238 |
+
edge_case_handling=edge_case_notes,
|
| 239 |
+
)
|
| 240 |
+
|
| 241 |
+
def _resolve_success_targets(
|
| 242 |
+
self,
|
| 243 |
+
success_condition: Dict[str, Any] | None,
|
| 244 |
+
fallback_max_steps: int,
|
| 245 |
+
edge_case_notes: List[str],
|
| 246 |
+
) -> SuccessTargets:
|
| 247 |
+
condition = dict(DEFAULT_SUCCESS_CONDITION)
|
| 248 |
+
if success_condition is not None:
|
| 249 |
+
condition.update(success_condition)
|
| 250 |
+
|
| 251 |
+
completion_raw = condition.get("completion_rate")
|
| 252 |
+
if completion_raw is None:
|
| 253 |
+
completion_raw = condition.get("completion_rate_min", 1.0)
|
| 254 |
+
|
| 255 |
+
completion_target = self._safe_ratio_target(
|
| 256 |
+
raw_value=completion_raw,
|
| 257 |
+
default=1.0,
|
| 258 |
+
metric_name="completion_rate_target",
|
| 259 |
+
edge_case_notes=edge_case_notes,
|
| 260 |
+
)
|
| 261 |
+
|
| 262 |
+
priority_raw = condition.get("high_priority_on_time_rate_min")
|
| 263 |
+
priority_target: float | None = None
|
| 264 |
+
if priority_raw is not None:
|
| 265 |
+
priority_target = self._safe_ratio_target(
|
| 266 |
+
raw_value=priority_raw,
|
| 267 |
+
default=1.0,
|
| 268 |
+
metric_name="high_priority_on_time_rate_target",
|
| 269 |
+
edge_case_notes=edge_case_notes,
|
| 270 |
+
)
|
| 271 |
+
|
| 272 |
+
max_steps_raw = condition.get("max_steps", fallback_max_steps)
|
| 273 |
+
try:
|
| 274 |
+
max_steps_target = int(max_steps_raw)
|
| 275 |
+
except Exception:
|
| 276 |
+
max_steps_target = fallback_max_steps
|
| 277 |
+
edge_case_notes.append("Invalid max_steps in success_condition; fallback max_steps used.")
|
| 278 |
+
if max_steps_target <= 0:
|
| 279 |
+
max_steps_target = max(1, fallback_max_steps)
|
| 280 |
+
edge_case_notes.append("Non-positive max_steps in success_condition; fallback max_steps used.")
|
| 281 |
+
|
| 282 |
+
invalid_rate_raw = condition.get("invalid_action_rate_max", 0.10)
|
| 283 |
+
invalid_action_rate_max = self._safe_ratio_target(
|
| 284 |
+
raw_value=invalid_rate_raw,
|
| 285 |
+
default=0.10,
|
| 286 |
+
metric_name="invalid_action_rate_max",
|
| 287 |
+
edge_case_notes=edge_case_notes,
|
| 288 |
+
)
|
| 289 |
+
|
| 290 |
+
battery_depletion_rule = condition.get("battery_depletion")
|
| 291 |
+
require_no_battery_depletion = battery_depletion_rule is False
|
| 292 |
+
|
| 293 |
+
deadline_raw = condition.get("high_priority_deadline_steps", DEFAULT_HIGH_PRIORITY_DEADLINE)
|
| 294 |
+
try:
|
| 295 |
+
deadline_steps = int(deadline_raw)
|
| 296 |
+
except Exception:
|
| 297 |
+
deadline_steps = DEFAULT_HIGH_PRIORITY_DEADLINE
|
| 298 |
+
edge_case_notes.append("Invalid high_priority_deadline_steps; default deadline used.")
|
| 299 |
+
if deadline_steps <= 0:
|
| 300 |
+
deadline_steps = DEFAULT_HIGH_PRIORITY_DEADLINE
|
| 301 |
+
edge_case_notes.append("Non-positive high_priority_deadline_steps; default deadline used.")
|
| 302 |
+
|
| 303 |
+
return SuccessTargets(
|
| 304 |
+
completion_target=completion_target,
|
| 305 |
+
priority_target=priority_target,
|
| 306 |
+
max_steps_target=max_steps_target,
|
| 307 |
+
invalid_action_rate_max=invalid_action_rate_max,
|
| 308 |
+
require_no_battery_depletion=require_no_battery_depletion,
|
| 309 |
+
high_priority_deadline_steps=deadline_steps,
|
| 310 |
+
)
|
| 311 |
+
|
| 312 |
+
def _safe_ratio_target(
|
| 313 |
+
self,
|
| 314 |
+
raw_value: Any,
|
| 315 |
+
default: float,
|
| 316 |
+
metric_name: str,
|
| 317 |
+
edge_case_notes: List[str],
|
| 318 |
+
) -> float:
|
| 319 |
+
try:
|
| 320 |
+
parsed = float(raw_value)
|
| 321 |
+
except Exception:
|
| 322 |
+
edge_case_notes.append(f"Invalid {metric_name}; default value used.")
|
| 323 |
+
return default
|
| 324 |
+
return clamp(parsed, 0.0, 1.0)
|
| 325 |
+
|
| 326 |
+
def _collect_episode_stats(self, trajectory: Sequence[StepResult]) -> EpisodeStats:
|
| 327 |
+
return EpisodeStats(
|
| 328 |
+
steps_taken=len(trajectory),
|
| 329 |
+
invalid_actions=sum(1 for item in trajectory if item.info.invalid_action),
|
| 330 |
+
delay_events=sum(1 for item in trajectory if item.info.delay_penalty_applied),
|
| 331 |
+
battery_depletion_events=sum(1 for item in trajectory if item.info.battery_depleted),
|
| 332 |
+
no_progress_events=sum(1 for item in trajectory if item.info.made_progress is False),
|
| 333 |
+
)
|
| 334 |
+
|
| 335 |
+
def _collect_delivered_steps(self, trajectory: Sequence[StepResult]) -> Dict[str, int]:
|
| 336 |
+
delivered_steps: Dict[str, int] = {}
|
| 337 |
+
for index, item in enumerate(trajectory, start=1):
|
| 338 |
+
delivered_id = item.info.delivered_order_id
|
| 339 |
+
if delivered_id is None:
|
| 340 |
+
continue
|
| 341 |
+
delivered_step = item.observation.step_count if item.observation.step_count > 0 else index
|
| 342 |
+
delivered_steps[delivered_id] = delivered_step
|
| 343 |
+
return delivered_steps
|
| 344 |
+
|
| 345 |
+
def _high_priority_metrics(
|
| 346 |
+
self,
|
| 347 |
+
order_catalog: Dict[str, Order],
|
| 348 |
+
delivered_steps: Dict[str, int],
|
| 349 |
+
deadline_steps: int,
|
| 350 |
+
edge_case_notes: List[str],
|
| 351 |
+
) -> tuple[int, int, int, float]:
|
| 352 |
+
high_priority_orders = [
|
| 353 |
+
order
|
| 354 |
+
for order in order_catalog.values()
|
| 355 |
+
if order.priority == OrderPriority.HIGH
|
| 356 |
+
]
|
| 357 |
+
high_priority_total = len(high_priority_orders)
|
| 358 |
+
|
| 359 |
+
if high_priority_total == 0:
|
| 360 |
+
edge_case_notes.append("No high-priority orders were present; high-priority SLA treated as 1.0.")
|
| 361 |
+
return 0, 0, 0, 1.0
|
| 362 |
+
|
| 363 |
+
high_priority_delivered = 0
|
| 364 |
+
high_priority_on_time = 0
|
| 365 |
+
|
| 366 |
+
for order in high_priority_orders:
|
| 367 |
+
delivered_step = delivered_steps.get(order.order_id)
|
| 368 |
+
if delivered_step is None:
|
| 369 |
+
continue
|
| 370 |
+
|
| 371 |
+
high_priority_delivered += 1
|
| 372 |
+
if order.accepted_step is None:
|
| 373 |
+
continue
|
| 374 |
+
|
| 375 |
+
if delivered_step - order.accepted_step <= deadline_steps:
|
| 376 |
+
high_priority_on_time += 1
|
| 377 |
+
|
| 378 |
+
on_time_rate = high_priority_on_time / high_priority_total
|
| 379 |
+
return high_priority_total, high_priority_delivered, high_priority_on_time, on_time_rate
|
| 380 |
+
|
| 381 |
+
def _efficiency_ratio(self, steps_taken: int, max_steps_target: int) -> float:
|
| 382 |
+
if max_steps_target <= 0:
|
| 383 |
+
return 1.0
|
| 384 |
+
if steps_taken <= 0:
|
| 385 |
+
return 1.0
|
| 386 |
+
return clamp((max_steps_target - steps_taken) / max_steps_target, 0.0, 1.0)
|
| 387 |
+
|
| 388 |
+
def _invalid_action_rate(self, invalid_actions: int, steps_taken: int) -> float:
|
| 389 |
+
if steps_taken <= 0:
|
| 390 |
+
return 0.0
|
| 391 |
+
return clamp(invalid_actions / steps_taken, 0.0, 1.0)
|
| 392 |
+
|
| 393 |
+
def _target_progress(
|
| 394 |
+
self,
|
| 395 |
+
value: float,
|
| 396 |
+
target: float,
|
| 397 |
+
edge_case_notes: List[str],
|
| 398 |
+
metric_name: str,
|
| 399 |
+
) -> float:
|
| 400 |
+
if target <= 0.0:
|
| 401 |
+
edge_case_notes.append(f"{metric_name} target was 0.0; component treated as 1.0.")
|
| 402 |
+
return 1.0
|
| 403 |
+
return clamp(value / target, 0.0, 1.0)
|
| 404 |
+
|
| 405 |
+
def _invalid_action_penalty_component(
|
| 406 |
+
self,
|
| 407 |
+
invalid_action_rate: float,
|
| 408 |
+
invalid_action_rate_max: float,
|
| 409 |
+
edge_case_notes: List[str],
|
| 410 |
+
) -> float:
|
| 411 |
+
if invalid_action_rate_max <= 0.0:
|
| 412 |
+
if invalid_action_rate > 0.0:
|
| 413 |
+
edge_case_notes.append(
|
| 414 |
+
"invalid_action_rate_max is 0.0 and invalid actions occurred; penalty set to 0.0."
|
| 415 |
+
)
|
| 416 |
+
return 0.0
|
| 417 |
+
return 1.0
|
| 418 |
+
|
| 419 |
+
normalized = invalid_action_rate / invalid_action_rate_max
|
| 420 |
+
if invalid_action_rate <= invalid_action_rate_max:
|
| 421 |
+
return clamp(1.0 - 0.25 * normalized, 0.75, 1.0)
|
| 422 |
+
|
| 423 |
+
overflow = (invalid_action_rate - invalid_action_rate_max) / max(1.0 - invalid_action_rate_max, 1e-9)
|
| 424 |
+
return clamp(0.75 - 0.75 * overflow, 0.0, 0.75)
|
| 425 |
+
if total_orders == 0:
|
| 426 |
+
edge_case_notes.append("No orders were present; completion is evaluated as neutral.")
|
| 427 |
+
if penalty_stats.steps_taken == 0 and total_orders > 0:
|
| 428 |
+
edge_case_notes.append("No steps were recorded for a non-empty episode.")
|
| 429 |
+
|
| 430 |
+
unknown_orders = max(0, total_orders - len(order_catalog))
|
| 431 |
+
if unknown_orders > 0:
|
| 432 |
+
edge_case_notes.append(
|
| 433 |
+
"Some order geometries were missing from observations; overhead-only fallback used."
|
| 434 |
+
)
|
| 435 |
+
|
| 436 |
+
ordered_orders = [order_catalog[key] for key in sorted(order_catalog.keys())]
|
| 437 |
+
optimal_steps_known = optimal_steps_single_agent(
|
| 438 |
+
orders=ordered_orders,
|
| 439 |
+
start_location=self.start_location,
|
| 440 |
+
)
|
| 441 |
+
optimal_steps = optimal_steps_known + (unknown_orders * 3)
|
| 442 |
+
|
| 443 |
+
if total_orders > 0 and optimal_steps == 0:
|
| 444 |
+
optimal_steps = max(1, penalty_stats.steps_taken)
|
| 445 |
+
safeguards.append("optimal_steps_zero_guard")
|
| 446 |
+
|
| 447 |
+
delivered_weight, total_weight = self._weighted_completion_masses(
|
| 448 |
+
delivered_ids=delivered_ids,
|
| 449 |
+
all_ids=all_ids,
|
| 450 |
+
order_catalog=order_catalog,
|
| 451 |
+
)
|
| 452 |
+
|
| 453 |
+
completion, efficiency_component, penalty_component = self._compute_components(
|
| 454 |
+
delivered_weight=delivered_weight,
|
| 455 |
+
total_weight=total_weight,
|
| 456 |
+
penalty_stats=penalty_stats,
|
| 457 |
+
optimal_steps=optimal_steps,
|
| 458 |
+
)
|
| 459 |
+
|
| 460 |
+
raw_score = self._combine_components(
|
| 461 |
+
completion=completion,
|
| 462 |
+
efficiency_component=efficiency_component,
|
| 463 |
+
penalty_component=penalty_component,
|
| 464 |
+
)
|
| 465 |
+
|
| 466 |
+
raw_score = self._apply_safeguards(
|
| 467 |
+
raw_score=raw_score,
|
| 468 |
+
safeguards=safeguards,
|
| 469 |
+
total_orders=total_orders,
|
| 470 |
+
delivered_orders=delivered_orders,
|
| 471 |
+
steps_taken=penalty_stats.steps_taken,
|
| 472 |
+
invalid_actions=penalty_stats.invalid_actions,
|
| 473 |
+
delay_events=penalty_stats.delay_events,
|
| 474 |
+
battery_depletion_events=penalty_stats.battery_depletion_events,
|
| 475 |
+
no_progress_events=penalty_stats.no_progress_events,
|
| 476 |
+
completion=completion,
|
| 477 |
+
efficiency_component=efficiency_component,
|
| 478 |
+
penalty_component=penalty_component,
|
| 479 |
+
)
|
| 480 |
+
|
| 481 |
+
score = round_score(raw_score, decimals=4)
|
| 482 |
+
|
| 483 |
+
logic = (
|
| 484 |
+
"score = 0.50*completion + 0.30*efficiency + 0.20*penalty_quality; "
|
| 485 |
+
"completion is weighted by order priority (high>low), "
|
| 486 |
+
"efficiency uses optimal_steps/steps_taken and is gated by completion, "
|
| 487 |
+
"penalty_quality decreases with invalid, delay, battery-depletion, and sustained no-progress events."
|
| 488 |
+
)
|
| 489 |
+
|
| 490 |
+
if not edge_case_notes:
|
| 491 |
+
edge_case_notes.append("No edge-case adjustments were needed.")
|
| 492 |
+
|
| 493 |
+
return GradeReport(
|
| 494 |
+
score=score,
|
| 495 |
+
delivered_orders=delivered_orders,
|
| 496 |
+
total_orders=total_orders,
|
| 497 |
+
steps_taken=penalty_stats.steps_taken,
|
| 498 |
+
optimal_steps=optimal_steps,
|
| 499 |
+
completion_component=completion,
|
| 500 |
+
efficiency_component=efficiency_component,
|
| 501 |
+
penalty_component=penalty_component,
|
| 502 |
+
invalid_actions=penalty_stats.invalid_actions,
|
| 503 |
+
delay_events=penalty_stats.delay_events,
|
| 504 |
+
battery_depletion_events=penalty_stats.battery_depletion_events,
|
| 505 |
+
no_progress_events=penalty_stats.no_progress_events,
|
| 506 |
+
battery_remaining=final_observation.battery_level,
|
| 507 |
+
scoring_logic=logic,
|
| 508 |
+
safeguards_applied=safeguards,
|
| 509 |
+
edge_case_handling=edge_case_notes,
|
| 510 |
+
)
|
| 511 |
+
|
| 512 |
+
def _collect_order_catalog(
|
| 513 |
+
self,
|
| 514 |
+
trajectory: Sequence[StepResult],
|
| 515 |
+
final_observation: Observation,
|
| 516 |
+
) -> Dict[str, Order]:
|
| 517 |
+
orders: Dict[str, Order] = {}
|
| 518 |
+
|
| 519 |
+
for item in trajectory:
|
| 520 |
+
for order in item.observation.pending_orders:
|
| 521 |
+
orders[order.order_id] = order
|
| 522 |
+
if item.observation.current_order is not None:
|
| 523 |
+
orders[item.observation.current_order.order_id] = item.observation.current_order
|
| 524 |
+
|
| 525 |
+
for order in final_observation.pending_orders:
|
| 526 |
+
orders[order.order_id] = order
|
| 527 |
+
if final_observation.current_order is not None:
|
| 528 |
+
orders[final_observation.current_order.order_id] = final_observation.current_order
|
| 529 |
+
|
| 530 |
+
return orders
|
grader/metrics.py
ADDED
|
@@ -0,0 +1,178 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from decimal import ROUND_HALF_UP, Decimal
|
| 4 |
+
from typing import List, Sequence, Tuple
|
| 5 |
+
|
| 6 |
+
from env.models import Order
|
| 7 |
+
from env.utils import manhattan
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def clamp(value: float, low: float, high: float) -> float:
|
| 11 |
+
return max(low, min(high, value))
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def round_score(value: float, decimals: int = 4) -> float:
|
| 15 |
+
"""Round with HALF_UP behavior for stable scoring outputs."""
|
| 16 |
+
safe = clamp(value, 0.0, 1.0)
|
| 17 |
+
|
| 18 |
+
if decimals > 0:
|
| 19 |
+
epsilon = float(Decimal("1") / (Decimal(10) ** decimals))
|
| 20 |
+
if safe <= 0.0:
|
| 21 |
+
safe = epsilon
|
| 22 |
+
elif safe >= 1.0:
|
| 23 |
+
safe = 1.0 - epsilon
|
| 24 |
+
|
| 25 |
+
quantizer = Decimal("1") if decimals <= 0 else Decimal(f"1.{'0' * decimals}")
|
| 26 |
+
return float(Decimal(str(safe)).quantize(quantizer, rounding=ROUND_HALF_UP))
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def completion_rate(delivered_weight: float, total_weight: float) -> float:
|
| 30 |
+
if total_weight <= 0:
|
| 31 |
+
# Neutral when no work was available.
|
| 32 |
+
return 0.5
|
| 33 |
+
return clamp(delivered_weight / total_weight, 0.0, 1.0)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def efficiency_score(steps_taken: int, optimal_steps: int) -> float:
|
| 37 |
+
"""Compute bounded efficiency ratio from optimal-step lower bound."""
|
| 38 |
+
if optimal_steps <= 0:
|
| 39 |
+
return 1.0 if steps_taken <= 0 else 0.0
|
| 40 |
+
if steps_taken <= 0:
|
| 41 |
+
return 0.0
|
| 42 |
+
return clamp(optimal_steps / steps_taken, 0.0, 1.0)
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def penalty_score(
|
| 46 |
+
invalid_actions: int,
|
| 47 |
+
delay_events: int,
|
| 48 |
+
battery_depletion_events: int,
|
| 49 |
+
no_progress_events: int,
|
| 50 |
+
total_steps: int,
|
| 51 |
+
) -> float:
|
| 52 |
+
"""Map penalty events to a [0, 1] quality score where 1 is best."""
|
| 53 |
+
denominator = max(total_steps, 1)
|
| 54 |
+
|
| 55 |
+
weighted_penalties = (
|
| 56 |
+
invalid_actions * 1.0
|
| 57 |
+
+ delay_events * 0.6
|
| 58 |
+
+ battery_depletion_events * 2.0
|
| 59 |
+
+ no_progress_events * 0.25
|
| 60 |
+
)
|
| 61 |
+
base_rate = weighted_penalties / denominator
|
| 62 |
+
|
| 63 |
+
# Add deterministic duration pressure for sustained non-progress behavior so
|
| 64 |
+
# longer stalled episodes are graded worse than short stalls.
|
| 65 |
+
stalled_fraction = no_progress_events / denominator
|
| 66 |
+
stall_duration_pressure = stalled_fraction * min(0.35, denominator * 0.02)
|
| 67 |
+
|
| 68 |
+
rate = base_rate + stall_duration_pressure
|
| 69 |
+
return clamp(1.0 - rate, 0.0, 1.0)
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def optimal_steps_single_agent(
|
| 73 |
+
orders: Sequence[Order],
|
| 74 |
+
start_location: Tuple[int, int] = (0, 0),
|
| 75 |
+
max_exact_orders: int = 10,
|
| 76 |
+
) -> int:
|
| 77 |
+
"""
|
| 78 |
+
Estimate minimal steps for one-agent execution with one active order at a time.
|
| 79 |
+
|
| 80 |
+
For small order counts this is exact DP over completion order and chosen drop location.
|
| 81 |
+
For larger sets it uses a deterministic nearest-first fallback to stay efficient.
|
| 82 |
+
"""
|
| 83 |
+
order_list: List[Order] = list(orders)
|
| 84 |
+
count = len(order_list)
|
| 85 |
+
if count == 0:
|
| 86 |
+
return 0
|
| 87 |
+
|
| 88 |
+
if count > max_exact_orders:
|
| 89 |
+
return _heuristic_optimal_steps(order_list, start_location)
|
| 90 |
+
|
| 91 |
+
pickups = [(o.pickup.x, o.pickup.y) for o in order_list]
|
| 92 |
+
drop_choices = [
|
| 93 |
+
[(d.x, d.y) for d in (o.delivery_locations or [o.dropoff])]
|
| 94 |
+
for o in order_list
|
| 95 |
+
]
|
| 96 |
+
|
| 97 |
+
max_mask = 1 << count
|
| 98 |
+
inf = float("inf")
|
| 99 |
+
|
| 100 |
+
dp = [
|
| 101 |
+
[
|
| 102 |
+
[inf for _ in drop_choices[i]]
|
| 103 |
+
for i in range(count)
|
| 104 |
+
]
|
| 105 |
+
for _ in range(max_mask)
|
| 106 |
+
]
|
| 107 |
+
|
| 108 |
+
for i in range(count):
|
| 109 |
+
for k, drop in enumerate(drop_choices[i]):
|
| 110 |
+
travel = manhattan(start_location, pickups[i]) + manhattan(pickups[i], drop)
|
| 111 |
+
dp[1 << i][i][k] = float(travel)
|
| 112 |
+
|
| 113 |
+
for mask in range(max_mask):
|
| 114 |
+
for last in range(count):
|
| 115 |
+
for last_k, current in enumerate(dp[mask][last]):
|
| 116 |
+
if current == inf:
|
| 117 |
+
continue
|
| 118 |
+
last_drop = drop_choices[last][last_k]
|
| 119 |
+
|
| 120 |
+
for nxt in range(count):
|
| 121 |
+
if mask & (1 << nxt):
|
| 122 |
+
continue
|
| 123 |
+
|
| 124 |
+
transition_to_pickup = manhattan(last_drop, pickups[nxt])
|
| 125 |
+
next_mask = mask | (1 << nxt)
|
| 126 |
+
|
| 127 |
+
for nxt_k, nxt_drop in enumerate(drop_choices[nxt]):
|
| 128 |
+
candidate = (
|
| 129 |
+
current
|
| 130 |
+
+ transition_to_pickup
|
| 131 |
+
+ manhattan(pickups[nxt], nxt_drop)
|
| 132 |
+
)
|
| 133 |
+
if candidate < dp[next_mask][nxt][nxt_k]:
|
| 134 |
+
dp[next_mask][nxt][nxt_k] = candidate
|
| 135 |
+
|
| 136 |
+
full_mask = max_mask - 1
|
| 137 |
+
optimal_travel = min(
|
| 138 |
+
dp[full_mask][last][last_k]
|
| 139 |
+
for last in range(count)
|
| 140 |
+
for last_k in range(len(drop_choices[last]))
|
| 141 |
+
)
|
| 142 |
+
|
| 143 |
+
if optimal_travel == inf:
|
| 144 |
+
return _heuristic_optimal_steps(order_list, start_location)
|
| 145 |
+
|
| 146 |
+
# Minimal action overhead per order: accept + pickup-action + deliver-action.
|
| 147 |
+
action_overhead = 3 * count
|
| 148 |
+
return int(optimal_travel + action_overhead)
|
| 149 |
+
|
| 150 |
+
|
| 151 |
+
def _heuristic_optimal_steps(
|
| 152 |
+
orders: Sequence[Order],
|
| 153 |
+
start_location: Tuple[int, int],
|
| 154 |
+
) -> int:
|
| 155 |
+
remaining = list(orders)
|
| 156 |
+
current = start_location
|
| 157 |
+
travel = 0
|
| 158 |
+
|
| 159 |
+
while remaining:
|
| 160 |
+
best_index = 0
|
| 161 |
+
best_cost = float("inf")
|
| 162 |
+
best_drop = current
|
| 163 |
+
|
| 164 |
+
for idx, order in enumerate(remaining):
|
| 165 |
+
pickup = (order.pickup.x, order.pickup.y)
|
| 166 |
+
drop_candidates = [(d.x, d.y) for d in (order.delivery_locations or [order.dropoff])]
|
| 167 |
+
nearest_drop = min(drop_candidates, key=lambda d: manhattan(pickup, d))
|
| 168 |
+
cost = manhattan(current, pickup) + manhattan(pickup, nearest_drop)
|
| 169 |
+
if cost < best_cost:
|
| 170 |
+
best_cost = cost
|
| 171 |
+
best_index = idx
|
| 172 |
+
best_drop = nearest_drop
|
| 173 |
+
|
| 174 |
+
travel += int(best_cost)
|
| 175 |
+
current = best_drop
|
| 176 |
+
remaining.pop(best_index)
|
| 177 |
+
|
| 178 |
+
return int(travel + 3 * len(orders))
|
inference.py
ADDED
|
@@ -0,0 +1,786 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import argparse
|
| 4 |
+
import json
|
| 5 |
+
import os
|
| 6 |
+
import re
|
| 7 |
+
import time
|
| 8 |
+
from dataclasses import dataclass
|
| 9 |
+
from typing import Any, Dict, List, Optional, Sequence, Tuple
|
| 10 |
+
|
| 11 |
+
from env.environment import LastMileDeliveryEnvironment
|
| 12 |
+
from env.models import Action, ActionType, Direction, Observation, Order, OrderPriority, OrderStatus, SimulatorAction, StepInfo, StepResult
|
| 13 |
+
from env.utils import in_bounds, manhattan, move_coordinate
|
| 14 |
+
from grader.grader import DeliveryEpisodeGrader
|
| 15 |
+
from tasks.registry import get_task_config, get_task_definition, list_tasks
|
| 16 |
+
|
| 17 |
+
try:
|
| 18 |
+
from openai import OpenAI
|
| 19 |
+
except Exception: # pragma: no cover - import guard for environments without dependency installed
|
| 20 |
+
OpenAI = None # type: ignore[assignment]
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
DEFAULT_SEED = 42
|
| 24 |
+
MAX_RUNTIME_SECONDS = 19 * 60
|
| 25 |
+
DEFAULT_TASK = "easy"
|
| 26 |
+
DEFAULT_MODEL_NAME = "gpt-4o-mini"
|
| 27 |
+
|
| 28 |
+
DIRECTION_PRIORITY: List[Direction] = [
|
| 29 |
+
Direction.UP,
|
| 30 |
+
Direction.LEFT,
|
| 31 |
+
Direction.RIGHT,
|
| 32 |
+
Direction.DOWN,
|
| 33 |
+
]
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
@dataclass(frozen=True)
|
| 37 |
+
class RuntimeSettings:
|
| 38 |
+
openai_api_key: str
|
| 39 |
+
model_name: str
|
| 40 |
+
api_base_url: Optional[str]
|
| 41 |
+
use_remote_model: bool
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def _read_env_first(*names: str) -> str:
|
| 45 |
+
for name in names:
|
| 46 |
+
value = (os.getenv(name) or "").strip()
|
| 47 |
+
if value:
|
| 48 |
+
return value
|
| 49 |
+
return ""
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def _require_any_env(*names: str) -> str:
|
| 53 |
+
value = _read_env_first(*names)
|
| 54 |
+
if value:
|
| 55 |
+
return value
|
| 56 |
+
raise RuntimeError(f"Missing required environment variable. Set one of: {', '.join(names)}")
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def _read_bool_env(name: str, default: bool = False) -> bool:
|
| 60 |
+
raw = (os.getenv(name) or "").strip().lower()
|
| 61 |
+
if not raw:
|
| 62 |
+
return default
|
| 63 |
+
return raw in {"1", "true", "yes", "on"}
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def _resolve_model_name() -> str:
|
| 67 |
+
value = _read_env_first("MODEL_NAME", "OPENAI_MODEL")
|
| 68 |
+
return value or DEFAULT_MODEL_NAME
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def load_runtime_settings() -> RuntimeSettings:
|
| 72 |
+
base_url = _read_env_first("API_BASE_URL", "OPENAI_BASE_URL") or None
|
| 73 |
+
model_name = _resolve_model_name()
|
| 74 |
+
|
| 75 |
+
return RuntimeSettings(
|
| 76 |
+
openai_api_key=_require_any_env("HF_TOKEN", "OPENAI_API_KEY"),
|
| 77 |
+
model_name=model_name,
|
| 78 |
+
api_base_url=base_url,
|
| 79 |
+
use_remote_model=_read_bool_env("OPENAI_USE_MODEL", default=False),
|
| 80 |
+
)
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
class DeterministicHeuristicPolicy:
|
| 84 |
+
def __init__(self, settings: RuntimeSettings, seed: int):
|
| 85 |
+
if OpenAI is None:
|
| 86 |
+
raise RuntimeError("openai package is required but not installed")
|
| 87 |
+
|
| 88 |
+
self.model_name = settings.model_name
|
| 89 |
+
self.use_remote_model = settings.use_remote_model
|
| 90 |
+
self.seed = seed
|
| 91 |
+
self._recent_positions: List[Tuple[int, int]] = []
|
| 92 |
+
self._previous_position: Optional[Tuple[int, int]] = None
|
| 93 |
+
# Keep a deterministic fallback path so local validation does not depend on external API availability.
|
| 94 |
+
self.client: Any | None = None
|
| 95 |
+
self.client_init_error: str | None = None
|
| 96 |
+
try:
|
| 97 |
+
client_kwargs: Dict[str, Any] = {
|
| 98 |
+
"api_key": settings.openai_api_key,
|
| 99 |
+
"timeout": 20.0,
|
| 100 |
+
}
|
| 101 |
+
if settings.api_base_url is not None:
|
| 102 |
+
client_kwargs["base_url"] = settings.api_base_url
|
| 103 |
+
self.client = OpenAI(**client_kwargs)
|
| 104 |
+
except Exception as exc:
|
| 105 |
+
self.client_init_error = type(exc).__name__
|
| 106 |
+
|
| 107 |
+
def _update_navigation_memory(self, observation: Observation, step_index: int) -> None:
|
| 108 |
+
current = (observation.agent_location.x, observation.agent_location.y)
|
| 109 |
+
|
| 110 |
+
# Step index resets to 1 at each new episode; clear memory so repeated single-step
|
| 111 |
+
# inference calls remain deterministic.
|
| 112 |
+
if step_index <= 1:
|
| 113 |
+
self._recent_positions = [current]
|
| 114 |
+
self._previous_position = None
|
| 115 |
+
return
|
| 116 |
+
|
| 117 |
+
self._recent_positions.append(current)
|
| 118 |
+
if len(self._recent_positions) > 8:
|
| 119 |
+
self._recent_positions.pop(0)
|
| 120 |
+
|
| 121 |
+
if len(self._recent_positions) >= 2:
|
| 122 |
+
self._previous_position = self._recent_positions[-2]
|
| 123 |
+
else:
|
| 124 |
+
self._previous_position = None
|
| 125 |
+
|
| 126 |
+
def _is_local_oscillation(self) -> bool:
|
| 127 |
+
if len(self._recent_positions) < 4:
|
| 128 |
+
return False
|
| 129 |
+
|
| 130 |
+
a, b, c, d = self._recent_positions[-4:]
|
| 131 |
+
return a == c and b == d and c != d
|
| 132 |
+
|
| 133 |
+
def decide(self, observation: Observation, step_index: int) -> Tuple[SimulatorAction, Optional[str]]:
|
| 134 |
+
self._update_navigation_memory(observation=observation, step_index=step_index)
|
| 135 |
+
break_oscillation = self._is_local_oscillation()
|
| 136 |
+
|
| 137 |
+
if self.use_remote_model and self.client is not None:
|
| 138 |
+
model_action, model_error = self._try_model_action(observation, step_index=step_index)
|
| 139 |
+
if model_action is not None:
|
| 140 |
+
return model_action, model_error
|
| 141 |
+
|
| 142 |
+
try:
|
| 143 |
+
action = _heuristic_action(
|
| 144 |
+
observation,
|
| 145 |
+
previous_position=self._previous_position,
|
| 146 |
+
break_oscillation=break_oscillation,
|
| 147 |
+
)
|
| 148 |
+
return action, self.client_init_error if self.use_remote_model and self.client is None else None
|
| 149 |
+
except Exception as exc:
|
| 150 |
+
fallback = SimulatorAction(action_type=ActionType.WAIT)
|
| 151 |
+
return fallback, f"policy_error:{type(exc).__name__}"
|
| 152 |
+
|
| 153 |
+
def _try_model_action(self, observation: Observation, step_index: int) -> Tuple[Optional[SimulatorAction], Optional[str]]:
|
| 154 |
+
if self.client is None:
|
| 155 |
+
return None, "openai_client_unavailable"
|
| 156 |
+
|
| 157 |
+
system_prompt = (
|
| 158 |
+
"You are a deterministic last-mile dispatch policy. "
|
| 159 |
+
"Return one JSON object matching this schema exactly: "
|
| 160 |
+
"{\"move\": string|null, \"accept_order\": integer|null, \"deliver_order\": boolean, \"wait\": boolean}. "
|
| 161 |
+
"Exactly one intent must be active. Do not include markdown."
|
| 162 |
+
)
|
| 163 |
+
payload = {
|
| 164 |
+
"step_index": step_index,
|
| 165 |
+
"observation": observation.model_dump(),
|
| 166 |
+
}
|
| 167 |
+
|
| 168 |
+
try:
|
| 169 |
+
response = self.client.responses.create(
|
| 170 |
+
model=self.model_name,
|
| 171 |
+
temperature=0,
|
| 172 |
+
max_output_tokens=120,
|
| 173 |
+
input=[
|
| 174 |
+
{
|
| 175 |
+
"role": "system",
|
| 176 |
+
"content": [{"type": "text", "text": system_prompt}],
|
| 177 |
+
},
|
| 178 |
+
{
|
| 179 |
+
"role": "user",
|
| 180 |
+
"content": [{"type": "text", "text": json.dumps(payload, separators=(",", ":"), sort_keys=True)}],
|
| 181 |
+
},
|
| 182 |
+
],
|
| 183 |
+
)
|
| 184 |
+
except Exception as exc:
|
| 185 |
+
return None, f"model_request_error:{type(exc).__name__}"
|
| 186 |
+
|
| 187 |
+
output_text = getattr(response, "output_text", "") or ""
|
| 188 |
+
parsed_json = self._extract_json_object(output_text)
|
| 189 |
+
if parsed_json is None:
|
| 190 |
+
return None, "model_parse_error"
|
| 191 |
+
|
| 192 |
+
try:
|
| 193 |
+
action = Action.model_validate(parsed_json).to_simulator_action()
|
| 194 |
+
except Exception as exc:
|
| 195 |
+
return None, f"model_action_invalid:{type(exc).__name__}"
|
| 196 |
+
|
| 197 |
+
return action, None
|
| 198 |
+
|
| 199 |
+
@staticmethod
|
| 200 |
+
def _extract_json_object(text: str) -> Optional[Dict[str, Any]]:
|
| 201 |
+
candidate = text.strip()
|
| 202 |
+
if not candidate:
|
| 203 |
+
return None
|
| 204 |
+
|
| 205 |
+
if candidate.startswith("```"):
|
| 206 |
+
candidate = re.sub(r"^```[a-zA-Z0-9_\-]*\s*", "", candidate)
|
| 207 |
+
candidate = re.sub(r"\s*```$", "", candidate).strip()
|
| 208 |
+
|
| 209 |
+
try:
|
| 210 |
+
value = json.loads(candidate)
|
| 211 |
+
if isinstance(value, dict):
|
| 212 |
+
return value
|
| 213 |
+
except Exception:
|
| 214 |
+
pass
|
| 215 |
+
|
| 216 |
+
start = candidate.find("{")
|
| 217 |
+
end = candidate.rfind("}")
|
| 218 |
+
if start == -1 or end == -1 or end <= start:
|
| 219 |
+
return None
|
| 220 |
+
|
| 221 |
+
try:
|
| 222 |
+
value = json.loads(candidate[start : end + 1])
|
| 223 |
+
if isinstance(value, dict):
|
| 224 |
+
return value
|
| 225 |
+
except Exception:
|
| 226 |
+
return None
|
| 227 |
+
|
| 228 |
+
return None
|
| 229 |
+
|
| 230 |
+
|
| 231 |
+
_POLICY: DeterministicHeuristicPolicy | None = None
|
| 232 |
+
|
| 233 |
+
|
| 234 |
+
def _get_policy(seed: int = DEFAULT_SEED) -> DeterministicHeuristicPolicy:
|
| 235 |
+
global _POLICY
|
| 236 |
+
if _POLICY is None:
|
| 237 |
+
settings = load_runtime_settings()
|
| 238 |
+
_POLICY = DeterministicHeuristicPolicy(settings=settings, seed=seed)
|
| 239 |
+
return _POLICY
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
def _priority_rank(order: Order) -> int:
|
| 243 |
+
return 0 if order.priority == OrderPriority.HIGH else 1
|
| 244 |
+
|
| 245 |
+
|
| 246 |
+
def _choose_pending_order(observation: Observation) -> Optional[Order]:
|
| 247 |
+
agent = (observation.agent_location.x, observation.agent_location.y)
|
| 248 |
+
pending = [order for order in observation.pending_orders if order.status == OrderStatus.PENDING]
|
| 249 |
+
if not pending:
|
| 250 |
+
return None
|
| 251 |
+
|
| 252 |
+
pending.sort(
|
| 253 |
+
key=lambda order: (
|
| 254 |
+
_priority_rank(order),
|
| 255 |
+
manhattan(agent, (order.pickup.x, order.pickup.y)),
|
| 256 |
+
order.order_id,
|
| 257 |
+
)
|
| 258 |
+
)
|
| 259 |
+
return pending[0]
|
| 260 |
+
|
| 261 |
+
|
| 262 |
+
def _current_target(observation: Observation) -> Optional[Tuple[int, int]]:
|
| 263 |
+
current_order = observation.current_order
|
| 264 |
+
if current_order is None:
|
| 265 |
+
return None
|
| 266 |
+
|
| 267 |
+
agent = (observation.agent_location.x, observation.agent_location.y)
|
| 268 |
+
if current_order.status == OrderStatus.ACCEPTED:
|
| 269 |
+
return (current_order.pickup.x, current_order.pickup.y)
|
| 270 |
+
|
| 271 |
+
if current_order.status == OrderStatus.PICKED_UP:
|
| 272 |
+
targets = current_order.delivery_locations or [current_order.dropoff]
|
| 273 |
+
target_coords = [(target.x, target.y) for target in targets]
|
| 274 |
+
target_coords.sort(key=lambda coord: (manhattan(agent, coord), coord[0], coord[1]))
|
| 275 |
+
return target_coords[0]
|
| 276 |
+
|
| 277 |
+
if current_order.status == OrderStatus.PENDING:
|
| 278 |
+
return (current_order.pickup.x, current_order.pickup.y)
|
| 279 |
+
|
| 280 |
+
return None
|
| 281 |
+
|
| 282 |
+
|
| 283 |
+
def _nearest_charging_station(observation: Observation) -> Optional[Tuple[int, int]]:
|
| 284 |
+
if not observation.charging_stations:
|
| 285 |
+
return None
|
| 286 |
+
|
| 287 |
+
agent = (observation.agent_location.x, observation.agent_location.y)
|
| 288 |
+
station_coords = [(item.x, item.y) for item in observation.charging_stations]
|
| 289 |
+
station_coords.sort(key=lambda coord: (manhattan(agent, coord), coord[0], coord[1]))
|
| 290 |
+
return station_coords[0]
|
| 291 |
+
|
| 292 |
+
|
| 293 |
+
def _best_move_direction(
|
| 294 |
+
observation: Observation,
|
| 295 |
+
target: Tuple[int, int],
|
| 296 |
+
previous_position: Optional[Tuple[int, int]] = None,
|
| 297 |
+
break_oscillation: bool = False,
|
| 298 |
+
) -> Optional[Direction]:
|
| 299 |
+
current = (observation.agent_location.x, observation.agent_location.y)
|
| 300 |
+
blocked = {(item.x, item.y) for item in observation.obstacles}
|
| 301 |
+
blocked.update((item.x, item.y) for item in observation.dynamic_obstacles)
|
| 302 |
+
|
| 303 |
+
candidates: List[Tuple[int, int, Direction, Tuple[int, int]]] = []
|
| 304 |
+
for index, direction in enumerate(DIRECTION_PRIORITY):
|
| 305 |
+
nxt = move_coordinate(current, direction)
|
| 306 |
+
if not in_bounds(nxt, observation.grid_width, observation.grid_height):
|
| 307 |
+
continue
|
| 308 |
+
if nxt in blocked:
|
| 309 |
+
continue
|
| 310 |
+
|
| 311 |
+
distance = manhattan(nxt, target)
|
| 312 |
+
candidates.append((distance, index, direction, nxt))
|
| 313 |
+
|
| 314 |
+
if not candidates:
|
| 315 |
+
return None
|
| 316 |
+
|
| 317 |
+
candidates.sort(key=lambda item: (item[0], item[1]))
|
| 318 |
+
|
| 319 |
+
if previous_position is not None and len(candidates) > 1:
|
| 320 |
+
best_distance = candidates[0][0]
|
| 321 |
+
for distance, _, direction, nxt in candidates:
|
| 322 |
+
if nxt == previous_position:
|
| 323 |
+
continue
|
| 324 |
+
|
| 325 |
+
# Stronger anti-backtrack when we detect A-B-A-B style oscillation.
|
| 326 |
+
if break_oscillation:
|
| 327 |
+
return direction
|
| 328 |
+
|
| 329 |
+
# Prefer a non-backtracking move when it is near-optimal.
|
| 330 |
+
if distance <= best_distance + 1:
|
| 331 |
+
return direction
|
| 332 |
+
|
| 333 |
+
break
|
| 334 |
+
|
| 335 |
+
return candidates[0][2]
|
| 336 |
+
|
| 337 |
+
|
| 338 |
+
def _move_toward_target(
|
| 339 |
+
observation: Observation,
|
| 340 |
+
target: Tuple[int, int],
|
| 341 |
+
previous_position: Optional[Tuple[int, int]] = None,
|
| 342 |
+
break_oscillation: bool = False,
|
| 343 |
+
) -> SimulatorAction:
|
| 344 |
+
direction = _best_move_direction(
|
| 345 |
+
observation,
|
| 346 |
+
target,
|
| 347 |
+
previous_position=previous_position,
|
| 348 |
+
break_oscillation=break_oscillation,
|
| 349 |
+
)
|
| 350 |
+
if direction is None:
|
| 351 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 352 |
+
return SimulatorAction(action_type=ActionType.MOVE, direction=direction)
|
| 353 |
+
|
| 354 |
+
|
| 355 |
+
def _battery_aware_action(
|
| 356 |
+
observation: Observation,
|
| 357 |
+
target: Tuple[int, int],
|
| 358 |
+
previous_position: Optional[Tuple[int, int]] = None,
|
| 359 |
+
break_oscillation: bool = False,
|
| 360 |
+
) -> Optional[SimulatorAction]:
|
| 361 |
+
battery_level = observation.battery_level
|
| 362 |
+
if battery_level is None:
|
| 363 |
+
return None
|
| 364 |
+
|
| 365 |
+
# Keep pickup behavior aggressive; only interrupt for charging when battery is near-empty.
|
| 366 |
+
if observation.current_order is None and battery_level > 3:
|
| 367 |
+
return None
|
| 368 |
+
|
| 369 |
+
charger_target = _nearest_charging_station(observation)
|
| 370 |
+
if charger_target is None:
|
| 371 |
+
return None
|
| 372 |
+
|
| 373 |
+
agent = (observation.agent_location.x, observation.agent_location.y)
|
| 374 |
+
distance_to_charger = manhattan(agent, charger_target)
|
| 375 |
+
distance_to_target = manhattan(agent, target)
|
| 376 |
+
|
| 377 |
+
# Trigger recharge only when battery is close to the minimum needed to reach a charger.
|
| 378 |
+
recharge_trigger = distance_to_charger + 1
|
| 379 |
+
if battery_level > recharge_trigger:
|
| 380 |
+
return None
|
| 381 |
+
|
| 382 |
+
if agent == charger_target:
|
| 383 |
+
# Recharge just enough to leave and make forward progress; avoid long waiting loops.
|
| 384 |
+
resume_level = min(12, max(8, distance_to_target + 2))
|
| 385 |
+
if battery_level < resume_level:
|
| 386 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 387 |
+
return None
|
| 388 |
+
|
| 389 |
+
return _move_toward_target(
|
| 390 |
+
observation,
|
| 391 |
+
charger_target,
|
| 392 |
+
previous_position=previous_position,
|
| 393 |
+
break_oscillation=break_oscillation,
|
| 394 |
+
)
|
| 395 |
+
|
| 396 |
+
|
| 397 |
+
def _heuristic_action(
|
| 398 |
+
observation: Observation,
|
| 399 |
+
previous_position: Optional[Tuple[int, int]] = None,
|
| 400 |
+
break_oscillation: bool = False,
|
| 401 |
+
) -> SimulatorAction:
|
| 402 |
+
agent = (observation.agent_location.x, observation.agent_location.y)
|
| 403 |
+
current_order = observation.current_order
|
| 404 |
+
|
| 405 |
+
if current_order is None:
|
| 406 |
+
selected = _choose_pending_order(observation)
|
| 407 |
+
if selected is None:
|
| 408 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 409 |
+
|
| 410 |
+
pickup_target = (selected.pickup.x, selected.pickup.y)
|
| 411 |
+
if agent == pickup_target:
|
| 412 |
+
return SimulatorAction(action_type=ActionType.ACCEPT_ORDER, order_id=selected.order_id)
|
| 413 |
+
|
| 414 |
+
recharge_action = _battery_aware_action(
|
| 415 |
+
observation,
|
| 416 |
+
pickup_target,
|
| 417 |
+
previous_position=previous_position,
|
| 418 |
+
break_oscillation=break_oscillation,
|
| 419 |
+
)
|
| 420 |
+
if recharge_action is not None:
|
| 421 |
+
return recharge_action
|
| 422 |
+
|
| 423 |
+
# Move toward the nearest pickup instead of idling while no order is selected.
|
| 424 |
+
return _move_toward_target(
|
| 425 |
+
observation,
|
| 426 |
+
pickup_target,
|
| 427 |
+
previous_position=previous_position,
|
| 428 |
+
break_oscillation=break_oscillation,
|
| 429 |
+
)
|
| 430 |
+
|
| 431 |
+
target = _current_target(observation)
|
| 432 |
+
if target is None:
|
| 433 |
+
selected = _choose_pending_order(observation)
|
| 434 |
+
if selected is None:
|
| 435 |
+
return SimulatorAction(action_type=ActionType.WAIT)
|
| 436 |
+
return _move_toward_target(
|
| 437 |
+
observation,
|
| 438 |
+
(selected.pickup.x, selected.pickup.y),
|
| 439 |
+
previous_position=previous_position,
|
| 440 |
+
break_oscillation=break_oscillation,
|
| 441 |
+
)
|
| 442 |
+
|
| 443 |
+
if agent == target:
|
| 444 |
+
return SimulatorAction(action_type=ActionType.DELIVER_ORDER)
|
| 445 |
+
|
| 446 |
+
recharge_action = _battery_aware_action(
|
| 447 |
+
observation,
|
| 448 |
+
target,
|
| 449 |
+
previous_position=previous_position,
|
| 450 |
+
break_oscillation=break_oscillation,
|
| 451 |
+
)
|
| 452 |
+
if recharge_action is not None:
|
| 453 |
+
return recharge_action
|
| 454 |
+
|
| 455 |
+
# With an active target, prefer movement over waiting whenever a valid move exists.
|
| 456 |
+
return _move_toward_target(
|
| 457 |
+
observation,
|
| 458 |
+
target,
|
| 459 |
+
previous_position=previous_position,
|
| 460 |
+
break_oscillation=break_oscillation,
|
| 461 |
+
)
|
| 462 |
+
|
| 463 |
+
|
| 464 |
+
def _bool_literal(value: bool) -> str:
|
| 465 |
+
return "true" if value else "false"
|
| 466 |
+
|
| 467 |
+
|
| 468 |
+
def _format_error(error: Optional[str]) -> str:
|
| 469 |
+
if error is None:
|
| 470 |
+
return "null"
|
| 471 |
+
return error.replace(" ", "_")
|
| 472 |
+
|
| 473 |
+
|
| 474 |
+
def _format_token(value: Optional[str], fallback: str) -> str:
|
| 475 |
+
token = (value or "").strip()
|
| 476 |
+
if not token:
|
| 477 |
+
token = fallback
|
| 478 |
+
return "_".join(token.split())
|
| 479 |
+
|
| 480 |
+
|
| 481 |
+
def _format_action(action: SimulatorAction) -> str:
|
| 482 |
+
payload = Action.from_simulator_action(action).model_dump()
|
| 483 |
+
return json.dumps(payload, separators=(",", ":"), sort_keys=True)
|
| 484 |
+
|
| 485 |
+
|
| 486 |
+
def format_start_line(task: str, env_name: str, model_name: str) -> str:
|
| 487 |
+
return (
|
| 488 |
+
f"[START] "
|
| 489 |
+
f"task={_format_token(task, fallback='unknown_task')} "
|
| 490 |
+
f"env={_format_token(env_name, fallback='unknown_env')} "
|
| 491 |
+
f"model={_format_token(model_name, fallback='unknown_model')}"
|
| 492 |
+
)
|
| 493 |
+
|
| 494 |
+
|
| 495 |
+
def format_step_line(step_index: int, action: SimulatorAction, reward: float, done: bool, error: Optional[str]) -> str:
|
| 496 |
+
return (
|
| 497 |
+
f"[STEP] step={step_index} "
|
| 498 |
+
f"action={_format_action(action)} "
|
| 499 |
+
f"reward={reward:.2f} "
|
| 500 |
+
f"done={_bool_literal(done)} "
|
| 501 |
+
f"error={_format_error(error)}"
|
| 502 |
+
)
|
| 503 |
+
|
| 504 |
+
|
| 505 |
+
def _format_reward_series(rewards: Sequence[float] | float) -> str:
|
| 506 |
+
if isinstance(rewards, (int, float)):
|
| 507 |
+
return f"{float(rewards):.2f}"
|
| 508 |
+
|
| 509 |
+
reward_values = list(rewards)
|
| 510 |
+
if not reward_values:
|
| 511 |
+
return "0.00"
|
| 512 |
+
return ",".join(f"{value:.2f}" for value in reward_values)
|
| 513 |
+
|
| 514 |
+
|
| 515 |
+
def format_end_line(success: bool, steps: int, rewards: Sequence[float] | float) -> str:
|
| 516 |
+
return f"[END] success={_bool_literal(success)} steps={steps} rewards={_format_reward_series(rewards)}"
|
| 517 |
+
|
| 518 |
+
|
| 519 |
+
def _episode_success(done: bool, observation: Observation) -> bool:
|
| 520 |
+
if not done:
|
| 521 |
+
return False
|
| 522 |
+
return len(observation.pending_orders) == 0 and observation.current_order is None
|
| 523 |
+
|
| 524 |
+
|
| 525 |
+
def _grade_score(task: str, trajectory: Sequence[StepResult], final_observation: Observation) -> float:
|
| 526 |
+
grader = DeliveryEpisodeGrader()
|
| 527 |
+
success_condition: Dict[str, Any] = {}
|
| 528 |
+
try:
|
| 529 |
+
task_definition = get_task_definition(task)
|
| 530 |
+
success_condition = dict(task_definition.get("success_condition", {}))
|
| 531 |
+
except Exception:
|
| 532 |
+
success_condition = {}
|
| 533 |
+
|
| 534 |
+
report = grader.grade_episode(
|
| 535 |
+
trajectory=trajectory,
|
| 536 |
+
final_observation=final_observation,
|
| 537 |
+
success_condition=success_condition,
|
| 538 |
+
)
|
| 539 |
+
return float(report.score)
|
| 540 |
+
|
| 541 |
+
|
| 542 |
+
def run_episode(
|
| 543 |
+
task: str = DEFAULT_TASK,
|
| 544 |
+
seed: int = DEFAULT_SEED,
|
| 545 |
+
max_steps: Optional[int] = None,
|
| 546 |
+
trace: bool = True,
|
| 547 |
+
) -> Dict[str, Any]:
|
| 548 |
+
steps = 0
|
| 549 |
+
reward_history: List[float] = []
|
| 550 |
+
trajectory: List[StepResult] = []
|
| 551 |
+
success = False
|
| 552 |
+
step_printed = False
|
| 553 |
+
startup_error: Optional[str] = None
|
| 554 |
+
runtime_error: Optional[str] = None
|
| 555 |
+
model_name = _resolve_model_name()
|
| 556 |
+
env_name = LastMileDeliveryEnvironment.__name__
|
| 557 |
+
settings: Optional[RuntimeSettings] = None
|
| 558 |
+
score = 0.0
|
| 559 |
+
env: Optional[LastMileDeliveryEnvironment] = None
|
| 560 |
+
observation: Optional[Observation] = None
|
| 561 |
+
|
| 562 |
+
try:
|
| 563 |
+
try:
|
| 564 |
+
settings = load_runtime_settings()
|
| 565 |
+
model_name = settings.model_name
|
| 566 |
+
except Exception as exc:
|
| 567 |
+
startup_error = f"startup_error:{type(exc).__name__}"
|
| 568 |
+
|
| 569 |
+
if trace:
|
| 570 |
+
print(format_start_line(task=task, env_name=env_name, model_name=model_name))
|
| 571 |
+
|
| 572 |
+
if settings is None:
|
| 573 |
+
action = SimulatorAction(action_type=ActionType.WAIT)
|
| 574 |
+
if trace:
|
| 575 |
+
print(
|
| 576 |
+
format_step_line(
|
| 577 |
+
step_index=1,
|
| 578 |
+
action=action,
|
| 579 |
+
reward=0.0,
|
| 580 |
+
done=False,
|
| 581 |
+
error=startup_error,
|
| 582 |
+
)
|
| 583 |
+
)
|
| 584 |
+
step_printed = True
|
| 585 |
+
return {
|
| 586 |
+
"success": False,
|
| 587 |
+
"steps": 0,
|
| 588 |
+
"rewards": [],
|
| 589 |
+
"total_reward": 0.0,
|
| 590 |
+
"score": 0.0,
|
| 591 |
+
}
|
| 592 |
+
|
| 593 |
+
policy = DeterministicHeuristicPolicy(settings=settings, seed=seed)
|
| 594 |
+
config = get_task_config(task_name=task, seed=seed)
|
| 595 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 596 |
+
observation = env.reset()
|
| 597 |
+
|
| 598 |
+
start_time = time.monotonic()
|
| 599 |
+
configured_limit = max_steps if max_steps is not None else config.max_steps
|
| 600 |
+
step_limit = max(1, min(configured_limit, config.max_steps))
|
| 601 |
+
|
| 602 |
+
for step_index in range(1, step_limit + 1):
|
| 603 |
+
if time.monotonic() - start_time >= MAX_RUNTIME_SECONDS:
|
| 604 |
+
runtime_error = "timeout"
|
| 605 |
+
break
|
| 606 |
+
|
| 607 |
+
action, action_error = policy.decide(observation, step_index=step_index)
|
| 608 |
+
try:
|
| 609 |
+
step_result = env.step_result(action)
|
| 610 |
+
trajectory.append(step_result)
|
| 611 |
+
observation = step_result.observation
|
| 612 |
+
reward = float(step_result.reward)
|
| 613 |
+
done = step_result.done
|
| 614 |
+
info = step_result.info.model_dump()
|
| 615 |
+
except Exception as exc:
|
| 616 |
+
reward = 0.0
|
| 617 |
+
done = False
|
| 618 |
+
info = {"invalid_action": True, "message": f"step_error:{type(exc).__name__}"}
|
| 619 |
+
action = SimulatorAction(action_type=ActionType.WAIT)
|
| 620 |
+
runtime_error = str(info["message"])
|
| 621 |
+
if env is not None:
|
| 622 |
+
fallback_observation = env.current_observation()
|
| 623 |
+
trajectory.append(
|
| 624 |
+
StepResult(
|
| 625 |
+
observation=fallback_observation,
|
| 626 |
+
reward=reward,
|
| 627 |
+
done=done,
|
| 628 |
+
info=StepInfo(invalid_action=True, message=str(info["message"])),
|
| 629 |
+
)
|
| 630 |
+
)
|
| 631 |
+
|
| 632 |
+
step_error = action_error
|
| 633 |
+
if step_error is None and bool(info.get("invalid_action")):
|
| 634 |
+
step_error = str(info.get("message", "invalid_action"))
|
| 635 |
+
|
| 636 |
+
steps = step_index
|
| 637 |
+
reward_history.append(float(reward))
|
| 638 |
+
if trace:
|
| 639 |
+
print(
|
| 640 |
+
format_step_line(
|
| 641 |
+
step_index=step_index,
|
| 642 |
+
action=action,
|
| 643 |
+
reward=reward,
|
| 644 |
+
done=done,
|
| 645 |
+
error=step_error,
|
| 646 |
+
)
|
| 647 |
+
)
|
| 648 |
+
step_printed = True
|
| 649 |
+
|
| 650 |
+
if done:
|
| 651 |
+
success = _episode_success(done=done, observation=observation)
|
| 652 |
+
break
|
| 653 |
+
if runtime_error is not None:
|
| 654 |
+
break
|
| 655 |
+
|
| 656 |
+
if env is not None:
|
| 657 |
+
final_observation = env.current_observation()
|
| 658 |
+
score = _grade_score(task=task, trajectory=trajectory, final_observation=final_observation)
|
| 659 |
+
|
| 660 |
+
return {
|
| 661 |
+
"success": success,
|
| 662 |
+
"steps": steps,
|
| 663 |
+
"rewards": [round(value, 2) for value in reward_history],
|
| 664 |
+
"total_reward": round(sum(reward_history), 2),
|
| 665 |
+
"score": round(score, 4),
|
| 666 |
+
}
|
| 667 |
+
|
| 668 |
+
except BaseException as exc:
|
| 669 |
+
runtime_error = f"fatal_error:{type(exc).__name__}"
|
| 670 |
+
return {
|
| 671 |
+
"success": False,
|
| 672 |
+
"steps": steps,
|
| 673 |
+
"rewards": [round(value, 2) for value in reward_history],
|
| 674 |
+
"total_reward": round(sum(reward_history), 2),
|
| 675 |
+
"score": 0.0,
|
| 676 |
+
}
|
| 677 |
+
finally:
|
| 678 |
+
if trace and not step_printed:
|
| 679 |
+
action = SimulatorAction(action_type=ActionType.WAIT)
|
| 680 |
+
fallback_error = runtime_error or startup_error
|
| 681 |
+
print(
|
| 682 |
+
format_step_line(
|
| 683 |
+
step_index=max(1, steps + 1),
|
| 684 |
+
action=action,
|
| 685 |
+
reward=0.0,
|
| 686 |
+
done=False,
|
| 687 |
+
error=fallback_error,
|
| 688 |
+
)
|
| 689 |
+
)
|
| 690 |
+
|
| 691 |
+
if trace:
|
| 692 |
+
print(format_end_line(success=success, steps=steps, rewards=reward_history))
|
| 693 |
+
|
| 694 |
+
|
| 695 |
+
def run_baseline_suite(seed: int = DEFAULT_SEED, max_steps: Optional[int] = None) -> Dict[str, Any]:
|
| 696 |
+
task_results: List[Dict[str, Any]] = []
|
| 697 |
+
for task_name in list_tasks():
|
| 698 |
+
episode_result = run_episode(task=task_name, seed=seed, max_steps=max_steps, trace=False)
|
| 699 |
+
task_results.append(
|
| 700 |
+
{
|
| 701 |
+
"task": task_name,
|
| 702 |
+
"success": bool(episode_result.get("success", False)),
|
| 703 |
+
"steps": int(episode_result.get("steps", 0)),
|
| 704 |
+
"total_reward": float(episode_result.get("total_reward", 0.0)),
|
| 705 |
+
"score": float(episode_result.get("score", 0.0)),
|
| 706 |
+
}
|
| 707 |
+
)
|
| 708 |
+
|
| 709 |
+
average_score = 0.0
|
| 710 |
+
if task_results:
|
| 711 |
+
average_score = round(sum(item["score"] for item in task_results) / len(task_results), 4)
|
| 712 |
+
|
| 713 |
+
return {
|
| 714 |
+
"seed": seed,
|
| 715 |
+
"model": _resolve_model_name(),
|
| 716 |
+
"tasks": task_results,
|
| 717 |
+
"average_score": average_score,
|
| 718 |
+
}
|
| 719 |
+
|
| 720 |
+
|
| 721 |
+
def inference(observation: Dict[str, Any]) -> Dict[str, Any]:
|
| 722 |
+
"""OpenEnv action entrypoint with deterministic heuristic policy and OpenAI client setup."""
|
| 723 |
+
parsed = Observation.model_validate(observation)
|
| 724 |
+
try:
|
| 725 |
+
action, _ = _get_policy(seed=DEFAULT_SEED).decide(parsed, step_index=1)
|
| 726 |
+
except Exception:
|
| 727 |
+
action = _heuristic_action(parsed)
|
| 728 |
+
return Action.from_simulator_action(action).model_dump()
|
| 729 |
+
|
| 730 |
+
|
| 731 |
+
def predict(observation: Dict[str, Any]) -> Dict[str, Any]:
|
| 732 |
+
"""Compatibility alias for runners expecting predict(...)."""
|
| 733 |
+
return inference(observation)
|
| 734 |
+
|
| 735 |
+
|
| 736 |
+
def _build_arg_parser() -> argparse.ArgumentParser:
|
| 737 |
+
parser = argparse.ArgumentParser(description="Run deterministic OpenEnv inference episode")
|
| 738 |
+
parser.add_argument("--task", default=DEFAULT_TASK, type=str, help="Task name from task registry")
|
| 739 |
+
parser.add_argument("--seed", default=DEFAULT_SEED, type=int, help="Deterministic random seed")
|
| 740 |
+
parser.add_argument(
|
| 741 |
+
"--all-tasks",
|
| 742 |
+
action="store_true",
|
| 743 |
+
help="Run reproducible baseline evaluation across easy, medium, and hard tasks",
|
| 744 |
+
)
|
| 745 |
+
parser.add_argument(
|
| 746 |
+
"--max-steps",
|
| 747 |
+
type=int,
|
| 748 |
+
default=None,
|
| 749 |
+
help="Optional cap on interaction steps (bounded by task max_steps)",
|
| 750 |
+
)
|
| 751 |
+
parser.add_argument(
|
| 752 |
+
"--json-summary",
|
| 753 |
+
action="store_true",
|
| 754 |
+
help="When used with --all-tasks, print compact JSON summary",
|
| 755 |
+
)
|
| 756 |
+
return parser
|
| 757 |
+
|
| 758 |
+
|
| 759 |
+
def main() -> None:
|
| 760 |
+
args = _build_arg_parser().parse_args()
|
| 761 |
+
if args.all_tasks:
|
| 762 |
+
summary = run_baseline_suite(seed=args.seed, max_steps=args.max_steps)
|
| 763 |
+
if args.json_summary:
|
| 764 |
+
print(json.dumps(summary, separators=(",", ":"), sort_keys=True))
|
| 765 |
+
return
|
| 766 |
+
|
| 767 |
+
for item in summary["tasks"]:
|
| 768 |
+
print(
|
| 769 |
+
f"[BASELINE] task={item['task']} "
|
| 770 |
+
f"score={item['score']:.4f} "
|
| 771 |
+
f"steps={item['steps']} "
|
| 772 |
+
f"success={_bool_literal(item['success'])}"
|
| 773 |
+
)
|
| 774 |
+
|
| 775 |
+
print(
|
| 776 |
+
f"[BASELINE_SUMMARY] seed={summary['seed']} "
|
| 777 |
+
f"tasks={len(summary['tasks'])} "
|
| 778 |
+
f"average_score={summary['average_score']:.4f}"
|
| 779 |
+
)
|
| 780 |
+
return
|
| 781 |
+
|
| 782 |
+
run_episode(task=args.task, seed=args.seed, max_steps=args.max_steps, trace=True)
|
| 783 |
+
|
| 784 |
+
|
| 785 |
+
if __name__ == "__main__":
|
| 786 |
+
main()
|
models/q_agent_easy.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
openenv.yaml
ADDED
|
@@ -0,0 +1,127 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
name: last-mile-delivery-optimization
|
| 2 |
+
description: OpenEnv-compatible grid delivery environment with tasks, grading, and baseline policy.
|
| 3 |
+
version: 1.0.0
|
| 4 |
+
entrypoint: app.main:app
|
| 5 |
+
|
| 6 |
+
runtime:
|
| 7 |
+
api_app: app.main:app
|
| 8 |
+
inference_entrypoint: inference.py
|
| 9 |
+
|
| 10 |
+
environment:
|
| 11 |
+
module: env.environment
|
| 12 |
+
class: LastMileDeliveryEnvironment
|
| 13 |
+
|
| 14 |
+
tasks:
|
| 15 |
+
registry: tasks.registry:get_all_task_metadata
|
| 16 |
+
config_loader: tasks.registry:get_task_config
|
| 17 |
+
|
| 18 |
+
grader:
|
| 19 |
+
module: grader.grader
|
| 20 |
+
class: DeliveryEpisodeGrader
|
| 21 |
+
|
| 22 |
+
baseline:
|
| 23 |
+
module: baseline.baseline_agent
|
| 24 |
+
class: BaselineGreedyAgent
|
| 25 |
+
|
| 26 |
+
action_schema: &action_schema
|
| 27 |
+
type: object
|
| 28 |
+
required:
|
| 29 |
+
- move
|
| 30 |
+
- accept_order
|
| 31 |
+
- deliver_order
|
| 32 |
+
- wait
|
| 33 |
+
additionalProperties: false
|
| 34 |
+
properties:
|
| 35 |
+
move:
|
| 36 |
+
type:
|
| 37 |
+
- string
|
| 38 |
+
- "null"
|
| 39 |
+
enum:
|
| 40 |
+
- up
|
| 41 |
+
- down
|
| 42 |
+
- left
|
| 43 |
+
- right
|
| 44 |
+
- stay
|
| 45 |
+
- null
|
| 46 |
+
accept_order:
|
| 47 |
+
type:
|
| 48 |
+
- integer
|
| 49 |
+
- "null"
|
| 50 |
+
deliver_order:
|
| 51 |
+
type: boolean
|
| 52 |
+
wait:
|
| 53 |
+
type: boolean
|
| 54 |
+
|
| 55 |
+
action_space: *action_schema
|
| 56 |
+
|
| 57 |
+
observation_schema: &observation_schema
|
| 58 |
+
type: object
|
| 59 |
+
required:
|
| 60 |
+
- agent_location
|
| 61 |
+
- pending_orders
|
| 62 |
+
- current_order
|
| 63 |
+
- obstacles
|
| 64 |
+
- dynamic_obstacles
|
| 65 |
+
- charging_stations
|
| 66 |
+
- battery_level
|
| 67 |
+
additionalProperties: true
|
| 68 |
+
properties:
|
| 69 |
+
agent_location:
|
| 70 |
+
type: object
|
| 71 |
+
required:
|
| 72 |
+
- x
|
| 73 |
+
- y
|
| 74 |
+
properties:
|
| 75 |
+
x:
|
| 76 |
+
type: integer
|
| 77 |
+
minimum: 0
|
| 78 |
+
y:
|
| 79 |
+
type: integer
|
| 80 |
+
minimum: 0
|
| 81 |
+
pending_orders:
|
| 82 |
+
type: array
|
| 83 |
+
items:
|
| 84 |
+
type: object
|
| 85 |
+
current_order:
|
| 86 |
+
oneOf:
|
| 87 |
+
- type: object
|
| 88 |
+
- type: "null"
|
| 89 |
+
obstacles:
|
| 90 |
+
type: array
|
| 91 |
+
items:
|
| 92 |
+
type: object
|
| 93 |
+
dynamic_obstacles:
|
| 94 |
+
type: array
|
| 95 |
+
items:
|
| 96 |
+
type: object
|
| 97 |
+
charging_stations:
|
| 98 |
+
type: array
|
| 99 |
+
items:
|
| 100 |
+
type: object
|
| 101 |
+
battery_level:
|
| 102 |
+
oneOf:
|
| 103 |
+
- type: integer
|
| 104 |
+
minimum: 0
|
| 105 |
+
- type: "null"
|
| 106 |
+
|
| 107 |
+
observation_space: *observation_schema
|
| 108 |
+
|
| 109 |
+
reward_schema: &reward_schema
|
| 110 |
+
type: object
|
| 111 |
+
required:
|
| 112 |
+
- value
|
| 113 |
+
additionalProperties: false
|
| 114 |
+
properties:
|
| 115 |
+
value:
|
| 116 |
+
type: number
|
| 117 |
+
components:
|
| 118 |
+
type: object
|
| 119 |
+
additionalProperties:
|
| 120 |
+
type: number
|
| 121 |
+
|
| 122 |
+
reward_space: *reward_schema
|
| 123 |
+
|
| 124 |
+
reward_type: dense
|
| 125 |
+
reward:
|
| 126 |
+
type: dense
|
| 127 |
+
schema: *reward_schema
|
pyproject.toml
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[build-system]
|
| 2 |
+
requires = ["setuptools>=68", "wheel"]
|
| 3 |
+
build-backend = "setuptools.build_meta"
|
| 4 |
+
|
| 5 |
+
[project]
|
| 6 |
+
name = "delivery-openenv"
|
| 7 |
+
version = "1.0.0"
|
| 8 |
+
description = "OpenEnv-compatible last-mile delivery optimization benchmark environment."
|
| 9 |
+
readme = "README.md"
|
| 10 |
+
requires-python = ">=3.10"
|
| 11 |
+
dependencies = [
|
| 12 |
+
"fastapi==0.115.2",
|
| 13 |
+
"uvicorn==0.30.6",
|
| 14 |
+
"pydantic==2.9.2",
|
| 15 |
+
"openai==1.51.2",
|
| 16 |
+
"openenv-core",
|
| 17 |
+
]
|
| 18 |
+
|
| 19 |
+
[project.scripts]
|
| 20 |
+
server = "server.app:main"
|
q_agent_easy.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
requirements.txt
CHANGED
|
@@ -1,2 +1,6 @@
|
|
| 1 |
-
fastapi
|
| 2 |
-
uvicorn
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
fastapi==0.115.2
|
| 2 |
+
uvicorn==0.30.6
|
| 3 |
+
pydantic==2.9.2
|
| 4 |
+
openai==1.51.2
|
| 5 |
+
openenv-core
|
| 6 |
+
pytest==8.3.3
|
scripts/train_q_agent.py
ADDED
|
@@ -0,0 +1,230 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import argparse
|
| 4 |
+
import json
|
| 5 |
+
import os
|
| 6 |
+
import random
|
| 7 |
+
import sys
|
| 8 |
+
from dataclasses import dataclass
|
| 9 |
+
from pathlib import Path
|
| 10 |
+
from typing import Dict, List
|
| 11 |
+
|
| 12 |
+
REPO_ROOT = Path(__file__).resolve().parents[1]
|
| 13 |
+
if str(REPO_ROOT) not in sys.path:
|
| 14 |
+
sys.path.insert(0, str(REPO_ROOT))
|
| 15 |
+
|
| 16 |
+
from baseline.trained_q_agent import NUM_ACTIONS, QTable, TrainedQAgent, action_from_index, observation_key
|
| 17 |
+
from env.environment import LastMileDeliveryEnvironment
|
| 18 |
+
from tasks.registry import get_task_config
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
@dataclass
|
| 22 |
+
class EpisodeStats:
|
| 23 |
+
total_reward: float
|
| 24 |
+
completion_rate: float
|
| 25 |
+
success: bool
|
| 26 |
+
steps: int
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def _new_q_row() -> List[float]:
|
| 30 |
+
return [0.0 for _ in range(NUM_ACTIONS)]
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def _argmax(values: List[float]) -> int:
|
| 34 |
+
best_index = 0
|
| 35 |
+
best_value = values[0]
|
| 36 |
+
for index in range(1, len(values)):
|
| 37 |
+
if values[index] > best_value:
|
| 38 |
+
best_value = values[index]
|
| 39 |
+
best_index = index
|
| 40 |
+
return best_index
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def _episode_summary(env: LastMileDeliveryEnvironment) -> EpisodeStats:
|
| 44 |
+
state = env.state()
|
| 45 |
+
observation = state.observation
|
| 46 |
+
if observation is None:
|
| 47 |
+
return EpisodeStats(total_reward=state.total_reward, completion_rate=0.0, success=False, steps=state.step_count)
|
| 48 |
+
|
| 49 |
+
delivered = state.delivered_orders
|
| 50 |
+
remaining = len(observation.pending_orders)
|
| 51 |
+
if observation.current_order is not None:
|
| 52 |
+
remaining += 1
|
| 53 |
+
|
| 54 |
+
total_orders = max(1, delivered + remaining)
|
| 55 |
+
completion_rate = delivered / total_orders
|
| 56 |
+
success = bool(state.done and remaining == 0)
|
| 57 |
+
return EpisodeStats(
|
| 58 |
+
total_reward=state.total_reward,
|
| 59 |
+
completion_rate=completion_rate,
|
| 60 |
+
success=success,
|
| 61 |
+
steps=state.step_count,
|
| 62 |
+
)
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
def train_q_agent(
|
| 66 |
+
task: str,
|
| 67 |
+
episodes: int,
|
| 68 |
+
seed: int,
|
| 69 |
+
alpha: float,
|
| 70 |
+
gamma: float,
|
| 71 |
+
epsilon_start: float,
|
| 72 |
+
epsilon_end: float,
|
| 73 |
+
epsilon_decay: float,
|
| 74 |
+
) -> QTable:
|
| 75 |
+
rng = random.Random(seed)
|
| 76 |
+
q_table: QTable = {}
|
| 77 |
+
|
| 78 |
+
running_reward = 0.0
|
| 79 |
+
running_completion = 0.0
|
| 80 |
+
|
| 81 |
+
for episode in range(episodes):
|
| 82 |
+
current_seed = seed + episode
|
| 83 |
+
config = get_task_config(task_name=task, seed=current_seed)
|
| 84 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 85 |
+
observation = env.reset(seed=current_seed)
|
| 86 |
+
|
| 87 |
+
epsilon = max(epsilon_end, epsilon_start * (epsilon_decay ** episode))
|
| 88 |
+
|
| 89 |
+
for _ in range(config.max_steps):
|
| 90 |
+
state_key = observation_key(observation)
|
| 91 |
+
row = q_table.setdefault(state_key, _new_q_row())
|
| 92 |
+
|
| 93 |
+
if rng.random() < epsilon:
|
| 94 |
+
action_index = rng.randrange(NUM_ACTIONS)
|
| 95 |
+
else:
|
| 96 |
+
action_index = _argmax(row)
|
| 97 |
+
|
| 98 |
+
action = action_from_index(action_index, observation)
|
| 99 |
+
next_observation, reward, done, _ = env.step(action)
|
| 100 |
+
|
| 101 |
+
next_key = observation_key(next_observation)
|
| 102 |
+
next_row = q_table.setdefault(next_key, _new_q_row())
|
| 103 |
+
|
| 104 |
+
td_target = reward if done else reward + gamma * max(next_row)
|
| 105 |
+
row[action_index] += alpha * (td_target - row[action_index])
|
| 106 |
+
|
| 107 |
+
observation = next_observation
|
| 108 |
+
if done:
|
| 109 |
+
break
|
| 110 |
+
|
| 111 |
+
summary = _episode_summary(env)
|
| 112 |
+
running_reward += summary.total_reward
|
| 113 |
+
running_completion += summary.completion_rate
|
| 114 |
+
|
| 115 |
+
if (episode + 1) % max(1, episodes // 10) == 0:
|
| 116 |
+
avg_reward = running_reward / (episode + 1)
|
| 117 |
+
avg_completion = running_completion / (episode + 1)
|
| 118 |
+
print(
|
| 119 |
+
f"[TRAIN] episode={episode + 1} "
|
| 120 |
+
f"epsilon={epsilon:.4f} "
|
| 121 |
+
f"avg_reward={avg_reward:.2f} "
|
| 122 |
+
f"avg_completion={avg_completion:.3f}",
|
| 123 |
+
flush=True,
|
| 124 |
+
)
|
| 125 |
+
|
| 126 |
+
return q_table
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
def evaluate_q_agent(task: str, q_table: QTable, episodes: int, seed: int) -> Dict[str, float]:
|
| 130 |
+
agent = TrainedQAgent(q_table=q_table)
|
| 131 |
+
|
| 132 |
+
total_reward = 0.0
|
| 133 |
+
total_completion = 0.0
|
| 134 |
+
success_count = 0
|
| 135 |
+
total_steps = 0
|
| 136 |
+
|
| 137 |
+
for episode in range(episodes):
|
| 138 |
+
current_seed = seed + 10_000 + episode
|
| 139 |
+
config = get_task_config(task_name=task, seed=current_seed)
|
| 140 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 141 |
+
observation = env.reset(seed=current_seed)
|
| 142 |
+
|
| 143 |
+
for _ in range(config.max_steps):
|
| 144 |
+
action = agent.act(observation)
|
| 145 |
+
observation, _, done, _ = env.step(action)
|
| 146 |
+
if done:
|
| 147 |
+
break
|
| 148 |
+
|
| 149 |
+
summary = _episode_summary(env)
|
| 150 |
+
total_reward += summary.total_reward
|
| 151 |
+
total_completion += summary.completion_rate
|
| 152 |
+
success_count += 1 if summary.success else 0
|
| 153 |
+
total_steps += summary.steps
|
| 154 |
+
|
| 155 |
+
count = max(1, episodes)
|
| 156 |
+
return {
|
| 157 |
+
"avg_reward": total_reward / count,
|
| 158 |
+
"avg_completion": total_completion / count,
|
| 159 |
+
"success_rate": success_count / count,
|
| 160 |
+
"avg_steps": total_steps / count,
|
| 161 |
+
}
|
| 162 |
+
|
| 163 |
+
|
| 164 |
+
def save_model(model_path: str, task: str, seed: int, episodes: int, q_table: QTable) -> None:
|
| 165 |
+
output_dir = os.path.dirname(model_path)
|
| 166 |
+
if output_dir:
|
| 167 |
+
os.makedirs(output_dir, exist_ok=True)
|
| 168 |
+
payload = {
|
| 169 |
+
"model_type": "tabular_q_learning",
|
| 170 |
+
"task": task,
|
| 171 |
+
"seed": seed,
|
| 172 |
+
"episodes": episodes,
|
| 173 |
+
"num_actions": NUM_ACTIONS,
|
| 174 |
+
"q_values": q_table,
|
| 175 |
+
}
|
| 176 |
+
with open(model_path, "w", encoding="utf-8") as handle:
|
| 177 |
+
json.dump(payload, handle, separators=(",", ":"), sort_keys=True)
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
def parse_args() -> argparse.Namespace:
|
| 181 |
+
parser = argparse.ArgumentParser(description="Train a tabular Q-learning agent for delivery-openenv")
|
| 182 |
+
parser.add_argument("--task", type=str, default="easy", choices=["easy", "medium", "hard"])
|
| 183 |
+
parser.add_argument("--episodes", type=int, default=800)
|
| 184 |
+
parser.add_argument("--eval-episodes", type=int, default=100)
|
| 185 |
+
parser.add_argument("--seed", type=int, default=42)
|
| 186 |
+
parser.add_argument("--alpha", type=float, default=0.20)
|
| 187 |
+
parser.add_argument("--gamma", type=float, default=0.98)
|
| 188 |
+
parser.add_argument("--epsilon-start", type=float, default=1.0)
|
| 189 |
+
parser.add_argument("--epsilon-end", type=float, default=0.05)
|
| 190 |
+
parser.add_argument("--epsilon-decay", type=float, default=0.995)
|
| 191 |
+
parser.add_argument("--output", type=str, default="models/q_agent_easy.json")
|
| 192 |
+
return parser.parse_args()
|
| 193 |
+
|
| 194 |
+
|
| 195 |
+
def main() -> None:
|
| 196 |
+
args = parse_args()
|
| 197 |
+
|
| 198 |
+
q_table = train_q_agent(
|
| 199 |
+
task=args.task,
|
| 200 |
+
episodes=args.episodes,
|
| 201 |
+
seed=args.seed,
|
| 202 |
+
alpha=args.alpha,
|
| 203 |
+
gamma=args.gamma,
|
| 204 |
+
epsilon_start=args.epsilon_start,
|
| 205 |
+
epsilon_end=args.epsilon_end,
|
| 206 |
+
epsilon_decay=args.epsilon_decay,
|
| 207 |
+
)
|
| 208 |
+
|
| 209 |
+
metrics = evaluate_q_agent(
|
| 210 |
+
task=args.task,
|
| 211 |
+
q_table=q_table,
|
| 212 |
+
episodes=args.eval_episodes,
|
| 213 |
+
seed=args.seed,
|
| 214 |
+
)
|
| 215 |
+
|
| 216 |
+
save_model(model_path=args.output, task=args.task, seed=args.seed, episodes=args.episodes, q_table=q_table)
|
| 217 |
+
|
| 218 |
+
print(
|
| 219 |
+
f"[EVAL] task={args.task} episodes={args.eval_episodes} "
|
| 220 |
+
f"avg_reward={metrics['avg_reward']:.2f} "
|
| 221 |
+
f"avg_completion={metrics['avg_completion']:.3f} "
|
| 222 |
+
f"success_rate={metrics['success_rate']:.3f} "
|
| 223 |
+
f"avg_steps={metrics['avg_steps']:.2f}",
|
| 224 |
+
flush=True,
|
| 225 |
+
)
|
| 226 |
+
print(f"[MODEL] saved={args.output} states={len(q_table)}", flush=True)
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
if __name__ == "__main__":
|
| 230 |
+
main()
|
scripts/validate.sh
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env bash
|
| 2 |
+
set -euo pipefail
|
| 3 |
+
|
| 4 |
+
REPO_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
|
| 5 |
+
PING_URL="${PING_URL:-}"
|
| 6 |
+
DOCKER_IMAGE_TAG="${DOCKER_IMAGE_TAG:-delivery-openenv:precheck}"
|
| 7 |
+
OPENENV_BIN_FALLBACK="$REPO_DIR/.venv/bin/openenv"
|
| 8 |
+
VENV_PYTHON="$REPO_DIR/.venv/bin/python"
|
| 9 |
+
|
| 10 |
+
pass() { echo "[PASS] $1"; }
|
| 11 |
+
info() { echo "[INFO] $1"; }
|
| 12 |
+
fail() { echo "[FAIL] $1"; exit 1; }
|
| 13 |
+
|
| 14 |
+
info "Step 1/4: Optional Hugging Face Space ping"
|
| 15 |
+
if [[ -n "$PING_URL" ]]; then
|
| 16 |
+
HTTP_CODE="$(
|
| 17 |
+
curl -sS -o /tmp/openenv_reset_response.txt -w "%{http_code}" \
|
| 18 |
+
-X POST "$PING_URL/reset" \
|
| 19 |
+
-H "Content-Type: application/json" \
|
| 20 |
+
-d '{"task":"easy","seed":101}' || true
|
| 21 |
+
)"
|
| 22 |
+
if [[ "$HTTP_CODE" == "200" ]]; then
|
| 23 |
+
pass "HF Space /reset returned HTTP 200"
|
| 24 |
+
else
|
| 25 |
+
fail "HF Space /reset returned HTTP $HTTP_CODE"
|
| 26 |
+
fi
|
| 27 |
+
else
|
| 28 |
+
info "PING_URL is not set; skipping remote endpoint check"
|
| 29 |
+
fi
|
| 30 |
+
|
| 31 |
+
info "Step 2/4: Docker build"
|
| 32 |
+
if ! command -v docker >/dev/null 2>&1; then
|
| 33 |
+
fail "docker command not found"
|
| 34 |
+
fi
|
| 35 |
+
|
| 36 |
+
if ! docker info >/dev/null 2>&1; then
|
| 37 |
+
fail "docker daemon is not running or not reachable"
|
| 38 |
+
fi
|
| 39 |
+
|
| 40 |
+
docker build -t "$DOCKER_IMAGE_TAG" "$REPO_DIR" >/tmp/openenv_docker_build.log 2>&1 || {
|
| 41 |
+
tail -20 /tmp/openenv_docker_build.log || true
|
| 42 |
+
fail "docker build failed"
|
| 43 |
+
}
|
| 44 |
+
pass "Docker build succeeded"
|
| 45 |
+
|
| 46 |
+
info "Step 3/4: OpenEnv validation"
|
| 47 |
+
if command -v openenv >/dev/null 2>&1; then
|
| 48 |
+
OPENENV_CMD=(openenv)
|
| 49 |
+
elif [[ -x "$OPENENV_BIN_FALLBACK" ]]; then
|
| 50 |
+
OPENENV_CMD=("$OPENENV_BIN_FALLBACK")
|
| 51 |
+
else
|
| 52 |
+
fail "openenv CLI not found; install with 'pip install openenv-core'"
|
| 53 |
+
fi
|
| 54 |
+
|
| 55 |
+
(cd "$REPO_DIR" && "${OPENENV_CMD[@]}" validate)
|
| 56 |
+
pass "openenv validate passed"
|
| 57 |
+
|
| 58 |
+
info "Step 4/4: Test suite"
|
| 59 |
+
if [[ -x "$VENV_PYTHON" ]]; then
|
| 60 |
+
PYTHON_BIN="$VENV_PYTHON"
|
| 61 |
+
elif command -v python3 >/dev/null 2>&1; then
|
| 62 |
+
PYTHON_BIN="$(command -v python3)"
|
| 63 |
+
else
|
| 64 |
+
fail "python3 not found"
|
| 65 |
+
fi
|
| 66 |
+
|
| 67 |
+
(cd "$REPO_DIR" && "$PYTHON_BIN" -m pytest -q)
|
| 68 |
+
pass "pytest passed"
|
| 69 |
+
|
| 70 |
+
echo
|
| 71 |
+
echo "[PASS] All validation checks completed"
|
server/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
"""Server entrypoint package for OpenEnv multi-mode deployment."""
|
server/app.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import os
|
| 4 |
+
|
| 5 |
+
import uvicorn
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
def main() -> None:
|
| 9 |
+
host = os.getenv("HOST", "0.0.0.0")
|
| 10 |
+
port_raw = os.getenv("PORT", "7860")
|
| 11 |
+
try:
|
| 12 |
+
port = int(port_raw)
|
| 13 |
+
except ValueError:
|
| 14 |
+
port = 7860
|
| 15 |
+
|
| 16 |
+
uvicorn.run("app.main:app", host=host, port=port, reload=False)
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
if __name__ == "__main__":
|
| 20 |
+
main()
|
tasks/__init__.py
ADDED
|
File without changes
|
tasks/easy.py
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from typing import Any, Dict
|
| 4 |
+
|
| 5 |
+
from env.models import EnvironmentConfig
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
TASK_NAME = "easy"
|
| 9 |
+
TASK_DIFFICULTY = 1
|
| 10 |
+
DEFAULT_SEED = 101
|
| 11 |
+
DESCRIPTION = "Simple single-order delivery on a compact map without operational disruptions."
|
| 12 |
+
OBJECTIVE = "Learn reliable movement, pickup, and dropoff sequencing in the simplest delivery setting."
|
| 13 |
+
SUCCESS_CONDITION: Dict[str, Any] = {
|
| 14 |
+
"completion_rate": 1.0,
|
| 15 |
+
"max_steps": 60,
|
| 16 |
+
"battery_depletion": False,
|
| 17 |
+
"invalid_action_rate_max": 0.12,
|
| 18 |
+
}
|
| 19 |
+
|
| 20 |
+
TASK_PARAMETERS: Dict[str, Any] = {
|
| 21 |
+
"grid_size": {"width": 6, "height": 6},
|
| 22 |
+
"orders": {"min": 1, "max": 1},
|
| 23 |
+
"variation": {
|
| 24 |
+
"seeded_layout": True,
|
| 25 |
+
"order_pattern": "single_short_trip",
|
| 26 |
+
},
|
| 27 |
+
"constraints": {
|
| 28 |
+
"obstacles": False,
|
| 29 |
+
"dynamic_obstacles": False,
|
| 30 |
+
"traffic_zones": False,
|
| 31 |
+
"order_priority": False,
|
| 32 |
+
"multiple_delivery_locations": False,
|
| 33 |
+
"battery_constraint": False,
|
| 34 |
+
"recharge": False,
|
| 35 |
+
"time_penalty": False,
|
| 36 |
+
},
|
| 37 |
+
}
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def get_config(seed: int | None = None) -> EnvironmentConfig:
|
| 41 |
+
effective_seed = DEFAULT_SEED if seed is None else seed
|
| 42 |
+
return EnvironmentConfig(
|
| 43 |
+
width=6,
|
| 44 |
+
height=6,
|
| 45 |
+
max_steps=60,
|
| 46 |
+
max_orders=1,
|
| 47 |
+
delivery_locations_per_order=1,
|
| 48 |
+
priority_high_ratio=0.0,
|
| 49 |
+
obstacle_density=0.0,
|
| 50 |
+
dynamic_obstacles_enabled=False,
|
| 51 |
+
dynamic_obstacle_ratio=0.0,
|
| 52 |
+
traffic_density=0.0,
|
| 53 |
+
traffic_extra_cost=1,
|
| 54 |
+
battery_enabled=False,
|
| 55 |
+
battery_capacity=100,
|
| 56 |
+
battery_recharge_rate=10,
|
| 57 |
+
seed=effective_seed,
|
| 58 |
+
)
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def get_definition() -> Dict[str, Any]:
|
| 62 |
+
return {
|
| 63 |
+
"name": TASK_NAME,
|
| 64 |
+
"difficulty": TASK_DIFFICULTY,
|
| 65 |
+
"deterministic": True,
|
| 66 |
+
"default_seed": DEFAULT_SEED,
|
| 67 |
+
"description": DESCRIPTION,
|
| 68 |
+
"objective": OBJECTIVE,
|
| 69 |
+
"success_condition": SUCCESS_CONDITION,
|
| 70 |
+
"parameters": TASK_PARAMETERS,
|
| 71 |
+
}
|
tasks/hard.py
ADDED
|
@@ -0,0 +1,79 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from typing import Any, Dict
|
| 4 |
+
|
| 5 |
+
from env.models import EnvironmentConfig
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
TASK_NAME = "hard"
|
| 9 |
+
TASK_DIFFICULTY = 3
|
| 10 |
+
DEFAULT_SEED = 303
|
| 11 |
+
DESCRIPTION = "High-pressure dispatch benchmark with strict energy limits, priority SLAs, and tight step budget."
|
| 12 |
+
OBJECTIVE = "Sustain high on-time completion for priority-heavy workloads under battery, traffic, and dynamic obstacle pressure."
|
| 13 |
+
SUCCESS_CONDITION: Dict[str, Any] = {
|
| 14 |
+
"completion_rate_min": 0.90,
|
| 15 |
+
"high_priority_on_time_rate_min": 0.88,
|
| 16 |
+
"battery_depletion": False,
|
| 17 |
+
"max_steps": 165,
|
| 18 |
+
"invalid_action_rate_max": 0.04,
|
| 19 |
+
}
|
| 20 |
+
|
| 21 |
+
TASK_PARAMETERS: Dict[str, Any] = {
|
| 22 |
+
"grid_size": {"width": 16, "height": 16},
|
| 23 |
+
"orders": {"min": 6, "max": 8},
|
| 24 |
+
"variation": {
|
| 25 |
+
"seeded_layout": True,
|
| 26 |
+
"order_mix": "priority_heavy",
|
| 27 |
+
"traffic_profile": "peak_hour",
|
| 28 |
+
"battery_profile": "tight_capacity",
|
| 29 |
+
"step_budget": "tight",
|
| 30 |
+
},
|
| 31 |
+
"constraints": {
|
| 32 |
+
"obstacles": True,
|
| 33 |
+
"dynamic_obstacles": True,
|
| 34 |
+
"traffic_zones": True,
|
| 35 |
+
"order_priority": True,
|
| 36 |
+
"multiple_delivery_locations": True,
|
| 37 |
+
"battery_constraint": True,
|
| 38 |
+
"recharge": True,
|
| 39 |
+
"time_penalty": True,
|
| 40 |
+
},
|
| 41 |
+
}
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def get_config(seed: int | None = None) -> EnvironmentConfig:
|
| 45 |
+
effective_seed = DEFAULT_SEED if seed is None else seed
|
| 46 |
+
order_count = 6 + (effective_seed % 3)
|
| 47 |
+
delivery_locations = 2 + (effective_seed % 2)
|
| 48 |
+
dynamic_ratio = 0.50 + (0.05 * (effective_seed % 3))
|
| 49 |
+
|
| 50 |
+
return EnvironmentConfig(
|
| 51 |
+
width=16,
|
| 52 |
+
height=16,
|
| 53 |
+
max_steps=165,
|
| 54 |
+
max_orders=order_count,
|
| 55 |
+
delivery_locations_per_order=delivery_locations,
|
| 56 |
+
priority_high_ratio=0.75,
|
| 57 |
+
obstacle_density=0.20,
|
| 58 |
+
dynamic_obstacles_enabled=True,
|
| 59 |
+
dynamic_obstacle_ratio=min(0.65, dynamic_ratio),
|
| 60 |
+
traffic_density=0.22,
|
| 61 |
+
traffic_extra_cost=5,
|
| 62 |
+
battery_enabled=True,
|
| 63 |
+
battery_capacity=55,
|
| 64 |
+
battery_recharge_rate=8,
|
| 65 |
+
seed=effective_seed,
|
| 66 |
+
)
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
def get_definition() -> Dict[str, Any]:
|
| 70 |
+
return {
|
| 71 |
+
"name": TASK_NAME,
|
| 72 |
+
"difficulty": TASK_DIFFICULTY,
|
| 73 |
+
"deterministic": True,
|
| 74 |
+
"default_seed": DEFAULT_SEED,
|
| 75 |
+
"description": DESCRIPTION,
|
| 76 |
+
"objective": OBJECTIVE,
|
| 77 |
+
"success_condition": SUCCESS_CONDITION,
|
| 78 |
+
"parameters": TASK_PARAMETERS,
|
| 79 |
+
}
|
tasks/medium.py
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from typing import Any, Dict
|
| 4 |
+
|
| 5 |
+
from env.models import EnvironmentConfig
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
TASK_NAME = "medium"
|
| 9 |
+
TASK_DIFFICULTY = 2
|
| 10 |
+
DEFAULT_SEED = 202
|
| 11 |
+
DESCRIPTION = "Multi-order delivery with static obstacles and realistic traffic friction."
|
| 12 |
+
OBJECTIVE = "Optimize route throughput across multiple concurrent orders while navigating blocked streets."
|
| 13 |
+
SUCCESS_CONDITION: Dict[str, Any] = {
|
| 14 |
+
"completion_rate_min": 0.85,
|
| 15 |
+
"max_steps": 140,
|
| 16 |
+
"invalid_action_rate_max": 0.08,
|
| 17 |
+
}
|
| 18 |
+
|
| 19 |
+
TASK_PARAMETERS: Dict[str, Any] = {
|
| 20 |
+
"grid_size": {"width": 10, "height": 10},
|
| 21 |
+
"orders": {"min": 3, "max": 4},
|
| 22 |
+
"variation": {
|
| 23 |
+
"seeded_layout": True,
|
| 24 |
+
"order_mix": "standard_multi_order",
|
| 25 |
+
"traffic_profile": "weekday_traffic",
|
| 26 |
+
},
|
| 27 |
+
"constraints": {
|
| 28 |
+
"obstacles": True,
|
| 29 |
+
"dynamic_obstacles": False,
|
| 30 |
+
"traffic_zones": True,
|
| 31 |
+
"order_priority": False,
|
| 32 |
+
"multiple_delivery_locations": False,
|
| 33 |
+
"battery_constraint": False,
|
| 34 |
+
"recharge": False,
|
| 35 |
+
"time_penalty": False,
|
| 36 |
+
},
|
| 37 |
+
}
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def get_config(seed: int | None = None) -> EnvironmentConfig:
|
| 41 |
+
effective_seed = DEFAULT_SEED if seed is None else seed
|
| 42 |
+
order_count = 3 + (effective_seed % 2)
|
| 43 |
+
traffic_cost = 2 + (effective_seed % 2)
|
| 44 |
+
|
| 45 |
+
return EnvironmentConfig(
|
| 46 |
+
width=10,
|
| 47 |
+
height=10,
|
| 48 |
+
max_steps=140,
|
| 49 |
+
max_orders=order_count,
|
| 50 |
+
delivery_locations_per_order=1,
|
| 51 |
+
priority_high_ratio=0.0,
|
| 52 |
+
obstacle_density=0.12,
|
| 53 |
+
dynamic_obstacles_enabled=False,
|
| 54 |
+
dynamic_obstacle_ratio=0.0,
|
| 55 |
+
traffic_density=0.11,
|
| 56 |
+
traffic_extra_cost=traffic_cost,
|
| 57 |
+
battery_enabled=False,
|
| 58 |
+
battery_capacity=100,
|
| 59 |
+
battery_recharge_rate=10,
|
| 60 |
+
seed=effective_seed,
|
| 61 |
+
)
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def get_definition() -> Dict[str, Any]:
|
| 65 |
+
return {
|
| 66 |
+
"name": TASK_NAME,
|
| 67 |
+
"difficulty": TASK_DIFFICULTY,
|
| 68 |
+
"deterministic": True,
|
| 69 |
+
"default_seed": DEFAULT_SEED,
|
| 70 |
+
"description": DESCRIPTION,
|
| 71 |
+
"objective": OBJECTIVE,
|
| 72 |
+
"success_condition": SUCCESS_CONDITION,
|
| 73 |
+
"parameters": TASK_PARAMETERS,
|
| 74 |
+
}
|
tasks/registry.py
ADDED
|
@@ -0,0 +1,73 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from typing import Any, Callable, Dict, List
|
| 4 |
+
|
| 5 |
+
from env.models import EnvironmentConfig
|
| 6 |
+
from tasks import easy, hard, medium
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
TaskFactory = Callable[[int | None], EnvironmentConfig]
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
TASK_REGISTRY: Dict[str, Dict[str, Any]] = {
|
| 13 |
+
easy.TASK_NAME: {
|
| 14 |
+
"definition": easy.get_definition(),
|
| 15 |
+
"factory": easy.get_config,
|
| 16 |
+
},
|
| 17 |
+
medium.TASK_NAME: {
|
| 18 |
+
"definition": medium.get_definition(),
|
| 19 |
+
"factory": medium.get_config,
|
| 20 |
+
},
|
| 21 |
+
hard.TASK_NAME: {
|
| 22 |
+
"definition": hard.get_definition(),
|
| 23 |
+
"factory": hard.get_config,
|
| 24 |
+
},
|
| 25 |
+
}
|
| 26 |
+
|
| 27 |
+
TASK_DIFFICULTY_ORDER: List[str] = [easy.TASK_NAME, medium.TASK_NAME, hard.TASK_NAME]
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def _ordered_task_names() -> List[str]:
|
| 31 |
+
return [name for name in TASK_DIFFICULTY_ORDER if name in TASK_REGISTRY]
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def list_tasks() -> List[str]:
|
| 35 |
+
return _ordered_task_names()
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def get_task_registry() -> Dict[str, Dict[str, Any]]:
|
| 39 |
+
return {
|
| 40 |
+
name: dict(item["definition"])
|
| 41 |
+
for name in _ordered_task_names()
|
| 42 |
+
for item in [TASK_REGISTRY[name]]
|
| 43 |
+
}
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def get_task_definition(task_name: str) -> Dict[str, Any]:
|
| 47 |
+
item = TASK_REGISTRY.get(task_name)
|
| 48 |
+
if item is None:
|
| 49 |
+
raise ValueError(f"Unknown task '{task_name}'. Available: {list_tasks()}")
|
| 50 |
+
return dict(item["definition"])
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
def get_task_config(task_name: str, seed: int | None = None) -> EnvironmentConfig:
|
| 54 |
+
item = TASK_REGISTRY.get(task_name)
|
| 55 |
+
if item is None:
|
| 56 |
+
raise ValueError(f"Unknown task '{task_name}'. Available: {list_tasks()}")
|
| 57 |
+
|
| 58 |
+
factory = item["factory"]
|
| 59 |
+
if not callable(factory):
|
| 60 |
+
raise RuntimeError(f"Task '{task_name}' has invalid factory")
|
| 61 |
+
return factory(seed)
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def get_all_task_metadata() -> List[Dict[str, Any]]:
|
| 65 |
+
return [
|
| 66 |
+
dict(TASK_REGISTRY[name]["definition"])
|
| 67 |
+
for name in _ordered_task_names()
|
| 68 |
+
]
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def list_tasks_with_metadata() -> List[Dict[str, Any]]:
|
| 72 |
+
"""Compatibility alias for tooling that expects this function name."""
|
| 73 |
+
return get_all_task_metadata()
|
tests/test_api.py
ADDED
|
@@ -0,0 +1,263 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from fastapi.testclient import TestClient
|
| 4 |
+
|
| 5 |
+
from app.main import app
|
| 6 |
+
from app.routes import env_service
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
client = TestClient(app)
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def test_root_endpoint_returns_service_info() -> None:
|
| 13 |
+
response = client.get("/")
|
| 14 |
+
assert response.status_code == 200
|
| 15 |
+
payload = response.json()
|
| 16 |
+
assert payload["status"] == "ok"
|
| 17 |
+
assert payload["health"] == "/health"
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def test_tasks_endpoint_returns_json() -> None:
|
| 21 |
+
response = client.get("/tasks")
|
| 22 |
+
assert response.status_code == 200
|
| 23 |
+
payload = response.json()
|
| 24 |
+
assert "tasks" in payload
|
| 25 |
+
assert "action_schema" in payload
|
| 26 |
+
assert isinstance(payload["tasks"], list)
|
| 27 |
+
assert isinstance(payload["action_schema"], dict)
|
| 28 |
+
|
| 29 |
+
for task in payload["tasks"]:
|
| 30 |
+
assert "name" in task
|
| 31 |
+
assert "description" in task
|
| 32 |
+
assert "difficulty" in task
|
| 33 |
+
assert "parameters" in task
|
| 34 |
+
assert "success_condition" in task
|
| 35 |
+
|
| 36 |
+
action_schema = payload["action_schema"]
|
| 37 |
+
assert action_schema["type"] == "object"
|
| 38 |
+
assert action_schema["required"] == ["move", "accept_order", "deliver_order", "wait"]
|
| 39 |
+
assert action_schema["additionalProperties"] is False
|
| 40 |
+
assert "properties" in action_schema
|
| 41 |
+
assert "move" in action_schema["properties"]
|
| 42 |
+
assert "accept_order" in action_schema["properties"]
|
| 43 |
+
assert "deliver_order" in action_schema["properties"]
|
| 44 |
+
assert "wait" in action_schema["properties"]
|
| 45 |
+
assert set(action_schema["properties"]["move"]["enum"]) == {
|
| 46 |
+
"up",
|
| 47 |
+
"down",
|
| 48 |
+
"left",
|
| 49 |
+
"right",
|
| 50 |
+
"stay",
|
| 51 |
+
None,
|
| 52 |
+
}
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def test_state_before_reset_returns_uninitialized_state() -> None:
|
| 56 |
+
env_service._env = None
|
| 57 |
+
env_service._task = None
|
| 58 |
+
env_service._success_condition = {}
|
| 59 |
+
env_service._trajectory = []
|
| 60 |
+
response = client.get("/state")
|
| 61 |
+
assert response.status_code == 200
|
| 62 |
+
payload = response.json()
|
| 63 |
+
assert payload["initialized"] is False
|
| 64 |
+
assert payload["done"] is False
|
| 65 |
+
assert payload["step_count"] == 0
|
| 66 |
+
assert payload["observation"] is None
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
def test_reset_step_state_baseline_and_grader_flow() -> None:
|
| 70 |
+
reset_response = client.post("/reset", json={"task": "easy", "seed": 7})
|
| 71 |
+
assert reset_response.status_code == 200
|
| 72 |
+
|
| 73 |
+
baseline_response = client.get("/baseline", params={"task": "easy", "seed": 7, "max_steps": 20})
|
| 74 |
+
assert baseline_response.status_code == 200
|
| 75 |
+
baseline_payload = baseline_response.json()
|
| 76 |
+
assert baseline_payload["agent"] == "BaselineGreedyAgent"
|
| 77 |
+
assert baseline_payload["task"] == "easy"
|
| 78 |
+
assert baseline_payload["steps_executed"] <= 20
|
| 79 |
+
assert 0.0 <= baseline_payload["score"] <= 1.0
|
| 80 |
+
assert "report" in baseline_payload
|
| 81 |
+
|
| 82 |
+
step_response = client.post(
|
| 83 |
+
"/step",
|
| 84 |
+
json={"move": None, "accept_order": None, "deliver_order": False, "wait": True},
|
| 85 |
+
)
|
| 86 |
+
assert step_response.status_code == 200
|
| 87 |
+
|
| 88 |
+
state_response = client.get("/state")
|
| 89 |
+
assert state_response.status_code == 200
|
| 90 |
+
|
| 91 |
+
grader_response = client.get("/grader")
|
| 92 |
+
assert grader_response.status_code == 200
|
| 93 |
+
grader_payload = grader_response.json()
|
| 94 |
+
assert "report" in grader_payload
|
| 95 |
+
assert 0.0 <= grader_payload["report"]["score"] <= 1.0
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
def test_reset_without_body_uses_defaults() -> None:
|
| 99 |
+
response = client.post("/reset")
|
| 100 |
+
assert response.status_code == 200
|
| 101 |
+
payload = response.json()
|
| 102 |
+
assert "agent_location" in payload
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def test_step_accepts_direct_action_payload() -> None:
|
| 106 |
+
client.post("/reset", json={"task": "easy", "seed": 7})
|
| 107 |
+
response = client.post(
|
| 108 |
+
"/step",
|
| 109 |
+
json={"move": None, "accept_order": None, "deliver_order": False, "wait": True},
|
| 110 |
+
)
|
| 111 |
+
assert response.status_code == 200
|
| 112 |
+
payload = response.json()
|
| 113 |
+
assert {"observation", "reward", "done", "info"}.issubset(payload.keys())
|
| 114 |
+
assert isinstance(payload["reward"], float)
|
| 115 |
+
assert isinstance(payload["done"], bool)
|
| 116 |
+
assert isinstance(payload["info"], dict)
|
| 117 |
+
|
| 118 |
+
|
| 119 |
+
def test_step_rejects_wrapper_payload_for_strict_schema() -> None:
|
| 120 |
+
client.post("/reset", json={"task": "easy", "seed": 7})
|
| 121 |
+
response = client.post(
|
| 122 |
+
"/step",
|
| 123 |
+
json={
|
| 124 |
+
"action": {
|
| 125 |
+
"move": None,
|
| 126 |
+
"accept_order": None,
|
| 127 |
+
"deliver_order": False,
|
| 128 |
+
"wait": True,
|
| 129 |
+
}
|
| 130 |
+
},
|
| 131 |
+
)
|
| 132 |
+
assert response.status_code == 400
|
| 133 |
+
payload = response.json()
|
| 134 |
+
assert payload["detail"]["error"] == "bad_request"
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
def test_step_invalid_payload_returns_clean_error() -> None:
|
| 138 |
+
client.post("/reset", json={"task": "easy", "seed": 7})
|
| 139 |
+
response = client.post("/step", json={"move": "right"})
|
| 140 |
+
assert response.status_code == 400
|
| 141 |
+
payload = response.json()
|
| 142 |
+
assert payload["detail"]["error"] == "bad_request"
|
| 143 |
+
|
| 144 |
+
|
| 145 |
+
def test_grader_before_reset_returns_clean_error() -> None:
|
| 146 |
+
env_service._env = None
|
| 147 |
+
env_service._task = None
|
| 148 |
+
env_service._success_condition = {}
|
| 149 |
+
env_service._trajectory = []
|
| 150 |
+
response = client.get("/grader")
|
| 151 |
+
assert response.status_code == 400
|
| 152 |
+
payload = response.json()
|
| 153 |
+
assert payload["detail"]["error"] == "bad_request"
|
| 154 |
+
|
| 155 |
+
|
| 156 |
+
def test_baseline_runs_without_active_environment() -> None:
|
| 157 |
+
env_service._env = None
|
| 158 |
+
env_service._task = None
|
| 159 |
+
env_service._success_condition = {}
|
| 160 |
+
env_service._trajectory = []
|
| 161 |
+
|
| 162 |
+
response = client.get("/baseline", params={"task": "easy", "seed": 1, "max_steps": 10})
|
| 163 |
+
assert response.status_code == 200
|
| 164 |
+
|
| 165 |
+
payload = response.json()
|
| 166 |
+
assert payload["agent"] == "BaselineGreedyAgent"
|
| 167 |
+
assert payload["task"] == "easy"
|
| 168 |
+
assert payload["steps_executed"] <= 10
|
| 169 |
+
assert isinstance(payload["done"], bool)
|
| 170 |
+
assert 0.0 <= payload["score"] <= 1.0
|
| 171 |
+
|
| 172 |
+
|
| 173 |
+
def test_baseline_invalid_task_returns_clean_error() -> None:
|
| 174 |
+
response = client.get("/baseline", params={"task": "unknown_task"})
|
| 175 |
+
assert response.status_code == 400
|
| 176 |
+
payload = response.json()
|
| 177 |
+
assert payload["detail"]["error"] == "bad_request"
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
def test_reset_from_scenario_accepts_real_world_payload() -> None:
|
| 181 |
+
response = client.post(
|
| 182 |
+
"/reset_from_scenario",
|
| 183 |
+
json={
|
| 184 |
+
"seed": 99,
|
| 185 |
+
"success_condition": {
|
| 186 |
+
"completion_rate_min": 1.0,
|
| 187 |
+
"max_steps": 20,
|
| 188 |
+
"invalid_action_rate_max": 0.2,
|
| 189 |
+
"battery_depletion": False,
|
| 190 |
+
},
|
| 191 |
+
"scenario": {
|
| 192 |
+
"width": 6,
|
| 193 |
+
"height": 6,
|
| 194 |
+
"max_steps": 20,
|
| 195 |
+
"agent_start": {"x": 0, "y": 0},
|
| 196 |
+
"orders": [
|
| 197 |
+
{
|
| 198 |
+
"order_id": "r1",
|
| 199 |
+
"pickup": {"x": 1, "y": 0},
|
| 200 |
+
"dropoff": {"x": 2, "y": 0},
|
| 201 |
+
"delivery_locations": [{"x": 2, "y": 0}],
|
| 202 |
+
"priority": "high",
|
| 203 |
+
}
|
| 204 |
+
],
|
| 205 |
+
"obstacles": [{"x": 4, "y": 4}],
|
| 206 |
+
"dynamic_obstacles": [{"x": 4, "y": 5}],
|
| 207 |
+
"traffic_zones": [{"location": {"x": 3, "y": 0}, "extra_cost": 3}],
|
| 208 |
+
"charging_stations": [{"x": 0, "y": 0}],
|
| 209 |
+
"battery_profile": {
|
| 210 |
+
"enabled": True,
|
| 211 |
+
"capacity": 10,
|
| 212 |
+
"recharge_rate": 4,
|
| 213 |
+
"initial_level": 7,
|
| 214 |
+
},
|
| 215 |
+
},
|
| 216 |
+
},
|
| 217 |
+
)
|
| 218 |
+
assert response.status_code == 200
|
| 219 |
+
payload = response.json()
|
| 220 |
+
assert payload["grid_width"] == 6
|
| 221 |
+
assert payload["grid_height"] == 6
|
| 222 |
+
assert payload["battery_level"] == 7
|
| 223 |
+
assert len(payload["pending_orders"]) == 1
|
| 224 |
+
assert sorted((item["x"], item["y"]) for item in payload["dynamic_obstacles"]) == [(4, 5)]
|
| 225 |
+
|
| 226 |
+
step_response = client.post(
|
| 227 |
+
"/step",
|
| 228 |
+
json={"move": None, "accept_order": None, "deliver_order": False, "wait": True},
|
| 229 |
+
)
|
| 230 |
+
assert step_response.status_code == 200
|
| 231 |
+
|
| 232 |
+
grader_response = client.get("/grader")
|
| 233 |
+
assert grader_response.status_code == 200
|
| 234 |
+
grader_payload = grader_response.json()
|
| 235 |
+
assert grader_payload["task"] == "scenario"
|
| 236 |
+
assert grader_payload["report"]["success_condition_used"]["max_steps"] == 20
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
def test_reset_from_scenario_rejects_invalid_world_geometry() -> None:
|
| 240 |
+
response = client.post(
|
| 241 |
+
"/reset_from_scenario",
|
| 242 |
+
json={
|
| 243 |
+
"scenario": {
|
| 244 |
+
"width": 6,
|
| 245 |
+
"height": 6,
|
| 246 |
+
"max_steps": 20,
|
| 247 |
+
"orders": [
|
| 248 |
+
{
|
| 249 |
+
"pickup": {"x": 1, "y": 1},
|
| 250 |
+
"dropoff": {"x": 2, "y": 2},
|
| 251 |
+
}
|
| 252 |
+
],
|
| 253 |
+
"obstacles": [{"x": 9, "y": 9}],
|
| 254 |
+
"dynamic_obstacles": [],
|
| 255 |
+
"traffic_zones": [],
|
| 256 |
+
"charging_stations": [{"x": 0, "y": 0}],
|
| 257 |
+
"battery_profile": {"enabled": False},
|
| 258 |
+
}
|
| 259 |
+
},
|
| 260 |
+
)
|
| 261 |
+
assert response.status_code == 400
|
| 262 |
+
payload = response.json()
|
| 263 |
+
assert payload["detail"]["error"] == "bad_request"
|
tests/test_env.py
ADDED
|
@@ -0,0 +1,177 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from env.environment import LastMileDeliveryEnvironment
|
| 4 |
+
from env.models import ActionType, EnvironmentConfig, RewardConfig, ScenarioDefinition, SimulatorAction
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
def build_env() -> LastMileDeliveryEnvironment:
|
| 8 |
+
config = EnvironmentConfig(
|
| 9 |
+
width=6,
|
| 10 |
+
height=6,
|
| 11 |
+
max_steps=50,
|
| 12 |
+
max_orders=2,
|
| 13 |
+
obstacle_density=0.0,
|
| 14 |
+
traffic_density=0.0,
|
| 15 |
+
battery_enabled=False,
|
| 16 |
+
seed=7,
|
| 17 |
+
)
|
| 18 |
+
rewards = RewardConfig(
|
| 19 |
+
delivery_reward=50,
|
| 20 |
+
destination_reward=10,
|
| 21 |
+
step_penalty=-1,
|
| 22 |
+
invalid_action_penalty=-20,
|
| 23 |
+
delay_penalty=-5,
|
| 24 |
+
)
|
| 25 |
+
return LastMileDeliveryEnvironment(config=config, reward_config=rewards)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def _observation_snapshot(observation) -> tuple:
|
| 29 |
+
orders = tuple(
|
| 30 |
+
sorted(
|
| 31 |
+
(
|
| 32 |
+
order.order_id,
|
| 33 |
+
order.pickup.x,
|
| 34 |
+
order.pickup.y,
|
| 35 |
+
order.dropoff.x,
|
| 36 |
+
order.dropoff.y,
|
| 37 |
+
)
|
| 38 |
+
for order in observation.pending_orders
|
| 39 |
+
)
|
| 40 |
+
)
|
| 41 |
+
obstacles = tuple(sorted((item.x, item.y) for item in observation.obstacles))
|
| 42 |
+
traffic = tuple(sorted((zone.location.x, zone.location.y, zone.extra_cost) for zone in observation.traffic_zones))
|
| 43 |
+
return orders, obstacles, traffic
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def test_reset_and_state_contract() -> None:
|
| 47 |
+
env = build_env()
|
| 48 |
+
observation = env.reset()
|
| 49 |
+
|
| 50 |
+
assert observation.agent_location.x == 0
|
| 51 |
+
assert observation.agent_location.y == 0
|
| 52 |
+
assert isinstance(observation.pending_orders, list)
|
| 53 |
+
assert observation.current_order is None
|
| 54 |
+
|
| 55 |
+
state = env.state()
|
| 56 |
+
assert state.initialized is True
|
| 57 |
+
assert state.done is False
|
| 58 |
+
assert state.observation is not None
|
| 59 |
+
assert state.step_count == 0
|
| 60 |
+
assert state.total_reward == 0.0
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
def test_invalid_move_penalty() -> None:
|
| 64 |
+
env = build_env()
|
| 65 |
+
env.reset()
|
| 66 |
+
|
| 67 |
+
observation, reward, done, info = env.step(SimulatorAction(action_type=ActionType.MOVE, direction="left"))
|
| 68 |
+
|
| 69 |
+
assert observation.step_count == 1
|
| 70 |
+
assert done is False
|
| 71 |
+
assert info["invalid_action"] is True
|
| 72 |
+
assert reward <= -21
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
def test_observation_fields_present() -> None:
|
| 76 |
+
env = build_env()
|
| 77 |
+
observation = env.reset()
|
| 78 |
+
|
| 79 |
+
assert hasattr(observation, "agent_location")
|
| 80 |
+
assert hasattr(observation, "pending_orders")
|
| 81 |
+
assert hasattr(observation, "current_order")
|
| 82 |
+
assert hasattr(observation, "obstacles")
|
| 83 |
+
assert hasattr(observation, "battery_level")
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
def test_deterministic_reset_with_seed() -> None:
|
| 87 |
+
env = build_env()
|
| 88 |
+
first = env.reset()
|
| 89 |
+
second = env.reset()
|
| 90 |
+
|
| 91 |
+
assert _observation_snapshot(first) == _observation_snapshot(second)
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
def test_reset_accepts_seed_and_reproduces_environment_state() -> None:
|
| 95 |
+
config = EnvironmentConfig(
|
| 96 |
+
width=8,
|
| 97 |
+
height=8,
|
| 98 |
+
max_steps=80,
|
| 99 |
+
max_orders=3,
|
| 100 |
+
obstacle_density=0.2,
|
| 101 |
+
traffic_density=0.2,
|
| 102 |
+
battery_enabled=False,
|
| 103 |
+
seed=None,
|
| 104 |
+
)
|
| 105 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 106 |
+
|
| 107 |
+
seed_a_first = env.reset(seed=111)
|
| 108 |
+
seed_a_second = env.reset(seed=111)
|
| 109 |
+
seed_b = env.reset(seed=222)
|
| 110 |
+
|
| 111 |
+
assert _observation_snapshot(seed_a_first) == _observation_snapshot(seed_a_second)
|
| 112 |
+
assert _observation_snapshot(seed_a_first) != _observation_snapshot(seed_b)
|
| 113 |
+
|
| 114 |
+
|
| 115 |
+
def test_step_after_done_is_stable_noop() -> None:
|
| 116 |
+
config = EnvironmentConfig(
|
| 117 |
+
width=6,
|
| 118 |
+
height=6,
|
| 119 |
+
max_steps=1,
|
| 120 |
+
max_orders=1,
|
| 121 |
+
obstacle_density=0.0,
|
| 122 |
+
traffic_density=0.0,
|
| 123 |
+
battery_enabled=False,
|
| 124 |
+
seed=7,
|
| 125 |
+
)
|
| 126 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 127 |
+
env.reset()
|
| 128 |
+
|
| 129 |
+
_, _, done_first, _ = env.step(SimulatorAction(action_type=ActionType.WAIT))
|
| 130 |
+
assert done_first is True
|
| 131 |
+
|
| 132 |
+
_, reward_second, done_second, info_second = env.step(SimulatorAction(action_type=ActionType.WAIT))
|
| 133 |
+
assert done_second is True
|
| 134 |
+
assert reward_second == 0.0
|
| 135 |
+
assert info_second["message"] == "episode_already_done"
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
def test_reset_with_scenario_uses_external_world_state() -> None:
|
| 139 |
+
env = build_env()
|
| 140 |
+
scenario = ScenarioDefinition.model_validate(
|
| 141 |
+
{
|
| 142 |
+
"width": 7,
|
| 143 |
+
"height": 7,
|
| 144 |
+
"max_steps": 40,
|
| 145 |
+
"agent_start": {"x": 1, "y": 1},
|
| 146 |
+
"orders": [
|
| 147 |
+
{
|
| 148 |
+
"order_id": "custom_1",
|
| 149 |
+
"pickup": {"x": 2, "y": 1},
|
| 150 |
+
"dropoff": {"x": 4, "y": 1},
|
| 151 |
+
"delivery_locations": [{"x": 4, "y": 1}],
|
| 152 |
+
"priority": "high",
|
| 153 |
+
}
|
| 154 |
+
],
|
| 155 |
+
"obstacles": [{"x": 5, "y": 5}],
|
| 156 |
+
"dynamic_obstacles": [{"x": 5, "y": 4}],
|
| 157 |
+
"traffic_zones": [{"location": {"x": 3, "y": 1}, "extra_cost": 3}],
|
| 158 |
+
"charging_stations": [{"x": 1, "y": 1}],
|
| 159 |
+
"battery_profile": {
|
| 160 |
+
"enabled": True,
|
| 161 |
+
"capacity": 10,
|
| 162 |
+
"recharge_rate": 4,
|
| 163 |
+
"initial_level": 6,
|
| 164 |
+
},
|
| 165 |
+
}
|
| 166 |
+
)
|
| 167 |
+
|
| 168 |
+
observation = env.reset_with_scenario(scenario=scenario, seed=13)
|
| 169 |
+
|
| 170 |
+
assert observation.grid_width == 7
|
| 171 |
+
assert observation.grid_height == 7
|
| 172 |
+
assert observation.agent_location.x == 1
|
| 173 |
+
assert observation.agent_location.y == 1
|
| 174 |
+
assert observation.battery_level == 6
|
| 175 |
+
assert sorted((item.x, item.y) for item in observation.obstacles) == [(5, 4), (5, 5)]
|
| 176 |
+
assert len(observation.pending_orders) == 1
|
| 177 |
+
assert observation.pending_orders[0].order_id == "custom_1"
|
tests/test_grader.py
ADDED
|
@@ -0,0 +1,270 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from env.models import Observation, Order, OrderPriority, Position, StepInfo, StepResult
|
| 4 |
+
from grader.grader import DeliveryEpisodeGrader
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
def _order(
|
| 8 |
+
order_id: str,
|
| 9 |
+
pickup: tuple[int, int],
|
| 10 |
+
dropoff: tuple[int, int],
|
| 11 |
+
priority: OrderPriority = OrderPriority.LOW,
|
| 12 |
+
accepted_step: int | None = None,
|
| 13 |
+
) -> Order:
|
| 14 |
+
return Order(
|
| 15 |
+
order_id=order_id,
|
| 16 |
+
pickup=Position(x=pickup[0], y=pickup[1]),
|
| 17 |
+
dropoff=Position(x=dropoff[0], y=dropoff[1]),
|
| 18 |
+
delivery_locations=[Position(x=dropoff[0], y=dropoff[1])],
|
| 19 |
+
priority=priority,
|
| 20 |
+
accepted_step=accepted_step,
|
| 21 |
+
)
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def _obs(step: int, pending: list[Order]) -> Observation:
|
| 25 |
+
return Observation(
|
| 26 |
+
grid_width=6,
|
| 27 |
+
grid_height=6,
|
| 28 |
+
agent_location=Position(x=0, y=0),
|
| 29 |
+
pending_orders=pending,
|
| 30 |
+
current_order=None,
|
| 31 |
+
obstacles=[],
|
| 32 |
+
dynamic_obstacles=[],
|
| 33 |
+
traffic_zones=[],
|
| 34 |
+
charging_stations=[],
|
| 35 |
+
battery_level=100,
|
| 36 |
+
step_count=step,
|
| 37 |
+
total_reward=0.0,
|
| 38 |
+
)
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def test_grader_is_deterministic() -> None:
|
| 42 |
+
grader = DeliveryEpisodeGrader()
|
| 43 |
+
order_1 = _order("order_1", (1, 0), (2, 0), OrderPriority.HIGH, accepted_step=1)
|
| 44 |
+
order_2 = _order("order_2", (2, 1), (3, 1), OrderPriority.LOW)
|
| 45 |
+
|
| 46 |
+
trajectory: list[StepResult] = []
|
| 47 |
+
for step in range(1, 11):
|
| 48 |
+
delivered = None
|
| 49 |
+
if step == 5:
|
| 50 |
+
delivered = "order_1"
|
| 51 |
+
if step == 10:
|
| 52 |
+
delivered = "order_2"
|
| 53 |
+
|
| 54 |
+
trajectory.append(
|
| 55 |
+
StepResult(
|
| 56 |
+
observation=_obs(step, [order_1, order_2]),
|
| 57 |
+
reward=0.0,
|
| 58 |
+
done=False,
|
| 59 |
+
info=StepInfo(delivered_order_id=delivered, made_progress=True),
|
| 60 |
+
)
|
| 61 |
+
)
|
| 62 |
+
|
| 63 |
+
final_obs = _obs(10, [])
|
| 64 |
+
|
| 65 |
+
success_condition = {
|
| 66 |
+
"completion_rate_min": 1.0,
|
| 67 |
+
"high_priority_on_time_rate_min": 0.8,
|
| 68 |
+
"max_steps": 20,
|
| 69 |
+
"invalid_action_rate_max": 0.10,
|
| 70 |
+
}
|
| 71 |
+
report_one = grader.grade_episode(trajectory, final_obs, success_condition=success_condition)
|
| 72 |
+
report_two = grader.grade_episode(trajectory, final_obs, success_condition=success_condition)
|
| 73 |
+
|
| 74 |
+
assert report_one.score == report_two.score
|
| 75 |
+
assert 0.0 <= report_one.score <= 1.0
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def test_grader_completion_rate_matches_delivered_over_total() -> None:
|
| 79 |
+
grader = DeliveryEpisodeGrader()
|
| 80 |
+
order_1 = _order("order_1", (1, 0), (2, 0), accepted_step=1)
|
| 81 |
+
order_2 = _order("order_2", (2, 1), (3, 1), accepted_step=1)
|
| 82 |
+
trajectory = [
|
| 83 |
+
StepResult(
|
| 84 |
+
observation=_obs(step, [order_1, order_2]),
|
| 85 |
+
reward=0.0,
|
| 86 |
+
done=False,
|
| 87 |
+
info=StepInfo(delivered_order_id="order_1" if step == 3 else None, made_progress=True),
|
| 88 |
+
)
|
| 89 |
+
for step in range(1, 4)
|
| 90 |
+
]
|
| 91 |
+
report = grader.grade_episode(
|
| 92 |
+
trajectory,
|
| 93 |
+
_obs(3, [order_2]),
|
| 94 |
+
success_condition={"completion_rate_min": 1.0, "max_steps": 10, "invalid_action_rate_max": 0.1},
|
| 95 |
+
)
|
| 96 |
+
|
| 97 |
+
assert report.delivered_orders == 1
|
| 98 |
+
assert report.total_orders == 2
|
| 99 |
+
assert report.completion_rate == 0.5
|
| 100 |
+
assert report.completion_component == 0.5
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
def test_grader_high_priority_on_time_rate_impacts_score() -> None:
|
| 104 |
+
grader = DeliveryEpisodeGrader()
|
| 105 |
+
high = _order("order_high", (1, 0), (2, 0), OrderPriority.HIGH, accepted_step=1)
|
| 106 |
+
|
| 107 |
+
on_time_trajectory = [
|
| 108 |
+
StepResult(
|
| 109 |
+
observation=_obs(step, [high]),
|
| 110 |
+
reward=0.0,
|
| 111 |
+
done=False,
|
| 112 |
+
info=StepInfo(delivered_order_id="order_high" if step == 5 else None, made_progress=True),
|
| 113 |
+
)
|
| 114 |
+
for step in range(1, 6)
|
| 115 |
+
]
|
| 116 |
+
|
| 117 |
+
late_trajectory = [
|
| 118 |
+
StepResult(
|
| 119 |
+
observation=_obs(step, [high]),
|
| 120 |
+
reward=0.0,
|
| 121 |
+
done=False,
|
| 122 |
+
info=StepInfo(delivered_order_id="order_high" if step == 20 else None, made_progress=True),
|
| 123 |
+
)
|
| 124 |
+
for step in range(1, 21)
|
| 125 |
+
]
|
| 126 |
+
|
| 127 |
+
success_condition = {
|
| 128 |
+
"completion_rate_min": 1.0,
|
| 129 |
+
"high_priority_on_time_rate_min": 0.8,
|
| 130 |
+
"max_steps": 30,
|
| 131 |
+
"invalid_action_rate_max": 0.1,
|
| 132 |
+
}
|
| 133 |
+
|
| 134 |
+
report_on_time = grader.grade_episode(on_time_trajectory, _obs(5, []), success_condition=success_condition)
|
| 135 |
+
report_late = grader.grade_episode(late_trajectory, _obs(20, []), success_condition=success_condition)
|
| 136 |
+
|
| 137 |
+
assert report_on_time.high_priority_on_time_rate == 1.0
|
| 138 |
+
assert report_late.high_priority_on_time_rate == 0.0
|
| 139 |
+
assert report_on_time.priority_component > report_late.priority_component
|
| 140 |
+
assert report_on_time.score > report_late.score
|
| 141 |
+
|
| 142 |
+
|
| 143 |
+
def test_grader_efficiency_uses_steps_vs_max_steps() -> None:
|
| 144 |
+
grader = DeliveryEpisodeGrader()
|
| 145 |
+
order = _order("order_1", (1, 0), (2, 0), accepted_step=1)
|
| 146 |
+
|
| 147 |
+
short_trajectory = [
|
| 148 |
+
StepResult(
|
| 149 |
+
observation=_obs(step, [order]),
|
| 150 |
+
reward=0.0,
|
| 151 |
+
done=False,
|
| 152 |
+
info=StepInfo(delivered_order_id="order_1" if step == 3 else None, made_progress=True),
|
| 153 |
+
)
|
| 154 |
+
for step in range(1, 4)
|
| 155 |
+
]
|
| 156 |
+
|
| 157 |
+
long_trajectory = [
|
| 158 |
+
StepResult(
|
| 159 |
+
observation=_obs(step, [order]),
|
| 160 |
+
reward=0.0,
|
| 161 |
+
done=False,
|
| 162 |
+
info=StepInfo(delivered_order_id="order_1" if step == 9 else None, made_progress=True),
|
| 163 |
+
)
|
| 164 |
+
for step in range(1, 10)
|
| 165 |
+
]
|
| 166 |
+
|
| 167 |
+
success_condition = {"completion_rate_min": 1.0, "max_steps": 10, "invalid_action_rate_max": 0.1}
|
| 168 |
+
report_short = grader.grade_episode(short_trajectory, _obs(3, []), success_condition=success_condition)
|
| 169 |
+
report_long = grader.grade_episode(long_trajectory, _obs(9, []), success_condition=success_condition)
|
| 170 |
+
|
| 171 |
+
assert report_short.efficiency_ratio > report_long.efficiency_ratio
|
| 172 |
+
assert report_short.efficiency_component > report_long.efficiency_component
|
| 173 |
+
assert report_short.score > report_long.score
|
| 174 |
+
|
| 175 |
+
|
| 176 |
+
def test_grader_invalid_action_rate_reduces_penalty_component() -> None:
|
| 177 |
+
grader = DeliveryEpisodeGrader()
|
| 178 |
+
order = _order("order_1", (1, 0), (2, 0), accepted_step=1)
|
| 179 |
+
|
| 180 |
+
clean_trajectory = [
|
| 181 |
+
StepResult(
|
| 182 |
+
observation=_obs(step, [order]),
|
| 183 |
+
reward=0.0,
|
| 184 |
+
done=False,
|
| 185 |
+
info=StepInfo(delivered_order_id="order_1" if step == 5 else None, made_progress=True),
|
| 186 |
+
)
|
| 187 |
+
for step in range(1, 6)
|
| 188 |
+
]
|
| 189 |
+
|
| 190 |
+
noisy_trajectory = [
|
| 191 |
+
StepResult(
|
| 192 |
+
observation=_obs(step, [order]),
|
| 193 |
+
reward=0.0,
|
| 194 |
+
done=False,
|
| 195 |
+
info=StepInfo(
|
| 196 |
+
delivered_order_id="order_1" if step == 5 else None,
|
| 197 |
+
invalid_action=step in {1, 3, 4},
|
| 198 |
+
made_progress=step == 5,
|
| 199 |
+
),
|
| 200 |
+
)
|
| 201 |
+
for step in range(1, 6)
|
| 202 |
+
]
|
| 203 |
+
|
| 204 |
+
success_condition = {"completion_rate_min": 1.0, "max_steps": 10, "invalid_action_rate_max": 0.1}
|
| 205 |
+
clean_report = grader.grade_episode(clean_trajectory, _obs(5, []), success_condition=success_condition)
|
| 206 |
+
noisy_report = grader.grade_episode(noisy_trajectory, _obs(5, []), success_condition=success_condition)
|
| 207 |
+
|
| 208 |
+
assert clean_report.invalid_action_rate < noisy_report.invalid_action_rate
|
| 209 |
+
assert clean_report.penalty_component > noisy_report.penalty_component
|
| 210 |
+
assert clean_report.score > noisy_report.score
|
| 211 |
+
|
| 212 |
+
|
| 213 |
+
def test_grader_no_high_priority_orders_is_handled_cleanly() -> None:
|
| 214 |
+
grader = DeliveryEpisodeGrader()
|
| 215 |
+
low = _order("order_low", (1, 0), (2, 0), OrderPriority.LOW, accepted_step=1)
|
| 216 |
+
trajectory = [
|
| 217 |
+
StepResult(
|
| 218 |
+
observation=_obs(step, [low]),
|
| 219 |
+
reward=0.0,
|
| 220 |
+
done=False,
|
| 221 |
+
info=StepInfo(delivered_order_id="order_low" if step == 3 else None, made_progress=True),
|
| 222 |
+
)
|
| 223 |
+
for step in range(1, 4)
|
| 224 |
+
]
|
| 225 |
+
|
| 226 |
+
report = grader.grade_episode(
|
| 227 |
+
trajectory,
|
| 228 |
+
_obs(3, []),
|
| 229 |
+
success_condition={
|
| 230 |
+
"completion_rate_min": 1.0,
|
| 231 |
+
"high_priority_on_time_rate_min": 0.8,
|
| 232 |
+
"max_steps": 20,
|
| 233 |
+
"invalid_action_rate_max": 0.1,
|
| 234 |
+
},
|
| 235 |
+
)
|
| 236 |
+
|
| 237 |
+
assert report.high_priority_total_orders == 0
|
| 238 |
+
assert report.high_priority_on_time_rate == 1.0
|
| 239 |
+
assert any("No high-priority orders were present" in note for note in report.edge_case_handling)
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
def test_grader_is_not_constant_for_different_step_counts() -> None:
|
| 243 |
+
grader = DeliveryEpisodeGrader()
|
| 244 |
+
order = _order("order_1", (1, 0), (2, 0), OrderPriority.LOW)
|
| 245 |
+
|
| 246 |
+
short_trajectory = [
|
| 247 |
+
StepResult(
|
| 248 |
+
observation=_obs(1, [order]),
|
| 249 |
+
reward=0.0,
|
| 250 |
+
done=False,
|
| 251 |
+
info=StepInfo(made_progress=False),
|
| 252 |
+
)
|
| 253 |
+
]
|
| 254 |
+
|
| 255 |
+
long_trajectory = [
|
| 256 |
+
StepResult(
|
| 257 |
+
observation=_obs(step, [order]),
|
| 258 |
+
reward=0.0,
|
| 259 |
+
done=False,
|
| 260 |
+
info=StepInfo(made_progress=False),
|
| 261 |
+
)
|
| 262 |
+
for step in range(1, 4)
|
| 263 |
+
]
|
| 264 |
+
|
| 265 |
+
success_condition = {"completion_rate_min": 1.0, "max_steps": 10, "invalid_action_rate_max": 0.1}
|
| 266 |
+
short_report = grader.grade_episode(short_trajectory, _obs(1, [order]), success_condition=success_condition)
|
| 267 |
+
long_report = grader.grade_episode(long_trajectory, _obs(3, [order]), success_condition=success_condition)
|
| 268 |
+
|
| 269 |
+
assert short_report.score != long_report.score
|
| 270 |
+
assert long_report.score < short_report.score
|
tests/test_inference_format.py
ADDED
|
@@ -0,0 +1,117 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import re
|
| 4 |
+
|
| 5 |
+
from env.environment import LastMileDeliveryEnvironment
|
| 6 |
+
from env.models import ActionType, Direction, SimulatorAction
|
| 7 |
+
from inference import format_end_line, format_start_line, format_step_line, inference, run_baseline_suite, run_episode
|
| 8 |
+
from tasks.registry import get_task_config
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def test_start_line_exact_format() -> None:
|
| 12 |
+
line = format_start_line(task="easy", env_name="LastMileDeliveryEnvironment", model_name="test-model")
|
| 13 |
+
assert line == "[START] task=easy env=LastMileDeliveryEnvironment model=test-model"
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def test_step_line_exact_format() -> None:
|
| 17 |
+
action = SimulatorAction(action_type=ActionType.MOVE, direction=Direction.UP)
|
| 18 |
+
line = format_step_line(step_index=1, action=action, reward=0.0, done=False, error=None)
|
| 19 |
+
assert line == (
|
| 20 |
+
'[STEP] step=1 action={"accept_order":null,"deliver_order":false,"move":"up","wait":false} '
|
| 21 |
+
'reward=0.00 done=false error=null'
|
| 22 |
+
)
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def test_end_line_exact_format() -> None:
|
| 26 |
+
line = format_end_line(success=True, steps=12, rewards=[54.0, -1.0, 3.456])
|
| 27 |
+
assert line == "[END] success=true steps=12 rewards=54.00,-1.00,3.46"
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def test_run_episode_always_prints_end_on_error(monkeypatch, capsys) -> None:
|
| 31 |
+
monkeypatch.delenv("HF_TOKEN", raising=False)
|
| 32 |
+
monkeypatch.delenv("MODEL_NAME", raising=False)
|
| 33 |
+
monkeypatch.delenv("API_BASE_URL", raising=False)
|
| 34 |
+
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
| 35 |
+
monkeypatch.delenv("OPENAI_MODEL", raising=False)
|
| 36 |
+
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
| 37 |
+
monkeypatch.delenv("OPENAI_USE_MODEL", raising=False)
|
| 38 |
+
|
| 39 |
+
result = run_episode(task="easy", seed=42, max_steps=1)
|
| 40 |
+
output_lines = [line.strip() for line in capsys.readouterr().out.strip().splitlines() if line.strip()]
|
| 41 |
+
|
| 42 |
+
assert output_lines[0].startswith("[START] ")
|
| 43 |
+
assert any(line.startswith("[STEP] ") for line in output_lines)
|
| 44 |
+
assert output_lines[-1].startswith("[END] ")
|
| 45 |
+
assert " rewards=" in output_lines[-1]
|
| 46 |
+
assert result["success"] is False
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def test_run_episode_with_env_is_strict_and_has_no_extra_output(monkeypatch, capsys) -> None:
|
| 50 |
+
monkeypatch.setenv("HF_TOKEN", "test-key")
|
| 51 |
+
monkeypatch.setenv("MODEL_NAME", "test-model")
|
| 52 |
+
monkeypatch.setenv("OPENAI_API_KEY", "test-key")
|
| 53 |
+
monkeypatch.setenv("OPENAI_MODEL", "test-model")
|
| 54 |
+
monkeypatch.setenv("OPENAI_USE_MODEL", "0")
|
| 55 |
+
|
| 56 |
+
result = run_episode(task="easy", seed=1, max_steps=1)
|
| 57 |
+
output_lines = [line.strip() for line in capsys.readouterr().out.strip().splitlines() if line.strip()]
|
| 58 |
+
|
| 59 |
+
assert len(output_lines) == 3
|
| 60 |
+
assert output_lines[0] == "[START] task=easy env=LastMileDeliveryEnvironment model=test-model"
|
| 61 |
+
assert re.fullmatch(
|
| 62 |
+
r'\[STEP\] step=1 action=\{.*\} reward=-?\d+\.\d{2} done=(true|false) error=null',
|
| 63 |
+
output_lines[1],
|
| 64 |
+
)
|
| 65 |
+
assert re.fullmatch(
|
| 66 |
+
r'\[END\] success=(true|false) steps=\d+ rewards=-?\d+\.\d{2}(,-?\d+\.\d{2})*',
|
| 67 |
+
output_lines[2],
|
| 68 |
+
)
|
| 69 |
+
assert result["steps"] == 1
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def test_run_episode_is_deterministic_for_same_seed(monkeypatch, capsys) -> None:
|
| 73 |
+
monkeypatch.setenv("HF_TOKEN", "test-key")
|
| 74 |
+
monkeypatch.setenv("MODEL_NAME", "test-model")
|
| 75 |
+
monkeypatch.setenv("OPENAI_API_KEY", "test-key")
|
| 76 |
+
monkeypatch.setenv("OPENAI_MODEL", "test-model")
|
| 77 |
+
monkeypatch.setenv("OPENAI_USE_MODEL", "0")
|
| 78 |
+
|
| 79 |
+
result_one = run_episode(task="easy", seed=7, max_steps=5)
|
| 80 |
+
output_one = [line.strip() for line in capsys.readouterr().out.strip().splitlines() if line.strip()]
|
| 81 |
+
|
| 82 |
+
result_two = run_episode(task="easy", seed=7, max_steps=5)
|
| 83 |
+
output_two = [line.strip() for line in capsys.readouterr().out.strip().splitlines() if line.strip()]
|
| 84 |
+
|
| 85 |
+
assert result_one == result_two
|
| 86 |
+
assert output_one == output_two
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def test_inference_same_input_same_output(monkeypatch) -> None:
|
| 90 |
+
monkeypatch.setenv("HF_TOKEN", "test-key")
|
| 91 |
+
monkeypatch.setenv("MODEL_NAME", "test-model")
|
| 92 |
+
monkeypatch.setenv("OPENAI_API_KEY", "test-key")
|
| 93 |
+
monkeypatch.setenv("OPENAI_MODEL", "test-model")
|
| 94 |
+
monkeypatch.setenv("OPENAI_USE_MODEL", "0")
|
| 95 |
+
|
| 96 |
+
config = get_task_config(task_name="easy", seed=11)
|
| 97 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 98 |
+
observation = env.reset().model_dump()
|
| 99 |
+
|
| 100 |
+
action_one = inference(observation)
|
| 101 |
+
action_two = inference(observation)
|
| 102 |
+
|
| 103 |
+
assert action_one == action_two
|
| 104 |
+
|
| 105 |
+
|
| 106 |
+
def test_run_baseline_suite_is_deterministic_for_same_seed(monkeypatch) -> None:
|
| 107 |
+
monkeypatch.setenv("HF_TOKEN", "test-key")
|
| 108 |
+
monkeypatch.setenv("MODEL_NAME", "test-model")
|
| 109 |
+
monkeypatch.setenv("OPENAI_API_KEY", "test-key")
|
| 110 |
+
monkeypatch.setenv("OPENAI_MODEL", "test-model")
|
| 111 |
+
monkeypatch.setenv("OPENAI_USE_MODEL", "0")
|
| 112 |
+
|
| 113 |
+
suite_one = run_baseline_suite(seed=9, max_steps=20)
|
| 114 |
+
suite_two = run_baseline_suite(seed=9, max_steps=20)
|
| 115 |
+
|
| 116 |
+
assert [item["task"] for item in suite_one["tasks"]] == ["easy", "medium", "hard"]
|
| 117 |
+
assert suite_one == suite_two
|
tests/test_real_world_constraints.py
ADDED
|
@@ -0,0 +1,148 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from env.environment import LastMileDeliveryEnvironment
|
| 4 |
+
from env.models import ActionType, Direction, EnvironmentConfig, OrderPriority, SimulatorAction
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
def _move_towards(env: LastMileDeliveryEnvironment, target_x: int, target_y: int) -> tuple[float, bool]:
|
| 8 |
+
obs = env.current_observation()
|
| 9 |
+
if obs.agent_location.x < target_x:
|
| 10 |
+
_, reward, done, _ = env.step(SimulatorAction(action_type=ActionType.MOVE, direction=Direction.RIGHT))
|
| 11 |
+
return reward, done
|
| 12 |
+
if obs.agent_location.x > target_x:
|
| 13 |
+
_, reward, done, _ = env.step(SimulatorAction(action_type=ActionType.MOVE, direction=Direction.LEFT))
|
| 14 |
+
return reward, done
|
| 15 |
+
if obs.agent_location.y < target_y:
|
| 16 |
+
_, reward, done, _ = env.step(SimulatorAction(action_type=ActionType.MOVE, direction=Direction.DOWN))
|
| 17 |
+
return reward, done
|
| 18 |
+
_, reward, done, _ = env.step(SimulatorAction(action_type=ActionType.MOVE, direction=Direction.UP))
|
| 19 |
+
return reward, done
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def test_observation_includes_new_constraints() -> None:
|
| 23 |
+
config = EnvironmentConfig(
|
| 24 |
+
width=8,
|
| 25 |
+
height=8,
|
| 26 |
+
max_steps=80,
|
| 27 |
+
max_orders=2,
|
| 28 |
+
delivery_locations_per_order=2,
|
| 29 |
+
priority_high_ratio=1.0,
|
| 30 |
+
obstacle_density=0.2,
|
| 31 |
+
dynamic_obstacles_enabled=True,
|
| 32 |
+
dynamic_obstacle_ratio=0.5,
|
| 33 |
+
traffic_density=0.1,
|
| 34 |
+
battery_enabled=True,
|
| 35 |
+
battery_capacity=20,
|
| 36 |
+
battery_recharge_rate=5,
|
| 37 |
+
seed=9,
|
| 38 |
+
)
|
| 39 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 40 |
+
observation = env.reset()
|
| 41 |
+
|
| 42 |
+
assert isinstance(observation.dynamic_obstacles, list)
|
| 43 |
+
assert isinstance(observation.charging_stations, list)
|
| 44 |
+
assert len(observation.pending_orders) > 0
|
| 45 |
+
for order in observation.pending_orders:
|
| 46 |
+
assert len(order.delivery_locations) >= 1
|
| 47 |
+
assert order.priority in {OrderPriority.HIGH, OrderPriority.LOW}
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def test_dynamic_obstacles_are_deterministic_on_reset() -> None:
|
| 51 |
+
config = EnvironmentConfig(
|
| 52 |
+
width=8,
|
| 53 |
+
height=8,
|
| 54 |
+
max_steps=50,
|
| 55 |
+
max_orders=1,
|
| 56 |
+
obstacle_density=0.2,
|
| 57 |
+
dynamic_obstacles_enabled=True,
|
| 58 |
+
dynamic_obstacle_ratio=1.0,
|
| 59 |
+
seed=3,
|
| 60 |
+
)
|
| 61 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 62 |
+
|
| 63 |
+
first = env.reset()
|
| 64 |
+
second = env.reset()
|
| 65 |
+
|
| 66 |
+
first_dynamic = sorted((p.x, p.y) for p in first.dynamic_obstacles)
|
| 67 |
+
second_dynamic = sorted((p.x, p.y) for p in second.dynamic_obstacles)
|
| 68 |
+
assert first_dynamic == second_dynamic
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def test_battery_recharge_and_failure_flow() -> None:
|
| 72 |
+
config = EnvironmentConfig(
|
| 73 |
+
width=6,
|
| 74 |
+
height=6,
|
| 75 |
+
max_steps=30,
|
| 76 |
+
max_orders=1,
|
| 77 |
+
obstacle_density=0.0,
|
| 78 |
+
traffic_density=0.0,
|
| 79 |
+
battery_enabled=True,
|
| 80 |
+
battery_capacity=3,
|
| 81 |
+
battery_recharge_rate=3,
|
| 82 |
+
seed=5,
|
| 83 |
+
)
|
| 84 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 85 |
+
env.reset()
|
| 86 |
+
|
| 87 |
+
_, _, done1, _ = env.step(SimulatorAction(action_type=ActionType.MOVE, direction=Direction.RIGHT))
|
| 88 |
+
assert done1 is False
|
| 89 |
+
|
| 90 |
+
_, _, done2, _ = env.step(SimulatorAction(action_type=ActionType.MOVE, direction=Direction.LEFT))
|
| 91 |
+
assert done2 is False
|
| 92 |
+
|
| 93 |
+
observation, reward_wait, done3, info_wait = env.step(SimulatorAction(action_type=ActionType.WAIT))
|
| 94 |
+
assert done3 is False
|
| 95 |
+
assert info_wait["message"] == "recharged"
|
| 96 |
+
assert observation.battery_level == 3
|
| 97 |
+
assert reward_wait > 0.0
|
| 98 |
+
|
| 99 |
+
# Drain battery completely to trigger failure.
|
| 100 |
+
env.step(SimulatorAction(action_type=ActionType.MOVE, direction=Direction.RIGHT))
|
| 101 |
+
env.step(SimulatorAction(action_type=ActionType.MOVE, direction=Direction.RIGHT))
|
| 102 |
+
_, reward_fail, done_fail, info_fail = env.step(SimulatorAction(action_type=ActionType.MOVE, direction=Direction.RIGHT))
|
| 103 |
+
|
| 104 |
+
assert done_fail is True
|
| 105 |
+
assert info_fail["battery_depleted"] is True
|
| 106 |
+
assert reward_fail < 0.0
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
def test_high_priority_delivery_gets_bonus_reward() -> None:
|
| 110 |
+
config = EnvironmentConfig(
|
| 111 |
+
width=6,
|
| 112 |
+
height=6,
|
| 113 |
+
max_steps=120,
|
| 114 |
+
max_orders=1,
|
| 115 |
+
delivery_locations_per_order=2,
|
| 116 |
+
priority_high_ratio=1.0,
|
| 117 |
+
obstacle_density=0.0,
|
| 118 |
+
traffic_density=0.0,
|
| 119 |
+
battery_enabled=False,
|
| 120 |
+
seed=11,
|
| 121 |
+
)
|
| 122 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 123 |
+
observation = env.reset()
|
| 124 |
+
|
| 125 |
+
order = observation.pending_orders[0]
|
| 126 |
+
env.step(SimulatorAction(action_type=ActionType.ACCEPT_ORDER, order_id=order.order_id))
|
| 127 |
+
|
| 128 |
+
pickup = order.pickup
|
| 129 |
+
while True:
|
| 130 |
+
current = env.current_observation().agent_location
|
| 131 |
+
if current.x == pickup.x and current.y == pickup.y:
|
| 132 |
+
break
|
| 133 |
+
_, done = _move_towards(env, pickup.x, pickup.y)
|
| 134 |
+
assert done is False
|
| 135 |
+
|
| 136 |
+
# Move to the primary dropoff target.
|
| 137 |
+
drop = order.delivery_locations[0]
|
| 138 |
+
while True:
|
| 139 |
+
current = env.current_observation().agent_location
|
| 140 |
+
if current.x == drop.x and current.y == drop.y:
|
| 141 |
+
break
|
| 142 |
+
_, done = _move_towards(env, drop.x, drop.y)
|
| 143 |
+
assert done is False
|
| 144 |
+
|
| 145 |
+
_, reward, _, info = env.step(SimulatorAction(action_type=ActionType.DELIVER_ORDER))
|
| 146 |
+
|
| 147 |
+
assert info["delivered_order_id"] == order.order_id
|
| 148 |
+
assert reward >= 60.0
|
tests/test_reward_shaping.py
ADDED
|
@@ -0,0 +1,146 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from env.environment import LastMileDeliveryEnvironment
|
| 4 |
+
from env.models import ActionType, Direction, EnvironmentConfig, SimulatorAction
|
| 5 |
+
from env.utils import manhattan
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
def _direction_towards(current_x: int, current_y: int, target_x: int, target_y: int) -> Direction:
|
| 9 |
+
if target_x > current_x:
|
| 10 |
+
return Direction.RIGHT
|
| 11 |
+
if target_x < current_x:
|
| 12 |
+
return Direction.LEFT
|
| 13 |
+
if target_y > current_y:
|
| 14 |
+
return Direction.DOWN
|
| 15 |
+
return Direction.UP
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def _opposite(direction: Direction) -> Direction:
|
| 19 |
+
if direction == Direction.UP:
|
| 20 |
+
return Direction.DOWN
|
| 21 |
+
if direction == Direction.DOWN:
|
| 22 |
+
return Direction.UP
|
| 23 |
+
if direction == Direction.LEFT:
|
| 24 |
+
return Direction.RIGHT
|
| 25 |
+
return Direction.LEFT
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def _move_to(env: LastMileDeliveryEnvironment, target_x: int, target_y: int) -> None:
|
| 29 |
+
while True:
|
| 30 |
+
current = env.current_observation().agent_location
|
| 31 |
+
if current.x == target_x and current.y == target_y:
|
| 32 |
+
return
|
| 33 |
+
direction = _direction_towards(current.x, current.y, target_x, target_y)
|
| 34 |
+
env.step(SimulatorAction(action_type=ActionType.MOVE, direction=direction))
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def _run_delivery(seed: int, extra_waits: int) -> float:
|
| 38 |
+
config = EnvironmentConfig(
|
| 39 |
+
width=8,
|
| 40 |
+
height=8,
|
| 41 |
+
max_steps=200,
|
| 42 |
+
max_orders=1,
|
| 43 |
+
delivery_locations_per_order=1,
|
| 44 |
+
priority_high_ratio=0.0,
|
| 45 |
+
obstacle_density=0.0,
|
| 46 |
+
traffic_density=0.0,
|
| 47 |
+
battery_enabled=False,
|
| 48 |
+
seed=seed,
|
| 49 |
+
)
|
| 50 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 51 |
+
observation = env.reset()
|
| 52 |
+
order = observation.pending_orders[0]
|
| 53 |
+
|
| 54 |
+
env.step(SimulatorAction(action_type=ActionType.ACCEPT_ORDER, order_id=order.order_id))
|
| 55 |
+
_move_to(env, order.pickup.x, order.pickup.y)
|
| 56 |
+
env.step(SimulatorAction(action_type=ActionType.DELIVER_ORDER)) # pickup transition
|
| 57 |
+
|
| 58 |
+
for _ in range(extra_waits):
|
| 59 |
+
env.step(SimulatorAction(action_type=ActionType.WAIT))
|
| 60 |
+
|
| 61 |
+
drop = order.delivery_locations[0]
|
| 62 |
+
_move_to(env, drop.x, drop.y)
|
| 63 |
+
_, reward, _, info = env.step(SimulatorAction(action_type=ActionType.DELIVER_ORDER))
|
| 64 |
+
assert info["delivered_order_id"] == order.order_id
|
| 65 |
+
return reward
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def test_progress_reward_is_dense_and_directional() -> None:
|
| 69 |
+
config = EnvironmentConfig(
|
| 70 |
+
width=8,
|
| 71 |
+
height=8,
|
| 72 |
+
max_steps=40,
|
| 73 |
+
max_orders=2,
|
| 74 |
+
obstacle_density=0.0,
|
| 75 |
+
traffic_density=0.0,
|
| 76 |
+
battery_enabled=False,
|
| 77 |
+
seed=13,
|
| 78 |
+
)
|
| 79 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 80 |
+
observation = env.reset()
|
| 81 |
+
|
| 82 |
+
start = observation.agent_location
|
| 83 |
+
targets = [(o.pickup.x, o.pickup.y) for o in observation.pending_orders]
|
| 84 |
+
nearest_target = min(targets, key=lambda t: manhattan((start.x, start.y), t))
|
| 85 |
+
|
| 86 |
+
direction_to_target = _direction_towards(start.x, start.y, nearest_target[0], nearest_target[1])
|
| 87 |
+
_, reward_towards, _, _ = env.step(SimulatorAction(action_type=ActionType.MOVE, direction=direction_to_target))
|
| 88 |
+
|
| 89 |
+
direction_away = _opposite(direction_to_target)
|
| 90 |
+
_, reward_away, _, _ = env.step(SimulatorAction(action_type=ActionType.MOVE, direction=direction_away))
|
| 91 |
+
|
| 92 |
+
assert reward_towards > reward_away
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def test_delay_penalty_accumulates_for_late_delivery() -> None:
|
| 96 |
+
config = EnvironmentConfig(
|
| 97 |
+
width=8,
|
| 98 |
+
height=8,
|
| 99 |
+
max_steps=80,
|
| 100 |
+
max_orders=1,
|
| 101 |
+
priority_high_ratio=1.0,
|
| 102 |
+
obstacle_density=0.0,
|
| 103 |
+
traffic_density=0.0,
|
| 104 |
+
battery_enabled=False,
|
| 105 |
+
seed=17,
|
| 106 |
+
)
|
| 107 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 108 |
+
observation = env.reset()
|
| 109 |
+
|
| 110 |
+
order = observation.pending_orders[0]
|
| 111 |
+
env.step(SimulatorAction(action_type=ActionType.ACCEPT_ORDER, order_id=order.order_id))
|
| 112 |
+
|
| 113 |
+
rewards = []
|
| 114 |
+
for _ in range(13):
|
| 115 |
+
_, reward, _, info = env.step(SimulatorAction(action_type=ActionType.WAIT))
|
| 116 |
+
rewards.append((reward, info["delay_penalty_applied"]))
|
| 117 |
+
|
| 118 |
+
assert rewards[0][1] is False
|
| 119 |
+
assert rewards[-1][1] is True
|
| 120 |
+
assert rewards[-1][0] < rewards[0][0]
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
def test_invalid_action_has_strong_penalty() -> None:
|
| 124 |
+
config = EnvironmentConfig(
|
| 125 |
+
width=6,
|
| 126 |
+
height=6,
|
| 127 |
+
max_steps=20,
|
| 128 |
+
max_orders=1,
|
| 129 |
+
obstacle_density=0.0,
|
| 130 |
+
traffic_density=0.0,
|
| 131 |
+
battery_enabled=False,
|
| 132 |
+
seed=21,
|
| 133 |
+
)
|
| 134 |
+
env = LastMileDeliveryEnvironment(config=config)
|
| 135 |
+
env.reset()
|
| 136 |
+
|
| 137 |
+
_, reward, _, info = env.step({"move": "up"})
|
| 138 |
+
assert info["invalid_action"] is True
|
| 139 |
+
assert reward <= -20.0
|
| 140 |
+
|
| 141 |
+
|
| 142 |
+
def test_efficient_delivery_gets_higher_bonus() -> None:
|
| 143 |
+
fast_reward = _run_delivery(seed=25, extra_waits=0)
|
| 144 |
+
slow_reward = _run_delivery(seed=25, extra_waits=6)
|
| 145 |
+
|
| 146 |
+
assert fast_reward > slow_reward
|
tests/test_tasks.py
ADDED
|
@@ -0,0 +1,92 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from tasks.registry import get_task_config, get_task_definition, get_task_registry, list_tasks
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
def test_registry_contains_three_tasks() -> None:
|
| 7 |
+
assert list_tasks() == ["easy", "medium", "hard"]
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def test_easy_definition_and_config() -> None:
|
| 11 |
+
definition = get_task_definition("easy")
|
| 12 |
+
config = get_task_config("easy", seed=1)
|
| 13 |
+
default_config = get_task_config("easy")
|
| 14 |
+
default_config_repeat = get_task_config("easy")
|
| 15 |
+
|
| 16 |
+
assert definition["parameters"]["orders"] == {"min": 1, "max": 1}
|
| 17 |
+
assert definition["parameters"]["constraints"]["obstacles"] is False
|
| 18 |
+
assert definition["parameters"]["constraints"]["traffic_zones"] is False
|
| 19 |
+
assert definition["parameters"]["constraints"]["order_priority"] is False
|
| 20 |
+
assert "objective" in definition
|
| 21 |
+
assert "success_condition" in definition
|
| 22 |
+
assert definition["deterministic"] is True
|
| 23 |
+
assert config.max_orders == 1
|
| 24 |
+
assert config.obstacle_density == 0.0
|
| 25 |
+
assert config.traffic_density == 0.0
|
| 26 |
+
assert config.max_steps == 60
|
| 27 |
+
assert default_config.model_dump() == default_config_repeat.model_dump()
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def test_medium_has_obstacles_and_multiple_orders() -> None:
|
| 31 |
+
config = get_task_config("medium", seed=1)
|
| 32 |
+
definition = get_task_definition("medium")
|
| 33 |
+
|
| 34 |
+
assert 3 <= config.max_orders <= 4
|
| 35 |
+
assert config.obstacle_density > 0.0
|
| 36 |
+
assert config.battery_enabled is False
|
| 37 |
+
assert config.priority_high_ratio == 0.0
|
| 38 |
+
assert definition["parameters"]["constraints"]["obstacles"] is True
|
| 39 |
+
assert definition["parameters"]["constraints"]["multiple_delivery_locations"] is False
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def test_hard_has_battery_and_time_penalty() -> None:
|
| 43 |
+
definition = get_task_definition("hard")
|
| 44 |
+
config = get_task_config("hard", seed=1)
|
| 45 |
+
|
| 46 |
+
assert definition["parameters"]["constraints"]["battery_constraint"] is True
|
| 47 |
+
assert definition["parameters"]["constraints"]["time_penalty"] is True
|
| 48 |
+
assert definition["parameters"]["constraints"]["order_priority"] is True
|
| 49 |
+
assert config.battery_enabled is True
|
| 50 |
+
assert config.dynamic_obstacles_enabled is True
|
| 51 |
+
assert config.priority_high_ratio >= 0.7
|
| 52 |
+
assert config.max_steps == 165
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def test_registry_payload_is_parameterized() -> None:
|
| 56 |
+
registry = get_task_registry()
|
| 57 |
+
assert set(registry.keys()) == {"easy", "medium", "hard"}
|
| 58 |
+
for name, item in registry.items():
|
| 59 |
+
assert item["name"] == name
|
| 60 |
+
assert "parameters" in item
|
| 61 |
+
assert "grid_size" in item["parameters"]
|
| 62 |
+
assert "objective" in item
|
| 63 |
+
assert "success_condition" in item
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def test_difficulty_progression_is_increasing() -> None:
|
| 67 |
+
easy = get_task_config("easy", seed=10)
|
| 68 |
+
medium = get_task_config("medium", seed=10)
|
| 69 |
+
hard = get_task_config("hard", seed=10)
|
| 70 |
+
|
| 71 |
+
assert easy.width < medium.width < hard.width
|
| 72 |
+
assert easy.height < medium.height < hard.height
|
| 73 |
+
assert easy.max_orders < medium.max_orders < hard.max_orders
|
| 74 |
+
assert easy.dynamic_obstacles_enabled is False
|
| 75 |
+
assert medium.dynamic_obstacles_enabled is False
|
| 76 |
+
assert hard.dynamic_obstacles_enabled is True
|
| 77 |
+
assert easy.battery_enabled is False
|
| 78 |
+
assert medium.battery_enabled is False
|
| 79 |
+
assert hard.battery_enabled is True
|
| 80 |
+
assert easy.priority_high_ratio == 0.0
|
| 81 |
+
assert medium.priority_high_ratio == 0.0
|
| 82 |
+
assert hard.priority_high_ratio > medium.priority_high_ratio
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
def test_seed_reproducibility_for_variable_tasks() -> None:
|
| 86 |
+
medium_a = get_task_config("medium", seed=77)
|
| 87 |
+
medium_b = get_task_config("medium", seed=77)
|
| 88 |
+
hard_a = get_task_config("hard", seed=91)
|
| 89 |
+
hard_b = get_task_config("hard", seed=91)
|
| 90 |
+
|
| 91 |
+
assert medium_a.model_dump() == medium_b.model_dump()
|
| 92 |
+
assert hard_a.model_dump() == hard_b.model_dump()
|
uv.lock
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|