Spaces:
Sleeping
Sleeping
File size: 4,764 Bytes
5716a3a 95707b2 5716a3a 2194233 5716a3a 2194233 95707b2 2194233 5716a3a 2194233 5716a3a 2194233 5716a3a 2194233 5716a3a 2194233 5716a3a 2194233 5716a3a 2194233 5716a3a 2194233 95707b2 2194233 95707b2 2194233 5716a3a 2194233 5716a3a 2194233 5716a3a 2194233 5716a3a 2194233 95707b2 5716a3a 2194233 | 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 112 113 114 115 116 | 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:
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:
"""
5-level grader with partial progress signals.
1.0 β exact normalized match (perfect fix)
0.7 β same token set, minor structural/whitespace differences
0.4 β most SQL keywords correct AND high token overlap
0.3 β partial keyword and structure match
0.2 β basic SELECT/FROM structure present
0.0 β not recognizable SQL
"""
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/alias diff) βββββ
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 "
f"({kw_overlap:.0%} keyword match, {token_overlap:.0%} token match)."
),
)
# ββ Level 3.5: Partial keyword and structure match ββββββββββββββββββββββββ
if kw_overlap >= 0.65 and token_overlap >= 0.50:
return SQLReward(
value=0.3,
reason=(
f"Partial keyword and structure match "
f"({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 the 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.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 ..."
|