Spaces:
Sleeping
Sleeping
File size: 4,558 Bytes
5716a3a 95707b2 5716a3a 95707b2 5716a3a a785d13 95707b2 a785d13 95707b2 5716a3a 95707b2 2112542 95707b2 5716a3a 95707b2 5716a3a 95707b2 5716a3a 95707b2 5716a3a 95707b2 5716a3a 2112542 5716a3a 95707b2 5716a3a 95707b2 5716a3a 95707b2 | 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 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 | import re
from sql_env.models import SQLAction, SQLTask, SQLReward
def _clamp(value: float) -> float:
"""Ensure score is strictly between 0 and 1 (gt=0.0, lt=1.0)."""
return max(0.001, min(0.999, value))
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:
normed = _normalize(query)
normed = re.sub(r"'[^']*'", '__STR__', normed) # normalize string literals
return set(re.findall(r"[A-Z0-9_'*.=><]+", normed))
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.
0.99 β 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.01 β completely wrong
All scores are clamped to be strictly between 0.0 and 1.0.
"""
agent = _normalize(action.corrected_query)
correct = _normalize(task.canonical_answer)
# ββ Level 1: Exact match βββββββββββββββββββββββββββββββββ
if agent == correct:
return SQLReward(
value=_clamp(0.999),
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=_clamp(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=_clamp(0.4),
reason=f"Most keywords correct ({kw_overlap:.0%} keyword match, {token_overlap:.0%} token match)."
)
# ββ Level 3.5: Partial keyword and structure match ββββββββ
if kw_overlap >= 0.7 and token_overlap >= 0.55:
return SQLReward(
value=_clamp(0.3),
reason=f"Partial keyword and structure match ({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=_clamp(0.2),
reason="Basic SELECT/FROM structure present but significant errors remain."
)
# ββ Level 0: No recognizable SQL βββββββββββββββββββββββββ
return SQLReward(value=_clamp(0.001), 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 >= 0.99:
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.3:
return "Partial match β right direction but several keywords or columns are off."
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 ..." |