Skip to content

Commit 9dc5630

Browse files
committed
feat: implement TileLang-based DSA split-K indexer loss, add regression tests, and clamp warmup_steps in ScheduleEditor
1 parent 4f8473a commit 9dc5630

13 files changed

Lines changed: 156280 additions & 378 deletions

backend.log

Lines changed: 152965 additions & 0 deletions
Large diffs are not rendered by default.

cppmega_mlx/nn/_tilelang/_mlx_runtime.py

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -810,6 +810,12 @@ def wrap_tilelang_metal_kernel(
810810
"output_buffer_names contains buffers missing from the Metal "
811811
f"signature: {missing_outputs!r}; parsed={buffer_names!r}"
812812
)
813+
aliased_names = set(input_sources) & set(output_sources)
814+
if aliased_names:
815+
raise MLXRuntimeError(
816+
f"Input and output buffer names must be mutually disjoint (aliasing is not supported): "
817+
f"aliased={sorted(list(aliased_names))}"
818+
)
813819
explicit_sources = input_sources + output_sources
814820
if len(set(explicit_sources)) != len(explicit_sources):
815821
raise MLXRuntimeError(
@@ -833,6 +839,14 @@ def wrap_tilelang_metal_kernel(
833839
f"{buffer_names!r}"
834840
)
835841

842+
# Ensure no input/output aliasing overlap
843+
aliased_names = set(input_sources) & set(output_sources)
844+
if aliased_names:
845+
raise MLXRuntimeError(
846+
f"Input and output buffer names must be mutually disjoint (aliasing is not supported): "
847+
f"aliased={sorted(list(aliased_names))}"
848+
)
849+
836850
input_names = tuple(f"inp{i}" for i in range(len(input_sources)))
837851
output_names = tuple(f"out{i}" for i in range(len(output_sources)))
838852
rename: dict[str, str] = {}

cppmega_mlx/nn/_tilelang/dsa_splitk_indexer_loss.py

Lines changed: 294 additions & 198 deletions
Large diffs are not rendered by default.

0 commit comments

Comments
 (0)