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
lint
  • Loading branch information
oraluben committed May 21, 2026
commit 8e8e83b6d02d235f59528e7e79491dac906e53a6
15 changes: 8 additions & 7 deletions src/backend/metal/op/copy.cc
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,8 @@ Stmt LowerSIMDGroupCopy(const CopyNode &op, const LowerArgs &T,
<< "simdgroup buffer size must be multiple of 64 (8x8), got "
<< total_elements;

TVM_FFI_ICHECK(op.src_range.size() == 2) << "Expected 2D source for simdgroup store";
TVM_FFI_ICHECK(op.src_range.size() == 2)
<< "Expected 2D source for simdgroup store";
TVM_FFI_ICHECK(op.dst_range.size() == 2)
<< "Expected 2D destination for simdgroup store";
PrimExpr dst_row_base = op.dst_range[0]->min;
Expand Down Expand Up @@ -112,12 +113,12 @@ Stmt LowerSIMDGroupCopy(const CopyNode &op, const LowerArgs &T,
dst_col_base + warp_n * (warp_col_tiles * kNPerWarp) + j * kNPerWarp;
PrimExpr ptr = Call(DataType::Handle(), builtin::address_of(),
{BufferLoad(op.dst, {row, col})});
stmts.push_back(Evaluate(Call(
DataType::Handle(), builtin::simdgroup_store(),
{op.src->data, IntImm(DataType::Int(32), tile_idx), ptr, dst_stride,
IntImm(DataType::Int(32), kMPerWarp),
IntImm(DataType::Int(32), kNPerWarp),
Cast(DataType::Bool(), IntImm(DataType::Int(32), 0))})));
stmts.push_back(Evaluate(
Call(DataType::Handle(), builtin::simdgroup_store(),
{op.src->data, IntImm(DataType::Int(32), tile_idx), ptr,
dst_stride, IntImm(DataType::Int(32), kMPerWarp),
IntImm(DataType::Int(32), kNPerWarp),
Cast(DataType::Bool(), IntImm(DataType::Int(32), 0))})));
}
}
if (stmts.size() == 1) {
Expand Down
3 changes: 2 additions & 1 deletion src/backend/metal/op/fill.cc
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,8 @@ struct Fill {
int region_elements = 1;
for (auto r : op.region) {
auto imm = r->extent.as<IntImmNode>();
TVM_FFI_ICHECK(imm) << "simdgroup fill region must have constant extents";
TVM_FFI_ICHECK(imm)
<< "simdgroup fill region must have constant extents";
region_elements *= imm->value;
}
TVM_FFI_ICHECK(region_elements % 64 == 0)
Expand Down
5 changes: 2 additions & 3 deletions src/backend/metal/op/gemm.cc
Original file line number Diff line number Diff line change
Expand Up @@ -24,9 +24,8 @@ namespace {

constexpr const char *kMetalSIMDGroup = "metal.simdgroup";

std::pair<int, int>
ComputeMetalWarpPartition(const GemmWarpPolicyNode &policy, int M, int N,
int num_warps) {
std::pair<int, int> ComputeMetalWarpPartition(const GemmWarpPolicyNode &policy,
int M, int N, int num_warps) {
int m_warp = 1, n_warp = 1;
constexpr int kMPerWarp = 8;
constexpr int kNPerWarp = 8;
Expand Down
17 changes: 9 additions & 8 deletions src/runtime/metal/metal_module.h
Original file line number Diff line number Diff line change
Expand Up @@ -14,25 +14,26 @@
namespace tvm {
namespace codegen {

inline ffi::Module MetalModuleCreate(ffi::Map<ffi::String, ffi::Bytes> smap,
ffi::Map<ffi::String, runtime::FunctionInfo> fmap,
ffi::String fmt, ffi::String source) {
inline ffi::Module
MetalModuleCreate(ffi::Map<ffi::String, ffi::Bytes> smap,
ffi::Map<ffi::String, runtime::FunctionInfo> fmap,
ffi::String fmt, ffi::String source) {
auto fcreate = ffi::Function::GetGlobal("ffi.Module.create.metal");
if (fcreate.has_value()) {
return (*fcreate)(smap, fmt, fmap,
ffi::Map<ffi::String, ffi::String>{{"metal", source}})
ffi::Map<ffi::String, ffi::String>{{"metal", source}})
.cast<ffi::Module>();
}
auto fallback = ffi::Function::GetGlobal("ffi.Module.create.metal_fallback");
if (fallback.has_value()) {
return (*fallback)(smap, fmt, fmap,
ffi::Map<ffi::String, ffi::String>{{"metal", source}})
ffi::Map<ffi::String, ffi::String>{{"metal", source}})
.cast<ffi::Module>();
}
LOG(FATAL) << "Metal module factory not available.";
}

} // namespace codegen
} // namespace tvm
} // namespace codegen
} // namespace tvm

#endif // TVM_RUNTIME_METAL_METAL_MODULE_H_
#endif // TVM_RUNTIME_METAL_METAL_MODULE_H_
17 changes: 10 additions & 7 deletions src/target/codegen_metal.cc
Original file line number Diff line number Diff line change
Expand Up @@ -334,7 +334,8 @@ void CodeGenTileLangMetal::VisitStmt_(const AllocBufferNode *op) {
size_t constant_size = 1;
for (const auto &dim : op->buffer->shape) {
const IntImmNode *dim_imm = dim.as<IntImmNode>();
TVM_FFI_ICHECK(dim_imm) << "Can only handle constant size stack allocation for now";
TVM_FFI_ICHECK(dim_imm)
<< "Can only handle constant size stack allocation for now";
constant_size *= dim_imm->value;
}
TVM_FFI_ICHECK_GT(constant_size, 0)
Expand All @@ -345,12 +346,13 @@ void CodeGenTileLangMetal::VisitStmt_(const AllocBufferNode *op) {
alloc_storage_scope_[op->buffer->data.get()] = scope;
if (scope == "metal.simdgroup") {
TVM_FFI_ICHECK(dtype == DataType::Float(16) ||
dtype == DataType::Float(32) ||
dtype == DataType::BFloat(16))
dtype == DataType::Float(32) ||
dtype == DataType::BFloat(16))
<< "Only float16, float32, and bfloat16 are supported, but got "
<< dtype;
TVM_FFI_ICHECK(constant_size % 64 == 0) << "Only 8x8 matrix is supported, but got "
<< constant_size << " bytes\n";
TVM_FFI_ICHECK(constant_size % 64 == 0)
<< "Only 8x8 matrix is supported, but got " << constant_size
<< " bytes\n";

std::ostringstream dtype_os;
PrintType(dtype, dtype_os);
Expand Down Expand Up @@ -404,7 +406,8 @@ void CodeGenTileLangMetal::VisitExpr_(const CallNode *op,
<< "but expression " << ffi::GetRef<Call>(op) << " calls PrimFunc "
<< op->op;
auto f_check_simdgroup_shape = [](PrimExpr col, PrimExpr row) {
TVM_FFI_ICHECK(col->IsInstance<IntImmNode>() && row->IsInstance<IntImmNode>())
TVM_FFI_ICHECK(col->IsInstance<IntImmNode>() &&
row->IsInstance<IntImmNode>())
<< "Only constant shape is supported for simdgroup matrix, but got "
<< col << "x" << row;
int col_val = col.as<IntImmNode>()->value;
Expand Down Expand Up @@ -516,7 +519,7 @@ ffi::Module BuildTileLangMetal(IRModule mod, Target target) {
}

return MetalModuleCreate(std::move(smap), ExtractFuncInfo(mod),
ffi::String(fmt), ffi::String(source_maker.str()));
ffi::String(fmt), ffi::String(source_maker.str()));
}

TVM_FFI_STATIC_INIT_BLOCK() {
Expand Down
2 changes: 1 addition & 1 deletion src/target/codegen_metal.h
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@ class CodeGenTileLangMetal final : public CodeGenC {
void PrintVecElemStore(const std::string &vec, DataType t, int i,
const std::string &value) final;
// overload visitor
void VisitStmt_(const AllocBufferNode *op) final; // NOLINT(*)
void VisitStmt_(const AllocBufferNode *op) final; // NOLINT(*)
void VisitExpr_(const SelectNode *op, std::ostream &os) final; // NOLINT(*)
void VisitExpr_(const BroadcastNode *op, std::ostream &os) final; // NOLINT(*)
void VisitExpr_(const CallNode *op, std::ostream &os) final; // NOLINT(*)
Expand Down
4 changes: 1 addition & 3 deletions tilelang/tileop/gemm/gemm_metal.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,9 +24,7 @@ def lower(
self, layout_map: dict, target: Target, thread_bounds: Range, thread_var: tir.Var, mbar_phase_expr: tir.PrimExpr | None = None
):
thread_nums = thread_bounds.extent
m_warp, n_warp = self.policy.compute_warp_partition(
self.M, self.N, thread_nums, target, GEMM_INST_METAL
)
m_warp, n_warp = self.policy.compute_warp_partition(self.M, self.N, thread_nums, target, GEMM_INST_METAL)
warp_row_tiles = int(self.M // m_warp)
warp_col_tiles = int(self.N // n_warp)
Comment on lines +28 to +29

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

int() cast on potentially symbolic self.M / self.N.

self.M and self.N originate from buffer shapes and may be tir.IntImm or symbolic PrimExpr. If symbolic, int(self.M // m_warp) will raise at runtime. Other GEMM backends (e.g., GemmMMA) handle this similarly, so this is likely fine for Metal's concrete-size use case, but worth noting.

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

In `@tilelang/tileop/gemm/gemm_metal.py` around lines 22 - 23, The int() cast on
potentially symbolic shapes self.M and self.N will fail at runtime for PrimExpr;
update the computation of warp_row_tiles and warp_col_tiles (currently
int(self.M // m_warp) and int(self.N // n_warp)) to preserve symbolic
expressions instead of forcing Python ints—either remove the int() and keep
self.M // m_warp and self.N // n_warp, or use tir.floordiv/tvm.tir.floordiv to
produce a PrimExpr; alternatively, if a concrete int is required, guard with an
isinstance check for tir.IntImm before casting. Ensure you change both
warp_row_tiles and warp_col_tiles and keep references to m_warp and n_warp.


Expand Down
2 changes: 1 addition & 1 deletion tilelang/transform/metal_fragment_to_simdgroup.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from tvm import tirx as tir
from tvm import IRModule
from tvm.ir import Op, PointerType
from tvm.tirx import SBlock, AllocBuffer
from tvm.tirx import SBlock
from tvm.tirx.transform import prim_func_pass

_GEMM_OPS = None
Expand Down