diff --git a/MaxKernel/evaluation/code_adapter/code_adapter.py b/MaxKernel/evaluation/code_adapter/code_adapter.py index 42ea2ddc..27b84acc 100644 --- a/MaxKernel/evaluation/code_adapter/code_adapter.py +++ b/MaxKernel/evaluation/code_adapter/code_adapter.py @@ -1,3 +1,4 @@ +import ast import logging import time from typing import List, Optional, Union @@ -11,12 +12,20 @@ from evaluation.custom_types.kernel_task import KernelTask from hitl_agent.constants import MODEL_NAME +# Adaptation reproduces the whole input script, so the response is long. Ask +# for the model's full output budget rather than relying on the default. +MAX_OUTPUT_TOKENS = 65536 + logging.basicConfig( level=logging.INFO, format="%(asctime)s [%(levelname)s] - %(message)s", ) +class OutputTruncatedError(Exception): + """The model ran out of output budget before finishing the script.""" + + class CodeAdapter: """ Uses an LLM to refactor raw Python/JAX code into a structured format @@ -55,13 +64,16 @@ def adapt( ) prompt = self._get_adapt_optimized_prompt(original_code, get_inputs_code) - config = genai.types.GenerateContentConfig(temperature=0.1) + config = genai.types.GenerateContentConfig( + temperature=0.1, max_output_tokens=MAX_OUTPUT_TOKENS + ) attempt = 0 while attempt < self.max_retries: try: response = self.client.models.generate_content( model=MODEL_NAME, contents=prompt, config=config ) + self._check_finish_reason(response) code = response.text.strip() if code.startswith("```python"): code = code[len("```python") :].strip() @@ -78,7 +90,21 @@ def adapt( ): raise ValueError("LLM output did not contain the required sections.") + # The section headers all appear near the top of the file, so they are + # still present when a long response is cut short. Parse the result to + # catch truncation before it reaches disk and fails inside the harness. + try: + ast.parse(code) + except SyntaxError as e: + raise ValueError( + f"LLM output is not valid Python (line {e.lineno}: {e.msg}). " + "The response was most likely truncated." + ) from e + return code + except OutputTruncatedError: + # Deterministic: a retry produces the same overlong response. + raise except Exception as e: attempt += 1 wait_time = 2**attempt @@ -92,6 +118,26 @@ def adapt( f"Failed to refactor code after {self.max_retries} retries." ) + def _check_finish_reason(self, response) -> None: + """Raises if the model stopped for any reason other than finishing.""" + candidates = getattr(response, "candidates", None) + if not candidates: + raise ValueError("LLM returned no candidates.") + + reason = getattr(candidates[0], "finish_reason", None) + if reason is None or reason == genai.types.FinishReason.STOP: + return + + if reason == genai.types.FinishReason.MAX_TOKENS: + usage = getattr(response, "usage_metadata", None) + raise OutputTruncatedError( + f"LLM hit the {MAX_OUTPUT_TOKENS}-token output limit, so the " + f"refactored code is truncated. The input script is too large to " + f"adapt in one response. (usage: {usage})" + ) + + raise ValueError(f"LLM stopped early with finish_reason={reason}.") + def generate_kernel_task( self, task_id: str, diff --git a/MaxKernel/evaluation/harness_code.py b/MaxKernel/evaluation/harness_code.py index e8f17b02..6ade8a5a 100644 --- a/MaxKernel/evaluation/harness_code.py +++ b/MaxKernel/evaluation/harness_code.py @@ -5,9 +5,11 @@ import json import jax import jax.numpy as jnp +import numpy as np import importlib import importlib.util import os +import sys import traceback import xprof_utils @@ -17,7 +19,16 @@ def load_module_from_path(module_name, file_path): if spec is None or spec.loader is None: raise ImportError(f"Could not load {module_name} from {file_path}") module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) + # Register before exec_module: modules using `from __future__ import + # annotations` turn every annotation into a string, and dataclasses then + # resolves them via sys.modules[cls.__module__], which raises + # AttributeError on None if the module was never registered. + sys.modules[module_name] = module + try: + spec.loader.exec_module(module) + except BaseException: + sys.modules.pop(module_name, None) + raise return module @@ -98,6 +109,44 @@ def run_xprof(): return avg_wall_time, xprof_time +def diff_metrics(b, o, chunk_elems=1 << 24): + \"\"\"Max absolute and max relative difference between two outputs. + + Computed on the host in bounded-size chunks. `b` and `o` have already been + pulled off the device by `jax.device_get`, so the naive + `jnp.max(jnp.abs((b - o) / b))` ships them straight back: the true division + promotes to float, and for a large integer output (e.g. a 4 GiB uint8 paged + KV cache) that is a 16 GiB argument plus a 16 GiB result, which does not fit + in HBM. These are diagnostic metrics only -- `is_correct` comes from + `jnp.allclose` -- but raising here used to fail the whole case. + + Chunking also fixes two latent issues with the old expression: the + subtraction no longer wraps around for unsigned dtypes, and a NaN in one + chunk no longer suppresses the maximum found in the others. + \"\"\" + fb = np.asarray(b).reshape(-1) + fo = np.asarray(o).reshape(-1) + max_abs = 0.0 + max_rel = 0.0 + for i in range(0, fb.size, chunk_elems): + x = fb[i:i + chunk_elems].astype(np.float64) + y = fo[i:i + chunk_elems].astype(np.float64) + d = np.abs(x - y) + # max() over an empty slice is undefined; size is never 0 here but guard + # anyway so a zero-sized output cannot take down the comparison. + if d.size == 0: + continue + max_abs = max(max_abs, float(np.max(d))) + with np.errstate(divide="ignore", invalid="ignore"): + r = d / np.abs(x) + # Matches the previous definition |(b - o) / b|: division by a zero + # reference stays +inf, and 0/0 stays NaN rather than being counted. + r = r[~np.isnan(r)] + if r.size: + max_rel = max(max_rel, float(np.max(r))) + return max_abs, max_rel + + def main(): try: # Load task configuration from task.json @@ -266,8 +315,9 @@ def main(): if b.shape != o.shape: raise ValueError(f"Shape mismatch: {b.shape} vs {o.shape}") is_correct = is_correct and bool(jnp.allclose(b, o, atol=curr_atol, rtol=curr_rtol)) - max_abs_diff = max(max_abs_diff, float(jnp.max(jnp.abs(b - o)))) - max_rel_diff = max(max_rel_diff, float(jnp.max(jnp.abs((b - o) / b)))) + leaf_abs, leaf_rel = diff_metrics(b, o) + max_abs_diff = max(max_abs_diff, leaf_abs) + max_rel_diff = max(max_rel_diff, leaf_rel) except Exception as e: harness_logs.append(f"Correctness check failed for input {idx}: {e}") is_correct = False