Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 47 additions & 1 deletion MaxKernel/evaluation/code_adapter/code_adapter.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import ast
import logging
import time
from typing import List, Optional, Union
Expand All @@ -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
Expand Down Expand Up @@ -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()
Expand All @@ -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
Expand All @@ -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,
Expand Down
56 changes: 53 additions & 3 deletions MaxKernel/evaluation/harness_code.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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
Comment on lines +26 to +31

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

If module_name was already present in sys.modules before calling load_module_from_path, unconditionally popping it on failure will delete the pre-existing module instead of restoring it. It is safer to capture the previous value of sys.modules[module_name] and restore it if execution fails.

Suggested change
sys.modules[module_name] = module
try:
spec.loader.exec_module(module)
except Exception:
sys.modules.pop(module_name, None)
raise
old_module = sys.modules.get(module_name)\n sys.modules[module_name] = module\n try:\n spec.loader.exec_module(module)\n except Exception:\n if old_module is None:\n sys.modules.pop(module_name, None)\n else:\n sys.modules[module_name] = old_module\n raise

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

If it gets to Exception, it will fail anyway.

return module


Expand Down Expand Up @@ -98,6 +109,44 @@ def run_xprof():
return avg_wall_time, xprof_time


def diff_metrics(b, o, chunk_elems=1 << 24):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this function added because the code can run into OOM during the diff calculation? It might be beneficial if we can make this comment easier to understand.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Config 3 (csa_decode_bs512) has a uint8 cache of (32769, 256, 4, 128) = 4 GiB. b and o are host numpy arrays by this point, so (b - o) / b is a numpy true division → promotes to float64 on the host. Then jnp.abs(...) ships it to the TPU as f32: a 16 GiB argument plus a 16 GiB result = 32.00 GiB against a 31.25 GiB chip, overflowing by 773 M.

\"\"\"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
Expand Down Expand Up @@ -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))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

If the diff can make OOM, why this line will not make OOM?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

OOM happens in abs calculation. allclose survives because it's jitted and XLA fuses the promotion

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
Expand Down
Loading