Spaces:
Sleeping
Sleeping
File size: 1,765 Bytes
852e969 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 | from fastapi import FastAPI, HTTPException
from baseline import baseline_scores, run_baseline
from env import IndianTrafficEnv
from grader import grade_rollout
from models import GraderRequest, ResetRequest, StepRequest, StepResult, TrafficState
from tasks import list_tasks
app = FastAPI(
title="Indian Traffic Signal OpenEnv",
version="1.0.0",
description="Seedable RL environment for Indian mixed-traffic signal control.",
)
ENV = IndianTrafficEnv()
@app.post("/reset", response_model=TrafficState)
def reset(request: ResetRequest):
try:
return ENV.reset(seed=request.seed, task_id=request.task_id)
except KeyError as exc:
raise HTTPException(status_code=404, detail=str(exc)) from exc
@app.post("/step", response_model=StepResult)
def step(request: StepRequest):
observation, reward, done, info = ENV.step(request.action)
return StepResult(observation=observation, reward=reward, done=done, info=info)
@app.get("/state", response_model=TrafficState)
def state():
return ENV.get_state()
@app.get("/tasks")
def tasks():
return list_tasks()
@app.post("/grader")
def grader(request: GraderRequest):
try:
return grade_rollout(
task_id=request.task_id,
seed=request.seed,
actions=request.actions,
max_steps=request.max_steps,
)
except KeyError as exc:
raise HTTPException(status_code=404, detail=str(exc)) from exc
@app.get("/baseline")
def baseline(task_id: str = "all", seed: int = 42):
if task_id == "all":
return baseline_scores(seed=seed)
try:
return run_baseline(task_id=task_id, seed=seed)
except KeyError as exc:
raise HTTPException(status_code=404, detail=str(exc)) from exc
|