Spaces:
Sleeping
Sleeping
File size: 3,717 Bytes
5716a3a | 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 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 | import re
from sql_env.models import SQLAction, SQLTask, SQLReward
def _normalize(query: str) -> str:
"""Uppercase, collapse whitespace, strip trailing semicolons."""
q = query.strip().upper()
q = re.sub(r'\s+', ' ', q)
q = q.rstrip(';').strip()
return q
def _tokenize(query: str) -> set:
return set(re.findall(r"[A-Z0-9_'*.=><]+", _normalize(query)))
def _sql_keywords_present(query: str) -> set:
keywords = {
'SELECT', 'FROM', 'WHERE', 'GROUP', 'BY', 'HAVING',
'ORDER', 'JOIN', 'INNER', 'LEFT', 'RIGHT', 'OUTER',
'BETWEEN', 'DESC', 'ASC', 'LIMIT', 'COUNT', 'SUM',
'AVG', 'MAX', 'MIN', 'AS', 'ON', 'AND', 'OR', 'NOT',
'IN', 'LIKE', 'IS', 'NULL', 'DISTINCT'
}
normed = _normalize(query)
found = set()
for kw in keywords:
if re.search(r'\b' + kw + r'\b', normed):
found.add(kw)
return found
def grade(action: SQLAction, task: SQLTask) -> SQLReward:
"""
4-level grader with partial progress signals.
1.0 β exact normalized match
0.7 β same tokens, minor whitespace/alias differences
0.4 β key SQL keywords all present and correct table/column names
0.2 β basic SELECT/FROM structure present
0.0 β completely wrong
"""
agent = _normalize(action.corrected_query)
correct = _normalize(task.canonical_answer)
# ββ Level 1: Exact match βββββββββββββββββββββββββββββββββ
if agent == correct:
return SQLReward(value=1.0, reason="Exact match β perfect correction.")
# ββ Level 2: Same token set (right words, minor ordering) β
agent_tokens = _tokenize(action.corrected_query)
correct_tokens = _tokenize(task.canonical_answer)
if agent_tokens == correct_tokens:
return SQLReward(
value=0.7,
reason="All correct tokens present but structure differs slightly."
)
# ββ Level 3: Most keywords correct + high token overlap βββ
correct_kws = _sql_keywords_present(task.canonical_answer)
agent_kws = _sql_keywords_present(action.corrected_query)
kw_overlap = len(correct_kws & agent_kws) / max(len(correct_kws), 1)
token_overlap = len(agent_tokens & correct_tokens) / max(len(correct_tokens), 1)
if kw_overlap >= 0.85 and token_overlap >= 0.75:
return SQLReward(
value=0.4,
reason=f"Most keywords correct ({kw_overlap:.0%} keyword match, {token_overlap:.0%} token match)."
)
# ββ Level 4: Basic structure present βββββββββββββββββββββ
if 'SELECT' in agent and 'FROM' in agent:
return SQLReward(
value=0.2,
reason="Basic SELECT/FROM structure present but significant errors remain."
)
# ββ Level 0: No recognizable SQL βββββββββββββββββββββββββ
return SQLReward(value=0.0, reason="Response is not valid SQL.")
def generate_feedback(action: SQLAction, task: SQLTask, reward: SQLReward) -> str:
"""Human-readable feedback shown in next observation."""
if reward.value == 1.0:
return "Correct! Query matches perfectly."
if reward.value >= 0.7:
return "Very close β check spacing or minor clause differences."
if reward.value >= 0.4:
return "Good progress β most keywords are right, but check for typos in keywords or column names."
if reward.value >= 0.2:
return "Basic structure is there β look carefully at every SQL keyword for typos."
return "The response doesn't look like valid SQL. Start with SELECT ... FROM ..."
|