Compilertetris / runtime_core.py
github-actions
Auto deploy from GitHub
eb408af
Raw
History Blame Contribute Delete
10.8 kB
"""
Shared runtime for the Gradio Space: mock compiler env + Deliverable 2 formatting.
Sourced from `compiler_optimization_grpo.ipynb` and
`role2_deliverable3_training_loop (2) (1) (1).ipynb`.
"""
from __future__ import annotations
import copy
import json
import re
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Tuple
# --- Deliverable 2 (LLM-facing pseudo-asm + pass-array parsing) -----------------
class Deliverable2_Formatter:
@staticmethod
def translate_state(raw_json: list) -> str:
"""Translate raw JSON IR into compact pseudo-assembly."""
if not isinstance(raw_json, list) or not raw_json:
return "; (empty program — 0 instructions)"
pseudo_assembly: list[str] = []
for i, instruction in enumerate(raw_json):
if not isinstance(instruction, dict):
pseudo_assembly.append(f"{i}. NOP")
continue
op = str(instruction.get("op", "UNKNOWN")).upper()
args = ", ".join(str(arg) for arg in instruction.get("args", []))
dest = instruction.get("dest", "")
if dest:
line = f"{i}. {dest} = {op} {args}".rstrip()
else:
line = f"{i}. {op} {args}".rstrip()
pseudo_assembly.append(line)
return "\n".join(pseudo_assembly)
@staticmethod
def extract_action_array(llm_output: str) -> list:
"""Best-effort extraction of JSON pass arrays from noisy LLM output."""
text = (llm_output or "").strip()
if not text:
raise ValueError("Invalid JSON format")
try:
parsed = json.loads(text)
if isinstance(parsed, list):
return parsed
except json.JSONDecodeError:
pass
cleaned = re.sub(r"```(?:json)?", "", text, flags=re.IGNORECASE).replace("```", "").strip()
if cleaned != text:
try:
parsed = json.loads(cleaned)
if isinstance(parsed, list):
return parsed
except json.JSONDecodeError:
pass
match = re.search(r"\[.*?\]", text, re.DOTALL)
if match:
candidate = match.group(0)
try:
parsed = json.loads(candidate)
if isinstance(parsed, list):
return parsed
except json.JSONDecodeError:
try:
parsed = json.loads(candidate.replace("'", '"'))
if isinstance(parsed, list):
return parsed
except json.JSONDecodeError:
pass
raise ValueError("Invalid JSON format")
# --- OpenEnv-style compiler environment (mock engine) --------------------------
class MCPEnvironment:
"""Minimal stub. In production: `from openenv import MCPEnvironment`."""
def reset(self, *args, **kwargs):
raise NotImplementedError
def step(self, *args, **kwargs):
raise NotImplementedError
def state(self):
raise NotImplementedError
@dataclass
class StepResult:
observation: str
reward: float
done: bool
info: Dict[str, Any] = field(default_factory=dict)
@dataclass
class EpisodeStats:
steps_taken: int = 0
total_reward: float = 0.0
passes_applied: List[str] = field(default_factory=list)
invalid_actions: int = 0
no_ops: int = 0
baseline_cycles: int = 0
final_cycles: int = 0
@property
def total_improvement_pct(self) -> float:
if self.baseline_cycles == 0:
return 0.0
return ((self.baseline_cycles - self.final_cycles) / self.baseline_cycles) * 100.0
class CompilerOptimizationEnv(MCPEnvironment):
TIME_TAX: float = 1.0
NO_OP_PENALTY: float = -2.0
INVALID_ACTION_PENALTY: float = -5.0
MAX_INVALID_ACTIONS: int = 3
TERMINAL_BONUS_SCALE: float = 0.5
def __init__(
self,
role1_engine,
role3_passes: Dict[str, Any],
max_steps: int = 10,
curriculum_level: int = 1,
):
self.engine = role1_engine
self.passes = role3_passes
self.max_steps = max_steps
self.curriculum_level = curriculum_level
self._valid_actions = frozenset(self.passes.keys())
self._stats: Optional[EpisodeStats] = None
self.original_program = None
self.current_program = None
self.previous_cycles = 0
self._consecutive_invalid = 0
def reset(self, new_program_json: List[Dict]) -> str:
self.original_program = copy.deepcopy(new_program_json)
self.current_program = copy.deepcopy(new_program_json)
self.previous_cycles = self._safe_count_cycles(self.current_program)
self._consecutive_invalid = 0
self._stats = EpisodeStats(
baseline_cycles=self.previous_cycles,
final_cycles=self.previous_cycles,
)
return self.state()
def state(self) -> str:
assert self.current_program is not None
return self._program_to_pseudoasm(self.current_program)
def step(self, action_string: str) -> StepResult:
assert self._stats is not None, "Call reset() before step()."
self._stats.steps_taken += 1
if action_string not in self._valid_actions:
return self._handle_invalid_action(action_string)
candidate_program = self.passes[action_string](copy.deepcopy(self.current_program))
is_valid = self.engine.verify_equivalence(self.original_program, candidate_program)
if not is_valid:
return self._handle_semantic_violation()
new_cycles = self._safe_count_cycles(candidate_program)
reward, info = self._compute_reward(action_string, new_cycles)
self.current_program = candidate_program
self.previous_cycles = new_cycles
self._stats.final_cycles = new_cycles
self._stats.total_reward += reward
self._stats.passes_applied.append(action_string)
self._consecutive_invalid = 0
done = self._stats.steps_taken >= self.max_steps
if done:
terminal_bonus = self._terminal_bonus()
reward += terminal_bonus
info["terminal_bonus"] = terminal_bonus
info["reason"] = "max_steps_reached"
info["episode_stats"] = self._episode_summary()
return StepResult(self.state(), reward, done, info)
def _compute_reward(self, action: str, new_cycles: int) -> Tuple[float, Dict]:
info: Dict[str, Any] = {"action": action}
if self.previous_cycles == 0:
return -self.TIME_TAX, {**info, "note": "zero_baseline"}
old_cycles = self.previous_cycles
delta_pct = ((old_cycles - new_cycles) / old_cycles) * 100.0
if new_cycles == old_cycles:
reward = self.NO_OP_PENALTY
if self._stats is not None:
self._stats.no_ops += 1
info["no_op"] = True
else:
reward = delta_pct - self.TIME_TAX
info["delta_pct"] = round(delta_pct, 3)
info["prev_cycles"] = old_cycles
info["new_cycles"] = new_cycles
return reward, info
def _terminal_bonus(self) -> float:
if self._stats is None:
return 0.0
return max(0.0, self._stats.total_improvement_pct * self.TERMINAL_BONUS_SCALE)
def _compute_crash_penalty(self) -> float:
return -2.0 * (100.0 * self.max_steps)
def _handle_invalid_action(self, action: str) -> StepResult:
self._consecutive_invalid += 1
if self._stats is not None:
self._stats.invalid_actions += 1
done = self._consecutive_invalid >= self.MAX_INVALID_ACTIONS
info = {
"error": f"Unknown action: '{action}'",
"valid_actions": sorted(self._valid_actions),
"consecutive_invalid": self._consecutive_invalid,
}
if done:
info["reason"] = "too_many_invalid_actions"
info["episode_stats"] = self._episode_summary()
return StepResult(self.state(), self.INVALID_ACTION_PENALTY, done, info)
def _handle_semantic_violation(self) -> StepResult:
return StepResult(
self.state(),
self._compute_crash_penalty(),
True,
{
"error": "Semantic equivalence check FAILED.",
"reason": "semantic_violation",
"episode_stats": self._episode_summary(),
},
)
@staticmethod
def _program_to_pseudoasm(program: List[Dict]) -> str:
if not program:
return "; (empty program)"
lines = []
for i, instr in enumerate(program):
op = instr.get("op", "NOP")
args = instr.get("args", [])
dest = instr.get("dest")
typ = instr.get("type", "")
arg_str = ", ".join(str(a) for a in args)
type_hint = f":{typ}" if typ else ""
if dest:
lines.append(f" {i:>3}: {dest}{type_hint} = {op} {arg_str}")
else:
lines.append(f" {i:>3}: {op} {arg_str}")
return "\n".join(lines)
def _safe_count_cycles(self, program: List[Dict]) -> int:
return max(0, int(self.engine.execute_and_count_cycles(program)))
def _episode_summary(self) -> Dict:
s = self._stats
if s is None:
return {}
return {
"steps": s.steps_taken,
"total_reward": round(s.total_reward, 3),
"passes_applied": s.passes_applied,
"invalid_actions": s.invalid_actions,
"no_ops": s.no_ops,
"baseline_cycles": s.baseline_cycles,
"final_cycles": s.final_cycles,
"total_improvement_pct": round(s.total_improvement_pct, 3),
}
def available_actions(self) -> List[str]:
return sorted(self._valid_actions)
class MockEngine:
"""Stub engine: cycles = instruction count, all programs semantically valid."""
def execute_and_count_cycles(self, program):
return len(program)
def verify_equivalence(self, original, candidate):
return True
MOCK_PASSES = {
"constant_folding": lambda p: p[:-1] if len(p) > 1 else p,
"dead_code_elimination": lambda p: p[:-1] if len(p) > 2 else p,
"loop_unrolling": lambda p: p,
}
SAMPLE_PROGRAM = [
{"op": "const", "dest": "x", "args": ["5"], "type": "int"},
{"op": "const", "dest": "y", "args": ["3"], "type": "int"},
{"op": "add", "dest": "z", "args": ["x", "y"], "type": "int"},
{"op": "mul", "dest": "w", "args": ["z", "x"], "type": "int"},
{"op": "ret", "args": ["w"]},
]