Spaces:
Sleeping
Sleeping
Upload folder using huggingface_hub
Browse files- client.py +7 -10
- server/environment.py +5 -2
client.py
CHANGED
|
@@ -51,15 +51,12 @@ class NL2SQLEnv:
|
|
| 51 |
return self._parse_result(resp.json())
|
| 52 |
|
| 53 |
def _parse_result(self, payload: Dict[str, Any]) -> StepResult:
|
| 54 |
-
obs_data = payload.get("observation",
|
| 55 |
|
| 56 |
-
#
|
| 57 |
-
|
| 58 |
-
|
| 59 |
|
| 60 |
-
raw_score = obs_data.get("score")
|
| 61 |
-
safe_score = float(raw_score) if raw_score is not None else 0.0
|
| 62 |
-
|
| 63 |
obs = NL2SQLObservation(
|
| 64 |
question=obs_data.get("question", ""),
|
| 65 |
schema_context=obs_data.get("schema_context", ""),
|
|
@@ -70,12 +67,12 @@ class NL2SQLEnv:
|
|
| 70 |
result_columns=obs_data.get("result_columns", []),
|
| 71 |
step=obs_data.get("step", 0),
|
| 72 |
max_steps=obs_data.get("max_steps", 5),
|
| 73 |
-
done=
|
| 74 |
reward=safe_reward,
|
| 75 |
-
score=
|
| 76 |
)
|
| 77 |
return StepResult(
|
| 78 |
observation=obs,
|
| 79 |
reward=safe_reward,
|
| 80 |
-
done=
|
| 81 |
)
|
|
|
|
| 51 |
return self._parse_result(resp.json())
|
| 52 |
|
| 53 |
def _parse_result(self, payload: Dict[str, Any]) -> StepResult:
|
| 54 |
+
obs_data = payload.get("observation", {})
|
| 55 |
|
| 56 |
+
# Read standard OpenEnv top-level keys safely
|
| 57 |
+
safe_reward = float(payload.get("reward", 0.0))
|
| 58 |
+
safe_done = bool(payload.get("done", False))
|
| 59 |
|
|
|
|
|
|
|
|
|
|
| 60 |
obs = NL2SQLObservation(
|
| 61 |
question=obs_data.get("question", ""),
|
| 62 |
schema_context=obs_data.get("schema_context", ""),
|
|
|
|
| 67 |
result_columns=obs_data.get("result_columns", []),
|
| 68 |
step=obs_data.get("step", 0),
|
| 69 |
max_steps=obs_data.get("max_steps", 5),
|
| 70 |
+
done=safe_done,
|
| 71 |
reward=safe_reward,
|
| 72 |
+
score=float(obs_data.get("score", 0.0)),
|
| 73 |
)
|
| 74 |
return StepResult(
|
| 75 |
observation=obs,
|
| 76 |
reward=safe_reward,
|
| 77 |
+
done=safe_done,
|
| 78 |
)
|
server/environment.py
CHANGED
|
@@ -152,7 +152,7 @@ class NL2SQLEnvironment(Environment):
|
|
| 152 |
self._last_obs = obs
|
| 153 |
return obs
|
| 154 |
|
| 155 |
-
def step(self, action: NL2SQLAction) -> NL2SQLObservation:
|
| 156 |
"""Execute the agent's SQL and return graded observation."""
|
| 157 |
if self._task is None or self._example is None:
|
| 158 |
# Called before reset — auto-reset
|
|
@@ -210,8 +210,11 @@ class NL2SQLEnvironment(Environment):
|
|
| 210 |
reward=reward,
|
| 211 |
score=score,
|
| 212 |
)
|
|
|
|
| 213 |
self._last_obs = obs
|
| 214 |
-
|
|
|
|
|
|
|
| 215 |
|
| 216 |
@property
|
| 217 |
def state(self) -> NL2SQLState:
|
|
|
|
| 152 |
self._last_obs = obs
|
| 153 |
return obs
|
| 154 |
|
| 155 |
+
def step(self, action: NL2SQLAction) -> tuple[NL2SQLObservation, float, bool, dict]:
|
| 156 |
"""Execute the agent's SQL and return graded observation."""
|
| 157 |
if self._task is None or self._example is None:
|
| 158 |
# Called before reset — auto-reset
|
|
|
|
| 210 |
reward=reward,
|
| 211 |
score=score,
|
| 212 |
)
|
| 213 |
+
|
| 214 |
self._last_obs = obs
|
| 215 |
+
# CRITICAL FIX: Return the standard OpenEnv 4-part tuple!
|
| 216 |
+
info = {"score": score, "error": error}
|
| 217 |
+
return obs, reward, done, info
|
| 218 |
|
| 219 |
@property
|
| 220 |
def state(self) -> NL2SQLState:
|