Team Ai
Apppublic

Kalletlamadhav/sql-optimization-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
reward.py139 linesDownload Raw Back to server
1# server/reward.py2 3import math4from .models import SQLOptReward, SQLOptAction5from .graders.speedup_grader import SpeedupGrader6from .graders.equivalence_grader import EquivalenceGrader7from .graders.antipattern_grader import AntiPatternGrader8from .graders.index_grader import IndexGrader9 10 11class RewardComposer:12 13    WEIGHTS = {14        'speedup': 0.35,15        'equivalence': 0.25,16        'pattern': 0.20,17        'index': 0.10,18        'simplicity': 0.10,19    }20 21    PENALTIES = {22        'syntax_error': -0.30,23        'wrong_results': -0.20,24        'slower_query': -0.10,25        'timeout_per_second': -0.05,26        'hack_detected': -0.40,27    }28 29    def __init__(self):30        self.speedup_grader = SpeedupGrader()31        self.equiv_grader = EquivalenceGrader()32        self.pattern_grader = AntiPatternGrader()33        self.index_grader = IndexGrader()34 35    def compute(36        self,37        task,38        action: SQLOptAction,39        orig_time: float,40        opt_time: float,41        opt_rows,42        opt_plan,43        hack,44        query_error45    ) -> SQLOptReward:46 47        penalty = 0.048 49        # ❗ Handle query failure immediately50        if query_error:51            return SQLOptReward(52                total=0.0,53                speedup_score=0.0,54                equivalence_score=0.0,55                pattern_score=0.0,56                index_score=0.0,57                simplicity_score=0.0,58                penalties=self.PENALTIES['syntax_error'],59                speedup_ratio=0.0,60                hack_detected=False,61                hack_type=None62            )63 64        # 1️⃣ Speedup Score65        speedup_ratio = orig_time / max(opt_time, 0.1)66 67        # Log scaling68        speedup_score = min(69            0.35,70            0.35 * (math.log10(max(speedup_ratio, 1)) / 2)71        )72 73        if speedup_ratio < 1.0:74            penalty += self.PENALTIES['slower_query']75 76        # 2️⃣ Equivalence Score77        equiv_score = self.equiv_grader.grade(task, action.optimized_query)78        equiv_score *= 0.2579 80        if equiv_score < 0.125:81            penalty += self.PENALTIES['wrong_results']82 83        # 3️⃣ Anti-pattern Score84        pattern_score = self.pattern_grader.grade(task, action) * 0.2085 86        # 4️⃣ Index Score87        index_score = self.index_grader.grade(task, action, opt_plan) * 0.1088 89        # 5️⃣ Simplicity Score90        simplicity_score = (91            self._simplicity_score(task.slow_query, action.optimized_query)92            * 0.1093        )94 95        # ⏱ Timeout penalty (>5s)96        if opt_time > 5000:97            penalty += self.PENALTIES['timeout_per_second'] * ((opt_time / 1000) - 5)98 99        # 🚨 Hack detection100        if hack:101            penalty += self.PENALTIES['hack_detected']102 103        # ✅ FIXED LINE BREAK104        raw_total = (105            speedup_score106            + equiv_score107            + pattern_score108            + index_score109            + simplicity_score110            + penalty111        )112 113        total = round(max(0.0, min(1.0, raw_total)), 4)114 115        return SQLOptReward(116            total=total,117            speedup_score=round(speedup_score, 4),118            equivalence_score=round(equiv_score, 4),119            pattern_score=round(pattern_score, 4),120            index_score=round(index_score, 4),121            simplicity_score=round(simplicity_score, 4),122            penalties=round(penalty, 4),123            speedup_ratio=round(speedup_ratio, 2),124            hack_detected=bool(hack),125            hack_type=hack126        )127 128    def _simplicity_score(self, original: str, optimized: str) -> float:129        """130        Reward simpler queries (only if logically correct)131        """132        orig_len = len(original.strip())133        opt_len = len(optimized.strip())134 135        if opt_len <= orig_len:136            return 1.0137 138        ratio = orig_len / max(opt_len, 1)139        return max(0.0, ratio)