Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
Prev Previous commit
Next Next commit
[Metal] Fix review findings: dedup warp partition, kernel_only, barri…
…er comment, immutable gemm ops

- Extract shared ComputeSquareWarpPartition to src/backend/metal/op/utils.h,
  used by both gemm.cc (Square policy) and copy.cc (simdgroup store lowering).
- Implement kernel_only parameter in MetalKernelAdapter.get_kernel_source.
- Add comment documenting that Metal's simdgroup_barrier synchronizes the
  full threadgroup (no per-simdgroup barrier exists).
- Replace module-level mutable global in metal_fragment_to_simdgroup.py
  with @lru_cache and remove CUDA-only gemm ops from the set.
  • Loading branch information
oraluben committed May 21, 2026
commit 7f7adfde43d259e4b28e288bfb4c5aef554bf6c0
28 changes: 2 additions & 26 deletions src/backend/metal/op/copy.cc
Original file line number Diff line number Diff line change
Expand Up @@ -11,10 +11,6 @@

#include <tvm/tirx/builtin.h>

#include <algorithm>
#include <cmath>
#include <limits>

namespace tvm {
namespace tl {

Expand Down Expand Up @@ -68,28 +64,8 @@ Stmt LowerSIMDGroupCopy(const CopyNode &op, const LowerArgs &T,

constexpr int kMPerWarp = 8;
constexpr int kNPerWarp = 8;
int m_warp = 1, n_warp = num_warps;
int max_m = M / kMPerWarp;
int max_n = N / kNPerWarp;
float ideal = N > 0 ? static_cast<float>(M) / N : 1.f;
float best_score = std::numeric_limits<float>::max();
for (int m = 1; m <= std::min(num_warps, max_m); ++m) {
if (num_warps % m != 0) {
continue;
}
int n = num_warps / m;
if (n > max_n) {
continue;
}
float m_per = static_cast<float>(M) / (m * kMPerWarp);
float n_per = static_cast<float>(N) / (n * kNPerWarp);
float score = std::abs(m_per / n_per - ideal);
if (score < best_score) {
best_score = score;
m_warp = m;
n_warp = n;
}
}
auto [m_warp, n_warp] =
ComputeSquareWarpPartition(num_warps, M, N, kMPerWarp, kNPerWarp);

TVM_FFI_ICHECK(M >= m_warp * kMPerWarp && N >= n_warp * kNPerWarp)
<< "Cannot partition " << M << "x" << N << " matrix across " << m_warp
Expand Down
33 changes: 3 additions & 30 deletions src/backend/metal/op/gemm.cc
Original file line number Diff line number Diff line change
Expand Up @@ -5,12 +5,11 @@

#include "op/gemm.h"

#include "backend/metal/op/utils.h"
#include "target/utils.h"

#include <tvm/runtime/logging.h>

#include <cmath>
#include <limits>
#include <utility>

namespace tvm {
Expand Down Expand Up @@ -58,34 +57,8 @@ std::pair<int, int> ComputeMetalWarpPartition(const GemmWarpPolicyNode &policy,
}
}
} else if (policy.isSquare()) {
int max_m_warps = M / kMPerWarp;
float ideal_ratio = N > 0 ? static_cast<float>(M) / N : 1.0f;

int best_m = 1;
int best_n = 1;
float best_balance = std::numeric_limits<float>::max();
for (int m = 1; m <= max_m_warps && m <= num_warps; m++) {
int n = num_warps / m;

float m_per_warp = static_cast<float>(M) / (m * kMPerWarp);
float n_per_warp = static_cast<float>(N) / (n * kNPerWarp);
if (m_per_warp < 1 || n_per_warp < 1) {
continue;
}
if (m * n != num_warps) {
continue;
}

float balance = std::abs(m_per_warp / n_per_warp - ideal_ratio);
if (balance < best_balance) {
best_balance = balance;
best_m = m;
best_n = n;
}
}

m_warp = best_m;
n_warp = best_n;
std::tie(m_warp, n_warp) =
ComputeSquareWarpPartition(num_warps, M, N, kMPerWarp, kNPerWarp);
} else {
TVM_FFI_ICHECK(0) << "Unknown GemmWarpPolicy";
}
Expand Down
32 changes: 32 additions & 0 deletions src/backend/metal/op/utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,10 @@
#ifndef TVM_TL_BACKEND_METAL_OP_UTILS_H_
#define TVM_TL_BACKEND_METAL_OP_UTILS_H_

#include <cmath>
#include <limits>
#include <utility>

#include "op/utils.h"

namespace tvm {
Expand All @@ -20,6 +24,34 @@ inline bool IsRegisterBuffer(const Buffer &buffer) {
return IsFragmentBuffer(buffer) || IsSIMDGroupBuffer(buffer);
}

inline std::pair<int, int> ComputeSquareWarpPartition(int num_warps, int M,
int N, int kMPerWarp,
int kNPerWarp) {
int max_m = M / kMPerWarp;
int max_n = N / kNPerWarp;
float ideal_ratio = N > 0 ? static_cast<float>(M) / N : 1.0f;

int best_m = 1, best_n = 1;
float best_balance = std::numeric_limits<float>::max();
for (int m = 1; m <= std::min(num_warps, max_m); ++m) {
if (num_warps % m != 0)
continue;
int n = num_warps / m;
if (n > max_n)
continue;

float m_per = static_cast<float>(M) / (m * kMPerWarp);
float n_per = static_cast<float>(N) / (n * kNPerWarp);
float balance = std::abs(m_per / n_per - ideal_ratio);
if (balance < best_balance) {
best_balance = balance;
best_m = m;
best_n = n;
}
}
return {best_m, best_n};
}

} // namespace metal
} // namespace tl
} // namespace tvm
Expand Down
4 changes: 4 additions & 0 deletions src/target/codegen_metal.cc
Original file line number Diff line number Diff line change
Expand Up @@ -282,6 +282,10 @@ void CodeGenTileLangMetal::PrintType(DataType t,
void CodeGenTileLangMetal::PrintStorageSync(const CallNode *op) {
const std::string &sync = op->args[0].as<StringImmNode>()->value;
if (sync == "warp") {
// Metal has no per-simdgroup barrier; simdgroup_barrier synchronizes
// the entire threadgroup (same as threadgroup_barrier). We emit the
// narrower intrinsic so the source documents the intended scope even
// though the hardware effect is identical.
this->PrintIndent();
this->stream << "simdgroup_barrier(mem_flags::mem_threadgroup);\n";
} else if (sync == "shared") {
Expand Down
6 changes: 6 additions & 0 deletions tilelang/jit/adapter/torch/metal.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,12 @@ def __init__(
_kernel = None

def get_kernel_source(self, kernel_only: bool = True) -> str:
if kernel_only:
# Return just the kernel function body, stripping Metal
# module-level boilerplate (includes, structs, etc.).
idx = self.kernel_global_source.find("kernel void ")
if idx >= 0:
return self.kernel_global_source[idx:]
return self.kernel_global_source

Comment on lines +56 to +64

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.

⚠️ Potential issue | 🟡 Minor

Fix the return type annotation and honour the kernel_only flag.

Two issues here:

  1. Return type mismatch — kernel_global_source is declared str | None (Line 23), so the method can return None, contradicting the -> str annotation. This will cause silent type errors for callers.

  2. Unused kernel_only parameter — Every peer adapter branches on this flag (base.py, nvrtc/adapter.py, cython/adapter.py). Silently ignoring it here means get_kernel_source(kernel_only=False) behaves identically to kernel_only=True, breaking the expected contract.

🛠️ Proposed fix
-    def get_kernel_source(self, kernel_only: bool = True) -> str:
-        return self.kernel_global_source
+    def get_kernel_source(self, kernel_only: bool = True) -> str | None:
+        # Metal has a single unified source; kernel_only has no distinct meaning here.
+        return self.kernel_global_source

If a non-None guarantee is truly required at call sites, add an explicit assertion:

-    def get_kernel_source(self, kernel_only: bool = True) -> str:
-        return self.kernel_global_source
+    def get_kernel_source(self, kernel_only: bool = True) -> str | None:
+        assert self.kernel_global_source is not None, "kernel_global_source is not available"
+        return self.kernel_global_source
🧰 Tools
🪛 Ruff (0.15.1)

[warning] 56-56: Unused method argument: kernel_only

(ARG002)

🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.

In `@tilelang/jit/adapter/torch/metal.py` around lines 56 - 58, The method
get_kernel_source currently claims to return str but may return None and ignores
the kernel_only flag; change its signature to -> str | None (or keep -> str but
assert/raise if kernel_global_source is None) and implement the kernel_only
branch: if kernel_only is True return self.kernel_global_source, otherwise
return the full Metal source (compose or return the attribute that holds the
complete module/source such as self.metal_source or self.full_source); ensure
you reference get_kernel_source and kernel_global_source and either assert
kernel_global_source is not None before returning a str or update callers/types
to accept Optional[str].

def _convert_torch_func(self) -> Callable:
Expand Down
16 changes: 6 additions & 10 deletions tilelang/transform/metal_fragment_to_simdgroup.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,24 +7,20 @@

from __future__ import annotations

from functools import lru_cache

from tvm import tirx as tir
from tvm import IRModule
from tvm.ir import Op, PointerType
from tvm.tirx import SBlock
from tvm.tirx.transform import prim_func_pass

_GEMM_OPS = None


@lru_cache(maxsize=1)
def _get_gemm_ops():
global _GEMM_OPS
if _GEMM_OPS is None:
_GEMM_OPS = {
Op.get("tl.tileop.gemm"),
Op.get("tl.tileop.wgmma_gemm"),
Op.get("tl.tileop.tcgen05_gemm"),
}
return _GEMM_OPS
return frozenset({
Op.get("tl.tileop.gemm"),
})


def _extract_buffer_var_from_region(region_call):
Expand Down