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
6 changes: 6 additions & 0 deletions maxkernel/adk/evaluation/custom_types/kernel_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,3 +9,9 @@ class KernelTask:
input_gen_code: Optional[str] = None
atol: Optional[Union[float, List[float]]] = None
rtol: Optional[Union[float, List[float]]] = None
# When true, the harness sorts every output leaf along its last axis before
# comparing reference and optimized outputs. Use this for outputs whose
# order along the last axis is not part of the contract (e.g. top-k index
# sets), so an otherwise correct kernel is not failed for emitting the same
# elements in a different order. Defaults to element-wise comparison.
sort_before_compare: bool = False
1 change: 1 addition & 0 deletions maxkernel/adk/evaluation/evaluation_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,7 @@ def load_kernel_task_from_yaml(yaml_path: str) -> KernelTask:
input_gen_code=data.get("input_gen_code"),
atol=data.get("atol"),
rtol=data.get("rtol"),
sort_before_compare=bool(data.get("sort_before_compare", False)),
Comment thread
NinaCai marked this conversation as resolved.
)


Expand Down
7 changes: 7 additions & 0 deletions maxkernel/adk/evaluation/harness_code.py
Original file line number Diff line number Diff line change
Expand Up @@ -216,6 +216,7 @@ def main():
input_gen_code = task_data.get("input_gen_code")
task_atol = task_data.get("atol")
task_rtol = task_data.get("rtol")
sort_before_compare = task_data.get("sort_before_compare", False)

if input_gen_code:
ldict = {}
Expand Down Expand Up @@ -370,6 +371,12 @@ def main():
for b, o in zip(out_base_flat, out_optimized_flat):
if b.shape != o.shape:
raise ValueError(f"Shape mismatch: {b.shape} vs {o.shape}")
if sort_before_compare and np.ndim(b) > 0:
# Order along the last axis is not part of this task's output
# contract (e.g. top-k index sets): compare rows as sorted
# multisets. Done on the host, so it never affects timing.
b = np.sort(b, axis=-1)
o = np.sort(o, axis=-1)
is_correct = is_correct and bool(jnp.allclose(b, o, atol=curr_atol, rtol=curr_rtol))
leaf_abs, leaf_rel = diff_metrics(b, o)
max_abs_diff = max(max_abs_diff, leaf_abs)
Expand Down
1 change: 1 addition & 0 deletions maxkernel/adk/evaluation/jax_kernel_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -442,6 +442,7 @@ def _build_task_json(
"input_gen_code": task.input_gen_code,
"atol": effective_atol,
"rtol": effective_rtol,
"sort_before_compare": task.sort_before_compare,
}
with open(local_path, "w") as f:
json.dump(task_info, f)
Expand Down
Loading