sravaniamere commited on
Commit
82a5b1b
Β·
1 Parent(s): 2be7532
README.md CHANGED
@@ -66,18 +66,18 @@ wrong answer across steps.
66
 
67
  | Score | Condition |
68
  |-------|-----------|
69
- | `1.0` | Exact match after normalization (perfect fix) |
70
  | `0.7` | All correct tokens present, structure slightly off |
71
  | `0.4` | Most keywords correct and token overlap is high (β‰₯85% keywords, β‰₯75% tokens) |
72
  | `0.3` | Partial keyword and structure match (β‰₯65% keywords, β‰₯50% tokens) |
73
  | `0.2` | Basic `SELECT ... FROM ...` structure present |
74
- | `0.0` | Response is not valid SQL |
75
 
76
  A **stagnation penalty** of `βˆ’0.1` is applied when the agent submits the same
77
  reward-equivalent answer for two or more consecutive steps, encouraging active
78
  correction rather than looping.
79
 
80
- Episodes terminate when reward = 1.0 (success) or max steps is reached.
81
 
82
  ---
83
 
 
66
 
67
  | Score | Condition |
68
  |-------|-----------|
69
+ | `0.99` | Exact match after normalization (perfect fix) |
70
  | `0.7` | All correct tokens present, structure slightly off |
71
  | `0.4` | Most keywords correct and token overlap is high (β‰₯85% keywords, β‰₯75% tokens) |
72
  | `0.3` | Partial keyword and structure match (β‰₯65% keywords, β‰₯50% tokens) |
73
  | `0.2` | Basic `SELECT ... FROM ...` structure present |
74
+ | `0.01` | Response is not valid SQL |
75
 
76
  A **stagnation penalty** of `βˆ’0.1` is applied when the agent submits the same
77
  reward-equivalent answer for two or more consecutive steps, encouraging active
78
  correction rather than looping.
79
 
80
+ Episodes terminate when reward reaches `0.99` (success) or max steps is reached.
81
 
82
  ---
83
 
openenv.yaml CHANGED
@@ -76,15 +76,15 @@ action_space:
76
  type: string
77
 
78
  reward:
79
- range: [0.0, 1.0]
80
  description: >
81
- 1.0 = exact match, 0.7 = right tokens minor structure diff,
82
  0.4 = most keywords correct, 0.3 = partial match,
83
- 0.2 = basic structure present, 0.0 = invalid SQL.
84
  Stagnation penalty of -0.1 applied after 2+ identical-reward steps.
85
 
86
  scoring:
87
- reward_range: [0.0, 1.0]
88
  success_threshold: 0.5
89
  score_formula: mean(step_rewards)
90
 
 
76
  type: string
77
 
78
  reward:
79
+ range: [0.01, 0.99]
80
  description: >
81
+ 0.99 = exact match, 0.7 = right tokens minor structure diff,
82
  0.4 = most keywords correct, 0.3 = partial match,
83
+ 0.2 = basic structure present, 0.01 = invalid SQL.
84
  Stagnation penalty of -0.1 applied after 2+ identical-reward steps.
85
 
86
  scoring:
87
+ reward_range: [0.01, 0.99]
88
  success_threshold: 0.5
89
  score_formula: mean(step_rewards)
90
 
sql_env/__pycache__/__init__.cpython-314.pyc CHANGED
Binary files a/sql_env/__pycache__/__init__.cpython-314.pyc and b/sql_env/__pycache__/__init__.cpython-314.pyc differ
 
sql_env/__pycache__/env.cpython-314.pyc CHANGED
Binary files a/sql_env/__pycache__/env.cpython-314.pyc and b/sql_env/__pycache__/env.cpython-314.pyc differ
 
sql_env/__pycache__/grader.cpython-314.pyc CHANGED
Binary files a/sql_env/__pycache__/grader.cpython-314.pyc and b/sql_env/__pycache__/grader.cpython-314.pyc differ
 
sql_env/__pycache__/models.cpython-314.pyc CHANGED
Binary files a/sql_env/__pycache__/models.cpython-314.pyc and b/sql_env/__pycache__/models.cpython-314.pyc differ
 
sql_env/env.py CHANGED
@@ -33,7 +33,7 @@ class SQLCorrectionEnv:
33
  self._done: bool = False
34
  self._previous_attempt: Optional[str] = None
35
  self._last_feedback: Optional[str] = None
36
- self._last_reward: float = 0.0
37
  self._stagnation_count: int = 0
38
 
39
  # ── OpenEnv Interface ─────────────────────────────────────────────────────
@@ -50,7 +50,7 @@ class SQLCorrectionEnv:
50
  self._done = False
51
  self._previous_attempt = None
52
  self._last_feedback = None
53
- self._last_reward = 0.0
54
  self._stagnation_count = 0
55
 
56
  return self._make_observation()
@@ -77,14 +77,16 @@ class SQLCorrectionEnv:
77
  reward = max(0.01, reward - 0.1)
78
  else:
79
  self._stagnation_count = 0
80
-
 
 
81
  self._last_reward = reward
82
 
83
  feedback = generate_feedback(action, self._task, reward_model)
84
  self._last_feedback = feedback
85
  self._previous_attempt = action.corrected_query
86
 
87
- done = reward_model.value >= 0.95 or self._step_count >= self._task.max_steps
88
  self._done = done
89
 
90
  obs = self._make_observation()
 
33
  self._done: bool = False
34
  self._previous_attempt: Optional[str] = None
35
  self._last_feedback: Optional[str] = None
36
+ self._last_reward: float = 0.01
37
  self._stagnation_count: int = 0
38
 
39
  # ── OpenEnv Interface ─────────────────────────────────────────────────────
 
50
  self._done = False
51
  self._previous_attempt = None
52
  self._last_feedback = None
53
+ self._last_reward = 0.01
54
  self._stagnation_count = 0
55
 
56
  return self._make_observation()
 
77
  reward = max(0.01, reward - 0.1)
78
  else:
79
  self._stagnation_count = 0
80
+ # Final Clamp
81
+ reward = max(0.01, min(0.98, reward))
82
+
83
  self._last_reward = reward
84
 
85
  feedback = generate_feedback(action, self._task, reward_model)
86
  self._last_feedback = feedback
87
  self._previous_attempt = action.corrected_query
88
 
89
+ done = reward >= 0.95 or self._step_count >= self._task.max_steps
90
  self._done = done
91
 
92
  obs = self._make_observation()
sql_env/grader.py CHANGED
@@ -34,7 +34,7 @@ def _sql_keywords_present(query: str) -> set:
34
 
35
  def _clamp(value: float) -> float:
36
  """Ensure reward is strictly within (0, 1) as required by the OpenEnv spec."""
37
- return max(0.01, min(0.99, value))
38
 
39
 
40
  def grade(action: SQLAction, task: SQLTask) -> SQLReward:
@@ -55,7 +55,7 @@ def grade(action: SQLAction, task: SQLTask) -> SQLReward:
55
  # ── Level 1: Exact match ─────────────────────────────────────────────────
56
  if agent == correct:
57
  return SQLReward(
58
- value=_clamp(0.99),
59
  reason="Exact match β€” perfect correction.",
60
  )
61
 
@@ -106,7 +106,7 @@ def grade(action: SQLAction, task: SQLTask) -> SQLReward:
106
 
107
  def generate_feedback(action: SQLAction, task: SQLTask, reward: SQLReward) -> str:
108
  """Human-readable feedback shown in the next observation."""
109
- if reward.value >= 0.99:
110
  return "Correct! Query matches perfectly."
111
  if reward.value >= 0.70:
112
  return "Very close β€” check spacing or minor clause differences."
 
34
 
35
  def _clamp(value: float) -> float:
36
  """Ensure reward is strictly within (0, 1) as required by the OpenEnv spec."""
37
+ return max(0.02, min(0.98, value))
38
 
39
 
40
  def grade(action: SQLAction, task: SQLTask) -> SQLReward:
 
55
  # ── Level 1: Exact match ─────────────────────────────────────────────────
56
  if agent == correct:
57
  return SQLReward(
58
+ value=_clamp(0.98),
59
  reason="Exact match β€” perfect correction.",
60
  )
61
 
 
106
 
107
  def generate_feedback(action: SQLAction, task: SQLTask, reward: SQLReward) -> str:
108
  """Human-readable feedback shown in the next observation."""
109
+ if reward.value >= 0.98:
110
  return "Correct! Query matches perfectly."
111
  if reward.value >= 0.70:
112
  return "Very close β€” check spacing or minor clause differences."
sql_env/server.py CHANGED
@@ -23,7 +23,7 @@ class SQLCorrectionEnvironment(Environment):
23
  self._current_task = None
24
  self._step_count = 0
25
  self._done = False
26
- self._last_reward = 0.0
27
  self._rewards_history = []
28
  self._stagnation_count = 0
29
 
@@ -38,7 +38,7 @@ class SQLCorrectionEnvironment(Environment):
38
  self._current_task = random.choice(tasks)
39
  self._step_count = 0
40
  self._done = False
41
- self._last_reward = 0.0
42
  self._rewards_history = []
43
  self._stagnation_count = 0
44
  return self._make_observation(previous_attempt=None, feedback=None)
@@ -58,15 +58,18 @@ class SQLCorrectionEnvironment(Environment):
58
  reward = max(0.01, reward - 0.1)
59
  else:
60
  self._stagnation_count = 0
61
-
 
 
62
  self._last_reward = reward
 
63
  self._rewards_history.append(reward)
64
 
65
- done = (reward_obj.value >= 0.95) or (
66
  self._step_count >= self._current_task.max_steps
67
  )
68
  self._done = done
69
- feedback = generate_feedback(action, self._current_task, reward_obj)
70
 
71
  return self._make_observation(
72
  previous_attempt=action.corrected_query,
@@ -82,7 +85,7 @@ class SQLCorrectionEnvironment(Environment):
82
  step_count=0,
83
  max_steps=0,
84
  done=False,
85
- last_reward=0.0,
86
  rewards_history=[],
87
  )
88
  return SQLState(
 
23
  self._current_task = None
24
  self._step_count = 0
25
  self._done = False
26
+ self._last_reward = 0.01
27
  self._rewards_history = []
28
  self._stagnation_count = 0
29
 
 
38
  self._current_task = random.choice(tasks)
39
  self._step_count = 0
40
  self._done = False
41
+ self._last_reward = 0.01
42
  self._rewards_history = []
43
  self._stagnation_count = 0
44
  return self._make_observation(previous_attempt=None, feedback=None)
 
58
  reward = max(0.01, reward - 0.1)
59
  else:
60
  self._stagnation_count = 0
61
+ # Final Clamp
62
+ reward = max(0.01, min(0.98, reward))
63
+
64
  self._last_reward = reward
65
+
66
  self._rewards_history.append(reward)
67
 
68
+ done = (reward >= 0.95) or (
69
  self._step_count >= self._current_task.max_steps
70
  )
71
  self._done = done
72
+ feedback = generate_feedback(action, self._current_task, reward)
73
 
74
  return self._make_observation(
75
  previous_attempt=action.corrected_query,
 
85
  step_count=0,
86
  max_steps=0,
87
  done=False,
88
+ last_reward=0.01,
89
  rewards_history=[],
90
  )
91
  return SQLState(
sql_env/tasks/__pycache__/__init__.cpython-314.pyc CHANGED
Binary files a/sql_env/tasks/__pycache__/__init__.cpython-314.pyc and b/sql_env/tasks/__pycache__/__init__.cpython-314.pyc differ
 
sql_env/tasks/__pycache__/easy.cpython-314.pyc CHANGED
Binary files a/sql_env/tasks/__pycache__/easy.cpython-314.pyc and b/sql_env/tasks/__pycache__/easy.cpython-314.pyc differ
 
sql_env/tasks/__pycache__/hard.cpython-314.pyc CHANGED
Binary files a/sql_env/tasks/__pycache__/hard.cpython-314.pyc and b/sql_env/tasks/__pycache__/hard.cpython-314.pyc differ
 
sql_env/tasks/__pycache__/medium.cpython-314.pyc CHANGED
Binary files a/sql_env/tasks/__pycache__/medium.cpython-314.pyc and b/sql_env/tasks/__pycache__/medium.cpython-314.pyc differ
 
tests/test_env.py CHANGED
@@ -29,13 +29,13 @@ def _action(query: str) -> SQLAction:
29
  # ── Grader unit tests ─────────────────────────────────────────────────────────
30
 
31
  class TestGrader:
32
- def test_exact_match_returns_one(self):
33
  task = _make_task(
34
  "SELECT * FORM users",
35
  "SELECT * FROM users",
36
  )
37
  reward = grade(_action("SELECT * FROM users"), task)
38
- assert reward.value == 1.0
39
 
40
  def test_exact_match_case_insensitive(self):
41
  task = _make_task(
@@ -43,7 +43,7 @@ class TestGrader:
43
  "SELECT * FROM users",
44
  )
45
  reward = grade(_action("select * from users"), task)
46
- assert reward.value == 1.0
47
 
48
  def test_exact_match_trailing_semicolon(self):
49
  task = _make_task(
@@ -51,23 +51,23 @@ class TestGrader:
51
  "SELECT * FROM users",
52
  )
53
  reward = grade(_action("SELECT * FROM users;"), task)
54
- assert reward.value == 1.0
55
 
56
- def test_wrong_answer_not_one(self):
57
  task = _make_task(
58
  "SELECT * FORM users",
59
  "SELECT * FROM users",
60
  )
61
  reward = grade(_action("SELECT * FORM users"), task)
62
- assert reward.value < 1.0
63
 
64
- def test_completely_wrong_returns_zero(self):
65
  task = _make_task(
66
  "SELECT * FORM users",
67
  "SELECT * FROM users",
68
  )
69
  reward = grade(_action("hello world"), task)
70
- assert reward.value == 0.0
71
 
72
  def test_basic_structure_returns_02(self):
73
  task = _make_task(
@@ -78,7 +78,7 @@ class TestGrader:
78
  reward = grade(_action("SELECT * FORM users WHERE id = 1"), task)
79
  assert reward.value == pytest.approx(0.2, abs=0.05)
80
 
81
- def test_reward_range(self):
82
  task = _make_task(
83
  "SELCT * FORM users WEHRE id = 1",
84
  "SELECT * FROM users WHERE id = 1",
@@ -90,8 +90,8 @@ class TestGrader:
90
  "select * from users where id = 1",
91
  ]:
92
  reward = grade(_action(query), task)
93
- assert 0.0 <= reward.value <= 1.0, (
94
- f"Reward {reward.value} out of [0, 1] for query: {query}"
95
  )
96
 
97
  def test_feedback_not_empty(self):
@@ -142,24 +142,24 @@ class TestTaskCatalogue:
142
  )
143
 
144
  def test_grading_canonical_answer_returns_perfect(self):
145
- """Every task must return 1.0 when given its own canonical answer."""
146
  for difficulty, tasks in ALL_TASKS.items():
147
  for task in tasks:
148
  action = _action(task.canonical_answer)
149
  reward = grade(action, task)
150
- assert reward.value == 1.0, (
151
- f"{task.task_id}: canonical answer did not score 1.0 "
152
  f"(got {reward.value})"
153
  )
154
 
155
  def test_grading_broken_query_below_perfect(self):
156
- """Broken queries must score below 1.0."""
157
  for difficulty, tasks in ALL_TASKS.items():
158
  for task in tasks:
159
  action = _action(task.broken_query)
160
  reward = grade(action, task)
161
- assert reward.value < 1.0, (
162
- f"{task.task_id}: broken query unexpectedly scored 1.0"
163
  )
164
 
165
 
@@ -182,7 +182,7 @@ class TestEnvironment:
182
  env = SQLCorrectionEnv(difficulty="easy")
183
  await env.reset()
184
  result = await env.step(_action("SELECT * FROM users WHERE id = 1"))
185
- assert 0.0 <= result.reward <= 1.0
186
  assert isinstance(result.done, bool)
187
  assert result.observation.step_number == 1
188
 
@@ -204,7 +204,7 @@ class TestEnvironment:
204
  canonical = EASY_TASKS[0].canonical_answer
205
  result = await env.step(_action(canonical))
206
  assert result.done is True
207
- assert result.reward == pytest.approx(1.0)
208
 
209
  asyncio.run(run())
210
 
 
29
  # ── Grader unit tests ─────────────────────────────────────────────────────────
30
 
31
  class TestGrader:
32
+ def test_exact_match_returns_099(self):
33
  task = _make_task(
34
  "SELECT * FORM users",
35
  "SELECT * FROM users",
36
  )
37
  reward = grade(_action("SELECT * FROM users"), task)
38
+ assert reward.value == 0.99
39
 
40
  def test_exact_match_case_insensitive(self):
41
  task = _make_task(
 
43
  "SELECT * FROM users",
44
  )
45
  reward = grade(_action("select * from users"), task)
46
+ assert reward.value == 0.99
47
 
48
  def test_exact_match_trailing_semicolon(self):
49
  task = _make_task(
 
51
  "SELECT * FROM users",
52
  )
53
  reward = grade(_action("SELECT * FROM users;"), task)
54
+ assert reward.value == 0.99
55
 
56
+ def test_wrong_answer_not_perfect(self):
57
  task = _make_task(
58
  "SELECT * FORM users",
59
  "SELECT * FROM users",
60
  )
61
  reward = grade(_action("SELECT * FORM users"), task)
62
+ assert reward.value < 0.99
63
 
64
+ def test_completely_wrong_returns_001(self):
65
  task = _make_task(
66
  "SELECT * FORM users",
67
  "SELECT * FROM users",
68
  )
69
  reward = grade(_action("hello world"), task)
70
+ assert reward.value == 0.01
71
 
72
  def test_basic_structure_returns_02(self):
73
  task = _make_task(
 
78
  reward = grade(_action("SELECT * FORM users WHERE id = 1"), task)
79
  assert reward.value == pytest.approx(0.2, abs=0.05)
80
 
81
+ def test_reward_range_is_strictly_open(self):
82
  task = _make_task(
83
  "SELCT * FORM users WEHRE id = 1",
84
  "SELECT * FROM users WHERE id = 1",
 
90
  "select * from users where id = 1",
91
  ]:
92
  reward = grade(_action(query), task)
93
+ assert 0.0 < reward.value < 1.0, (
94
+ f"Reward {reward.value} out of (0, 1) for query: {query}"
95
  )
96
 
97
  def test_feedback_not_empty(self):
 
142
  )
143
 
144
  def test_grading_canonical_answer_returns_perfect(self):
145
+ """Every task must return 0.99 when given its own canonical answer."""
146
  for difficulty, tasks in ALL_TASKS.items():
147
  for task in tasks:
148
  action = _action(task.canonical_answer)
149
  reward = grade(action, task)
150
+ assert reward.value == 0.99, (
151
+ f"{task.task_id}: canonical answer did not score 0.99 "
152
  f"(got {reward.value})"
153
  )
154
 
155
  def test_grading_broken_query_below_perfect(self):
156
+ """Broken queries must score below the perfect 0.99 score."""
157
  for difficulty, tasks in ALL_TASKS.items():
158
  for task in tasks:
159
  action = _action(task.broken_query)
160
  reward = grade(action, task)
161
+ assert reward.value < 0.99, (
162
+ f"{task.task_id}: broken query unexpectedly scored 0.99"
163
  )
164
 
165
 
 
182
  env = SQLCorrectionEnv(difficulty="easy")
183
  await env.reset()
184
  result = await env.step(_action("SELECT * FROM users WHERE id = 1"))
185
+ assert 0.0 < result.reward < 1.0
186
  assert isinstance(result.done, bool)
187
  assert result.observation.step_number == 1
188
 
 
204
  canonical = EASY_TASKS[0].canonical_answer
205
  result = await env.step(_action(canonical))
206
  assert result.done is True
207
+ assert result.reward == pytest.approx(0.99)
208
 
209
  asyncio.run(run())
210