Medical-Triage / git_diff.txt
Aspirant200715's picture
graders logic redefine
89be39e
Raw
History Blame Contribute Delete
34.3 kB
��diff --git a/triage_env/server/app.py b/triage_env/server/app.py
index aa7a2f9..f018042 100644
--- a/triage_env/server/app.py
+++ b/triage_env/server/app.py
@@ -39,6 +39,14 @@ from fastapi.openapi.docs import get_redoc_html, get_swagger_ui_html
from fastapi.openapi.utils import get_openapi
from fastapi.responses import JSONResponse, RedirectResponse
+try:
+ import gradio as gr
+except Exception: # pragma: no cover
+ gr = None
+
+GRADIO_PATH = "/gradio"
+_GRADIO_MOUNT_ERROR: str | None = None
+
try:
from ..models import TriageAction, TriageObservation
from .triage_env_environment import TriageEnvironment
@@ -56,22 +64,60 @@ app = create_app(
max_concurrent_envs=1, # increase this number to allow more concurrent WebSocket sessions
)
+if gr is not None:
+ try:
+ try:
+ from .gradio_ui import build_gradio_ui
+ except ModuleNotFoundError:
+ from server.gradio_ui import build_gradio_ui
+ app = gr.mount_gradio_app(app, build_gradio_ui(), path=GRADIO_PATH)
+ except Exception as exc:
+ # Keep API serving even if Gradio UI mount fails, but retain error for diagnostics.
+ _GRADIO_MOUNT_ERROR = f"{type(exc).__name__}: {exc}"
+
def _has_route(path: str) -> bool:
return any(getattr(route, "path", None) == path for route in app.routes)
+def _has_route_variant(path: str) -> bool:
+ return _has_route(path) or _has_route(f"{path}/")
+
+
@app.get("/", include_in_schema=False)
def root():
- # On Spaces, users land on "/" first. Redirect to docs when available.
+ # On Spaces, users land on "/" first. Prefer Gradio UI, then docs.
+ if _has_route_variant(GRADIO_PATH):
+ return RedirectResponse(url=GRADIO_PATH)
if _has_route("/docs"):
return RedirectResponse(url="/docs")
- return {
+ payload = {
"message": "MedicalTriage API is running",
"health": "/health",
"openapi": "/openapi.json",
+ "ui": GRADIO_PATH,
"docs": "/docs",
}
+ if _GRADIO_MOUNT_ERROR is not None:
+ payload["gradio_mount_error"] = _GRADIO_MOUNT_ERROR
+ return payload
+
+
+if not _has_route("/ui"):
+
+ @app.get("/ui", include_in_schema=False)
+ def ui_redirect():
+ if _has_route_variant(GRADIO_PATH):
+ return RedirectResponse(url=GRADIO_PATH)
+ if _has_route("/docs"):
+ return RedirectResponse(url="/docs")
+ return JSONResponse(
+ {
+ "message": "Gradio UI is unavailable",
+ "gradio_mount_error": _GRADIO_MOUNT_ERROR,
+ },
+ status_code=503,
+ )
if not _has_route("/openapi.json"):
diff --git a/triage_env/server/gradio_ui.py b/triage_env/server/gradio_ui.py
new file mode 100644
index 0000000..d5dfd04
--- /dev/null
+++ b/triage_env/server/gradio_ui.py
@@ -0,0 +1,115 @@
+from __future__ import annotations
+
+import threading
+import traceback
+from typing import Any
+
+import gradio as gr
+
+try:
+ from ..models import TriageAction
+ from .triage_env_environment import TriageEnvironment
+except ModuleNotFoundError:
+ from models import TriageAction
+ from server.triage_env_environment import TriageEnvironment
+
+
+_ENV_LOCK = threading.Lock()
+_ENV: TriageEnvironment | None = None
+
+
+def _ensure_env(task: str | None = None) -> TriageEnvironment:
+ global _ENV
+ if _ENV is None:
+ initial_task = task or "task2"
+ _ENV = TriageEnvironment(task=initial_task)
+ _ENV.reset(task=initial_task)
+ return _ENV
+
+
+def _obs_payload(env: TriageEnvironment, observation: Any) -> dict[str, Any]:
+ return {
+ "observation": observation.model_dump(mode="json"),
+ "state": env.state.model_dump(mode="json"),
+ }
+
+
+def _state_payload(env: TriageEnvironment) -> dict[str, Any]:
+ return {
+ "task": env.task_name,
+ "step_count": env.state.step_count,
+ "state": env.state.model_dump(mode="json"),
+ }
+
+
+def build_gradio_ui() -> gr.Blocks:
+ with gr.Blocks(title="Medical Triage") as demo:
+ gr.Markdown("## Medical Triage\nSimple controls for reset, step, and getState.")
+
+ with gr.Row():
+ task = gr.Dropdown(
+ choices=["task1", "task2", "task3"],
+ value="task2",
+ label="Task",
+ )
+ action_type = gr.Dropdown(
+ choices=["treat", "allocate_ventilator", "wait"],
+ value="wait",
+ label="Action Type",
+ )
+ patient_id = gr.Number(value=-1, precision=0, label="Patient ID")
+
+ with gr.Row():
+ reset_btn = gr.Button("Reset", variant="primary")
+ step_btn = gr.Button("Step")
+ get_state_btn = gr.Button("Get State")
+
+ output = gr.JSON(label="Response")
+
+ def on_reset(selected_task: str):
+ try:
+ with _ENV_LOCK:
+ env = _ensure_env(selected_task)
+ obs = env.reset(task=selected_task)
+ return _obs_payload(env, obs)
+ except Exception as exc:
+ return {"error": str(exc), "traceback": traceback.format_exc()}
+
+ def on_step(
+ selected_task: str,
+ selected_action: str,
+ selected_patient_id: float,
+ ):
+ try:
+ with _ENV_LOCK:
+ env = _ensure_env(selected_task)
+ if env.task_name != selected_task:
+ env.reset(task=selected_task)
+
+ pid = -1 if selected_action == "wait" else int(selected_patient_id or 0)
+ action = TriageAction(action_type=selected_action, patient_id=pid)
+ obs = env.step(action)
+ return _obs_payload(env, obs)
+ except Exception as exc:
+ return {"error": str(exc), "traceback": traceback.format_exc()}
+
+ def on_get_state():
+ try:
+ with _ENV_LOCK:
+ env = _ensure_env("task2")
+ return _state_payload(env)
+ except Exception as exc:
+ return {"error": str(exc), "traceback": traceback.format_exc()}
+
+ reset_btn.click(on_reset, inputs=[task], outputs=[output], queue=False)
+ step_btn.click(
+ on_step,
+ inputs=[task, action_type, patient_id],
+ outputs=[output],
+ queue=False,
+ )
+ get_state_btn.click(on_get_state, inputs=[], outputs=[output], queue=False)
+
+ demo.load(on_get_state, inputs=[], outputs=[output], queue=False)
+
+ return demo
diff --git a/triage_env/server/triage_env_environment.py b/triage_env/server/triage_env_environment.py
index 429f208..7e7dcd6 100644
--- a/triage_env/server/triage_env_environment.py
+++ b/triage_env/server/triage_env_environment.py
@@ -1,3 +1,4 @@
+import math
from uuid import uuid4
from openenv.core.env_server.interfaces import Environment
@@ -12,6 +13,7 @@ except ImportError:
class TriageEnvironment(Environment):
SUPPORTS_CONCURRENT_SESSIONS: bool = True
+ STEP_REWARD_SCALE: float = 25.0
def __init__(
self,
@@ -253,6 +255,8 @@ class TriageEnvironment(Environment):
reward += terminal_reward
reward_breakdown["episode_terminal_reward"] = terminal_reward
+ raw_reward = reward
+ reward = self._normalize_step_reward(raw_reward)
self._state.total_reward += reward
components = {
@@ -274,7 +278,11 @@ class TriageEnvironment(Environment):
reward=reward,
reward_detail=TriageReward(value=float(reward), components=components, penalties=penalties),
message=message,
- metadata=self._build_metadata(reward_breakdown=reward_breakdown),
+ metadata=self._build_metadata(
+ reward_breakdown=reward_breakdown,
+ raw_reward=raw_reward,
+ normalized_reward=reward,
+ ),
)
@property
@@ -385,6 +393,16 @@ class TriageEnvironment(Environment):
return terminal_reward, success_achieved
+ def _normalize_step_reward(self, raw_reward: float) -> float:
+ """Normalize raw reward to (0,1) and round to two decimals."""
+ if not math.isfinite(raw_reward):
+ raw_reward = 0.0
+
+ scaled = 0.5 + (0.5 * math.tanh(raw_reward / self.STEP_REWARD_SCALE))
+ rounded_cents = int(math.floor((scaled * 100.0) + 0.5))
+ bounded_cents = min(99, max(1, rounded_cents))
+ return bounded_cents / 100.0
+
def _is_done(self) -> bool:
if self._state.step_count >= self._state.max_steps:
return True
@@ -398,12 +416,17 @@ class TriageEnvironment(Environment):
return False
- def _build_metadata(self, reward_breakdown: dict) -> dict:
+ def _build_metadata(
+ self,
+ reward_breakdown: dict,
+ raw_reward: float | None = None,
+ normalized_reward: float | None = None,
+ ) -> dict:
alive_patients = [p for p in self._state.patients if p.alive]
critical_patients = [p for p in self._state.patients if p.severity == "critical"]
surviving_critical = [p for p in critical_patients if p.alive]
- return {
+ metadata = {
"episode_id": self._state.episode_id,
"task": self.task_name,
"total_reward": self._state.total_reward,
@@ -418,4 +441,11 @@ class TriageEnvironment(Environment):
"ventilators_available": self.task_config.ventilators_available,
},
"reward_breakdown": reward_breakdown,
- }
\ No newline at end of file
+ }
+
+ if raw_reward is not None:
+ metadata["raw_reward"] = float(raw_reward)
+ if normalized_reward is not None:
+ metadata["normalized_reward"] = float(normalized_reward)
+
+ return metadata
\ No newline at end of file
diff --git a/triage_env/tests/test_environment.py b/triage_env/tests/test_environment.py
index c2a6a0e..eaba830 100644
--- a/triage_env/tests/test_environment.py
+++ b/triage_env/tests/test_environment.py
@@ -5,6 +5,11 @@ from triage_env.tasks import TASK_CONFIGS
from triage_env.models import TriageAction
+def _assert_step_reward_contract(value: float) -> None:
+ assert 0.0 < value < 1.0
+ assert value == pytest.approx(round(value, 2), abs=1e-12)
+
+
@pytest.fixture
def env():
environment = TriageEnvironment(max_steps=20)
@@ -44,6 +49,8 @@ def test_treat_more_urgent_patient_better_than_wait():
env2.reset()
wait_obs = env2.step(TriageAction(action_type="wait"))
+ _assert_step_reward_contract(treat_obs.reward)
+ _assert_step_reward_contract(wait_obs.reward)
assert treat_obs.reward > wait_obs.reward
@@ -60,13 +67,16 @@ def test_treat_critical_better_than_moderate():
TriageAction(action_type="treat", patient_id=2)
).reward
+ _assert_step_reward_contract(critical_reward)
+ _assert_step_reward_contract(moderate_reward)
assert critical_reward > moderate_reward
def test_invalid_treatment_gets_penalty(env):
obs = env.step(TriageAction(action_type="treat", patient_id=999))
- assert obs.reward < 0
+ _assert_step_reward_contract(obs.reward)
+ assert obs.reward < 0.5
assert obs.message == "Invalid treatment action"
@@ -83,13 +93,16 @@ def test_allocate_ventilator_to_critical_is_better_than_moderate():
TriageAction(action_type="allocate_ventilator", patient_id=2)
)
+ _assert_step_reward_contract(critical_obs.reward)
+ _assert_step_reward_contract(moderate_obs.reward)
assert critical_obs.reward > moderate_obs.reward
def test_wait_penalty_when_urgent_patient_exists(env):
obs = env.step(TriageAction(action_type="wait"))
- assert obs.reward < 0
+ _assert_step_reward_contract(obs.reward)
+ assert obs.reward < 0.5
assert obs.message == "Waited one step"
@@ -105,6 +118,8 @@ def test_ignoring_critical_is_bad():
wait_env.reset()
wait_reward = wait_env.step(TriageAction(action_type="wait")).reward
+ _assert_step_reward_contract(moderate_reward)
+ _assert_step_reward_contract(wait_reward)
assert moderate_reward > wait_reward
@@ -127,15 +142,30 @@ def test_patient_can_die_if_ignored():
def test_patient_death_penalty():
- env = TriageEnvironment(max_steps=10)
- env.reset()
+ wait_env = TriageEnvironment(max_steps=10)
+ wait_env.reset()
+
+ wait_total = 0.0
+ for _ in range(10):
+ obs = wait_env.step(TriageAction(action_type="wait"))
+ _assert_step_reward_contract(obs.reward)
+ wait_total += obs.reward
+
+ treat_env = TriageEnvironment(max_steps=10)
+ treat_env.reset()
- total_reward = 0.0
+ treat_total = 0.0
for _ in range(10):
- obs = env.step(TriageAction(action_type="wait"))
- total_reward += obs.reward
+ alive = [p for p in treat_env.state.patients if p.alive]
+ if not alive:
+ break
+
+ target = min(alive, key=lambda p: p.health)
+ obs = treat_env.step(TriageAction(action_type="treat", patient_id=target.id))
+ _assert_step_reward_contract(obs.reward)
+ treat_total += obs.reward
- assert total_reward < -20
+ assert wait_total < treat_total
def test_environment_runs_multiple_steps_without_crashing():
@@ -154,10 +184,26 @@ def test_environment_runs_multiple_steps_without_crashing():
assert obs is not None
assert obs.step_count >= 1
assert isinstance(obs.reward, float)
+ _assert_step_reward_contract(obs.reward)
def test_reward_breakdown_present(env):
obs = env.step(TriageAction(action_type="treat", patient_id=0))
assert "reward_breakdown" in obs.metadata
- assert isinstance(obs.metadata["reward_breakdown"], dict)
\ No newline at end of file
+ assert isinstance(obs.metadata["reward_breakdown"], dict)
+
+
+@pytest.mark.parametrize("task_name", ["task1", "task2", "task3"])
+def test_step_reward_contract_across_all_tasks(task_name):
+ env = TriageEnvironment(task=task_name)
+ env.reset(task=task_name)
+
+ actions = [
+ TriageAction(action_type="wait", patient_id=-1),
+ TriageAction(action_type="treat", patient_id=0),
+ TriageAction(action_type="allocate_ventilator", patient_id=0),
+ ]
+ for action in actions:
+ obs = env.step(action)
+ _assert_step_reward_contract(obs.reward)
\ No newline at end of file
diff --git a/triage_env/tests/test_step.py b/triage_env/tests/test_step.py
index 1433214..3cc480c 100644
--- a/triage_env/tests/test_step.py
+++ b/triage_env/tests/test_step.py
@@ -32,12 +32,18 @@ from triage_env.server.triage_env_environment import TriageEnvironment
from triage_env.models import TriageAction
+def _assert_step_reward_contract(value: float) -> None:
+ assert 0.0 < value < 1.0
+ assert value == round(value, 2)
+
+
def test_step_increments_step_count():
env = TriageEnvironment()
env.reset()
obs = env.step(TriageAction(action_type="wait", patient_id=-1))
assert obs.step_count == 1
+ _assert_step_reward_contract(obs.reward)
def test_treat_action_returns_observation():
@@ -46,4 +52,5 @@ def test_treat_action_returns_observation():
obs = env.step(TriageAction(action_type="treat", patient_id=0))
assert isinstance(obs.reward, float)
+ _assert_step_reward_contract(obs.reward)
assert hasattr(obs, "patients")
\ No newline at end of file
diff --git a/triage_env/tests/test_task_reward_scaling.py b/triage_env/tests/test_task_reward_scaling.py
index 4bc5a65..823aaf7 100644
--- a/triage_env/tests/test_task_reward_scaling.py
+++ b/triage_env/tests/test_task_reward_scaling.py
@@ -2,6 +2,11 @@ from triage_env.models import TriageAction
from triage_env.server.triage_env_environment import TriageEnvironment
+def _assert_step_reward_contract(value: float) -> None:
+ assert 0.0 < value < 1.0
+ assert value == round(value, 2)
+
+
def test_wait_penalty_harsher_on_task3_than_task1():
env_easy = TriageEnvironment(task="task1")
env_easy.reset(task="task1")
@@ -11,4 +16,6 @@ def test_wait_penalty_harsher_on_task3_than_task1():
env_hard.reset(task="task3")
reward_hard = env_hard.step(TriageAction(action_type="wait", patient_id=-1)).reward
+ _assert_step_reward_contract(reward_easy)
+ _assert_step_reward_contract(reward_hard)
assert reward_hard < reward_easy