md896 commited on
Commit
9b71d1b
·
1 Parent(s): 8e7c622

Harden strict (0,1) scoring boundaries across runtime and config.

Browse files

Clamp reward outputs to non-boundary values, align schema ranges, and remove inference fallbacks that could emit 0.0 or 1.0 during evaluation.

ss:

Files changed (5) hide show
  1. inference.py +4 -4
  2. openenv.yaml +2 -2
  3. server/models.py +2 -2
  4. server/reward.py +9 -2
  5. tests/test_reward.py +10 -10
inference.py CHANGED
@@ -218,7 +218,7 @@ def run_task(
218
 
219
  rewards = []
220
  steps_taken = 0
221
- score = 0.0
222
  success = False
223
 
224
  with httpx.Client(base_url=ENV_BASE_URL, timeout=30.0) as http:
@@ -245,11 +245,11 @@ def run_task(
245
  step_resp.raise_for_status()
246
  step_result = step_resp.json()
247
  except Exception as e:
248
- log_step(step=step, action=str(action_dict), reward=0.0, done=False, error=str(e))
249
  continue
250
 
251
  obs = step_result["observation"]
252
- reward = float(step_result.get("reward") or 0.0)
253
  done = step_result["done"]
254
  error = None
255
  info = step_result.get("info") or {}
@@ -264,7 +264,7 @@ def run_task(
264
  rewards.append(reward)
265
  reward_history.append(reward)
266
  steps_taken = step
267
- score = float(info.get("grade_score") or obs.get("current_score") or 0.0)
268
 
269
  log_step(step=step, action=action_str, reward=reward, done=done, error=error)
270
 
 
218
 
219
  rewards = []
220
  steps_taken = 0
221
+ score = MIN_STRICT_SCORE
222
  success = False
223
 
224
  with httpx.Client(base_url=ENV_BASE_URL, timeout=30.0) as http:
 
245
  step_resp.raise_for_status()
246
  step_result = step_resp.json()
247
  except Exception as e:
248
+ log_step(step=step, action=str(action_dict), reward=MIN_STRICT_SCORE, done=False, error=str(e))
249
  continue
250
 
251
  obs = step_result["observation"]
252
+ reward = float(step_result.get("reward") or MIN_STRICT_SCORE)
253
  done = step_result["done"]
254
  error = None
255
  info = step_result.get("info") or {}
 
264
  rewards.append(reward)
265
  reward_history.append(reward)
266
  steps_taken = step
267
+ score = float(info.get("grade_score") or obs.get("current_score") or MIN_STRICT_SCORE)
268
 
269
  log_step(step=step, action=action_str, reward=reward, done=done, error=error)
270
 
openenv.yaml CHANGED
@@ -37,7 +37,7 @@ tasks:
37
  description: "Fix 5 bugs: correlated subquery, window function, duplicate rows, date logic, CTE scope"
38
 
39
  api:
40
- base_url: "https://YOUR-USERNAME-sql-debug-env.hf.space"
41
  reset: "/reset"
42
  step: "/step"
43
  state: "/state"
@@ -77,7 +77,7 @@ action_space:
77
  description: "Reset to original broken query (penalty: -0.05)"
78
 
79
  reward:
80
- range: [0.0, 1.0]
81
  components:
82
  - name: correctness
83
  range: [0.0, 0.6]
 
37
  description: "Fix 5 bugs: correlated subquery, window function, duplicate rows, date logic, CTE scope"
38
 
39
  api:
40
+ base_url: "https://md896-sql-debug-env.hf.space"
41
  reset: "/reset"
42
  step: "/step"
43
  state: "/state"
 
77
  description: "Reset to original broken query (penalty: -0.05)"
78
 
79
  reward:
80
+ range: [0.001, 0.999]
81
  components:
82
  - name: correctness
83
  range: [0.0, 0.6]
server/models.py CHANGED
@@ -86,7 +86,7 @@ class SQLDebugObservation(BaseModel):
86
  # Progress
87
  steps_taken: int
88
  steps_remaining: int
89
- current_score: float = Field(description="Current score 0.0-1.0 for this episode")
90
 
91
  # Contextual help (populated based on action type)
92
  schema_info: Optional[SchemaInfo] = None
@@ -112,7 +112,7 @@ class SQLDebugReward(BaseModel):
112
  - schema_bonus: 0.0-0.1 for queries that reference correct tables/columns
113
  - penalties: negative values for reset_query, infinite loops, destructive SQL
114
  """
115
- value: float = Field(ge=0.0, le=1.0, description="Total reward for this step")
116
  correctness: float = Field(ge=0.0, le=0.6)
117
  efficiency: float = Field(ge=0.0, le=0.2)
118
  syntax_progress: float = Field(ge=0.0, le=0.1)
 
86
  # Progress
87
  steps_taken: int
88
  steps_remaining: int
89
+ current_score: float = Field(description="Current score in strict range (0, 1) for this episode")
90
 
91
  # Contextual help (populated based on action type)
92
  schema_info: Optional[SchemaInfo] = None
 
112
  - schema_bonus: 0.0-0.1 for queries that reference correct tables/columns
113
  - penalties: negative values for reset_query, infinite loops, destructive SQL
114
  """
115
+ value: float = Field(ge=0.001, le=0.999, description="Total reward for this step")
116
  correctness: float = Field(ge=0.0, le=0.6)
117
  efficiency: float = Field(ge=0.0, le=0.2)
118
  syntax_progress: float = Field(ge=0.0, le=0.1)
server/reward.py CHANGED
@@ -16,6 +16,13 @@ Total range: 0.0 to 1.0 (clamped to [0.0, 1.0])
16
  from typing import Optional, List, Dict, Any
17
  from .models import SQLDebugReward
18
 
 
 
 
 
 
 
 
19
 
20
  def compute_reward(
21
  action_type: str,
@@ -33,7 +40,7 @@ def compute_reward(
33
  Args:
34
  action_type: The action taken this step
35
  query_result: Result dict from EpisodeDatabase.execute_query()
36
- grade_score: 0.0-1.0 score from task grader
37
  steps_taken: How many steps have been used (1-indexed)
38
  max_steps: Maximum steps for this task
39
  previous_best_score: Best grade score seen so far
@@ -103,7 +110,7 @@ def compute_reward(
103
  penalty += 0.03
104
 
105
  total_raw = correctness + efficiency + syntax_progress + schema_bonus - penalty
106
- total = round(max(0.0, min(1.0, total_raw)), 4)
107
 
108
  breakdown = (
109
  f"correctness={correctness:.3f} + "
 
16
  from typing import Optional, List, Dict, Any
17
  from .models import SQLDebugReward
18
 
19
+ MIN_STRICT_SCORE = 0.001
20
+ MAX_STRICT_SCORE = 0.999
21
+
22
+
23
+ def _strict_score(value: float) -> float:
24
+ return round(min(MAX_STRICT_SCORE, max(MIN_STRICT_SCORE, value)), 4)
25
+
26
 
27
  def compute_reward(
28
  action_type: str,
 
40
  Args:
41
  action_type: The action taken this step
42
  query_result: Result dict from EpisodeDatabase.execute_query()
43
+ grade_score: strict (0, 1) score from task grader
44
  steps_taken: How many steps have been used (1-indexed)
45
  max_steps: Maximum steps for this task
46
  previous_best_score: Best grade score seen so far
 
110
  penalty += 0.03
111
 
112
  total_raw = correctness + efficiency + syntax_progress + schema_bonus - penalty
113
+ total = _strict_score(total_raw)
114
 
115
  breakdown = (
116
  f"correctness={correctness:.3f} + "
tests/test_reward.py CHANGED
@@ -8,42 +8,42 @@ class TestReward(unittest.TestCase):
8
  reward = compute_reward(
9
  action_type="submit_query",
10
  query_result={"success": True},
11
- grade_score=1.0,
12
  steps_taken=1,
13
  max_steps=10,
14
- previous_best_score=0.0,
15
  schema_tables=["t1", "t2"],
16
  submitted_query="SELECT * FROM t1 JOIN t2",
17
  )
18
- self.assertAlmostEqual(reward.value, 1.0, places=4)
19
 
20
  def test_reset_query_penalty(self):
21
  reward = compute_reward(
22
  action_type="reset_query",
23
  query_result=None,
24
- grade_score=0.0,
25
  steps_taken=1,
26
  max_steps=10,
27
- previous_best_score=0.0,
28
  schema_tables=[],
29
  submitted_query=None,
30
  )
31
- self.assertAlmostEqual(reward.value, 0.0, places=4)
32
 
33
  def test_inspect_schema_urgency_penalty(self):
34
  # Make steps_remaining <= 2 and grade_score < 0.5 to trigger urgency penalty.
35
  reward = compute_reward(
36
  action_type="inspect_schema",
37
  query_result=None,
38
- grade_score=0.0,
39
  steps_taken=8,
40
  max_steps=9,
41
- previous_best_score=0.0,
42
  schema_tables=[],
43
  submitted_query=None,
44
  )
45
- # syntax_progress=0.01, penalty=0.03 => total_raw=-0.02, clamped to 0.0
46
- self.assertAlmostEqual(reward.value, 0.0, places=4)
47
 
48
 
49
  if __name__ == "__main__":
 
8
  reward = compute_reward(
9
  action_type="submit_query",
10
  query_result={"success": True},
11
+ grade_score=0.999,
12
  steps_taken=1,
13
  max_steps=10,
14
+ previous_best_score=0.001,
15
  schema_tables=["t1", "t2"],
16
  submitted_query="SELECT * FROM t1 JOIN t2",
17
  )
18
+ self.assertAlmostEqual(reward.value, 0.999, places=4)
19
 
20
  def test_reset_query_penalty(self):
21
  reward = compute_reward(
22
  action_type="reset_query",
23
  query_result=None,
24
+ grade_score=0.001,
25
  steps_taken=1,
26
  max_steps=10,
27
+ previous_best_score=0.001,
28
  schema_tables=[],
29
  submitted_query=None,
30
  )
31
+ self.assertAlmostEqual(reward.value, 0.001, places=4)
32
 
33
  def test_inspect_schema_urgency_penalty(self):
34
  # Make steps_remaining <= 2 and grade_score < 0.5 to trigger urgency penalty.
35
  reward = compute_reward(
36
  action_type="inspect_schema",
37
  query_result=None,
38
+ grade_score=0.001,
39
  steps_taken=8,
40
  max_steps=9,
41
+ previous_best_score=0.001,
42
  schema_tables=[],
43
  submitted_query=None,
44
  )
45
+ # syntax_progress=0.01, penalty=0.03 => total_raw=-0.02, clamped to strict min
46
+ self.assertAlmostEqual(reward.value, 0.001, places=4)
47
 
48
 
49
  if __name__ == "__main__":