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
fix: tir->tirx migration build fixes for metal backend
- Replace tir::Buffer with Buffer, tir:: with tirx:: in metal backend
- Update codegen_metal to use AllocBufferNode instead of AllocateNode
- Create runtime/metal/metal_module.h with MetalModuleCreate factory
- Fix ICHECK -> TVM_FFI_ICHECK macros
- Fix 'from tvm import tir' -> 'from tvm import tirx as tir' in Python
- Restore tilelang/backend/__init__.py for metal gemm registration
- Add metal backend import to tilelang/__init__.py
- Fix metal_fragment_to_simdgroup.py: Block->SBlock, Allocate->AllocBuffer
  • Loading branch information
oraluben committed May 21, 2026
commit 0eb2865ee99dc14cfb923f8c4c5e3fedc3c037c4
3 changes: 0 additions & 3 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -27,9 +27,6 @@ classifiers = [
]
dynamic = ["version"]
dependencies = [
# >=0.1.6 fixes a memory issue: tilelang#1502, but keep
# requirement as wide as possible to be compatible with other libraries
# pip will try to use latest version whenever possible.
"apache-tvm-ffi~=0.1.0,>=0.1.10",
# torch-c-dlpack-ext provides prebuilt torch extensions.
# Without it, TVM FFI may require JIT compilation on first import.
Expand Down
22 changes: 11 additions & 11 deletions src/backend/metal/op/copy.cc
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
#include "op/utils.h"
#include "target/utils.h"

#include <tvm/tir/builtin.h>
#include <tvm/tirx/builtin.h>

#include <algorithm>
#include <cmath>
Expand All @@ -32,36 +32,36 @@ bool CheckSIMDGroupCopy(const CopyNode &op) {
Stmt LowerSIMDGroupCopy(const CopyNode &op, const LowerArgs &T,
arith::Analyzer *analyzer) {
(void)analyzer;
ICHECK(IsSIMDGroupBuffer(op.src));
TVM_FFI_ICHECK(IsSIMDGroupBuffer(op.src));

int total_elements = 1;
for (auto s : op.src->shape) {
auto imm = s.as<IntImmNode>();
ICHECK(imm) << "simdgroup buffer must have constant shape";
TVM_FFI_ICHECK(imm) << "simdgroup buffer must have constant shape";
total_elements *= imm->value;
}
ICHECK(total_elements % 64 == 0)
TVM_FFI_ICHECK(total_elements % 64 == 0)
<< "simdgroup buffer size must be multiple of 64 (8x8), got "
<< total_elements;

ICHECK(op.src_range.size() == 2) << "Expected 2D source for simdgroup store";
ICHECK(op.dst_range.size() == 2)
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;
PrimExpr dst_col_base = op.dst_range[1]->min;
PrimExpr dst_stride = op.dst->shape[op.dst->shape.size() - 1];

int warp_size = TargetGetWarpSize(T.target);
const auto *block_size_imm = T.thread_bounds->extent.as<IntImmNode>();
ICHECK(block_size_imm)
TVM_FFI_ICHECK(block_size_imm)
<< "simdgroup copy requires constant thread bounds";
int block_size = block_size_imm->value;
int num_warps = block_size / warp_size;
PrimExpr warp_id = FloorDiv(T.thread_var, warp_size);

const auto *m_imm = op.src_range[0]->extent.as<IntImmNode>();
const auto *n_imm = op.src_range[1]->extent.as<IntImmNode>();
ICHECK(m_imm && n_imm) << "simdgroup copy requires constant extents";
TVM_FFI_ICHECK(m_imm && n_imm) << "simdgroup copy requires constant extents";
int M = m_imm->value;
int N = n_imm->value;

Expand Down Expand Up @@ -90,13 +90,13 @@ Stmt LowerSIMDGroupCopy(const CopyNode &op, const LowerArgs &T,
}
}

ICHECK(M >= m_warp * kMPerWarp && N >= n_warp * kNPerWarp)
TVM_FFI_ICHECK(M >= m_warp * kMPerWarp && N >= n_warp * kNPerWarp)
<< "Cannot partition " << M << "x" << N << " matrix across " << m_warp
<< "x" << n_warp << " warps with 8x8 simdgroup tiles";
int warp_row_tiles = M / m_warp / kMPerWarp;
int warp_col_tiles = N / n_warp / kNPerWarp;
ICHECK(warp_row_tiles > 0 && warp_col_tiles > 0);
ICHECK(warp_row_tiles * warp_col_tiles * 64 <= total_elements)
TVM_FFI_ICHECK(warp_row_tiles > 0 && warp_col_tiles > 0);
TVM_FFI_ICHECK(warp_row_tiles * warp_col_tiles * 64 <= total_elements)
<< "Warp partition produces more tiles than buffer capacity";

PrimExpr warp_m = FloorMod(warp_id, m_warp);
Expand Down
8 changes: 4 additions & 4 deletions src/backend/metal/op/fill.cc
Original file line number Diff line number Diff line change
Expand Up @@ -12,12 +12,12 @@
#include "transform/loop_partition.h"
#include "transform/loop_vectorize.h"

#include <tvm/tir/builtin.h>
#include <tvm/tirx/builtin.h>

namespace tvm {
namespace tl {

using namespace tir;
using namespace tirx;

namespace metal {

Expand All @@ -28,10 +28,10 @@ struct Fill {
int region_elements = 1;
for (auto r : op.region) {
auto imm = r->extent.as<IntImmNode>();
ICHECK(imm) << "simdgroup fill region must have constant extents";
TVM_FFI_ICHECK(imm) << "simdgroup fill region must have constant extents";
region_elements *= imm->value;
}
ICHECK(region_elements % 64 == 0)
TVM_FFI_ICHECK(region_elements % 64 == 0)
<< "simdgroup buffer size must be multiple of 64 (8x8), got "
<< region_elements;

Expand Down
14 changes: 8 additions & 6 deletions src/backend/metal/op/gemm.cc
Original file line number Diff line number Diff line change
Expand Up @@ -7,14 +7,16 @@

#include "target/utils.h"

#include <tvm/runtime/logging.h>

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

namespace tvm {
namespace tl {

using namespace tir;
using namespace tirx;

namespace metal {

Expand All @@ -29,9 +31,9 @@ ComputeMetalWarpPartition(const GemmWarpPolicyNode &policy, int M, int N,
constexpr int kMPerWarp = 8;
constexpr int kNPerWarp = 8;

ICHECK(M % kMPerWarp == 0)
TVM_FFI_ICHECK(M % kMPerWarp == 0)
<< "M must be divisible by " << kMPerWarp << ", but got " << M;
ICHECK(N % kNPerWarp == 0)
TVM_FFI_ICHECK(N % kNPerWarp == 0)
<< "N must be divisible by " << kNPerWarp << ", but got " << N;

if (policy.isFullRow()) {
Expand Down Expand Up @@ -86,10 +88,10 @@ ComputeMetalWarpPartition(const GemmWarpPolicyNode &policy, int M, int N,
m_warp = best_m;
n_warp = best_n;
} else {
ICHECK(0) << "Unknown GemmWarpPolicy";
TVM_FFI_ICHECK(0) << "Unknown GemmWarpPolicy";
}

ICHECK(m_warp * n_warp == num_warps)
TVM_FFI_ICHECK(m_warp * n_warp == num_warps)
<< "m_warp * n_warp must equal num_warps, m_warp: " << m_warp
<< ", n_warp: " << n_warp << ", num_warps: " << num_warps;
policy.m_warp = m_warp;
Expand All @@ -113,7 +115,7 @@ struct Gemm {
static std::pair<int, int>
ComputeWarpPartition(const GemmWarpPolicyNode &policy, int M, int N,
int block_size, Target target, String gemm_inst) {
ICHECK(gemm_inst == kMetalSIMDGroup)
TVM_FFI_ICHECK(gemm_inst == kMetalSIMDGroup)
<< "Unsupported Metal GEMM instruction: " << gemm_inst;
int num_warps = block_size / TargetGetWarpSize(target);
return ComputeMetalWarpPartition(policy, M, N, num_warps);
Expand Down
4 changes: 2 additions & 2 deletions src/backend/metal/op/utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -12,11 +12,11 @@ namespace tvm {
namespace tl {
namespace metal {

inline bool IsSIMDGroupBuffer(const tir::Buffer &buffer) {
inline bool IsSIMDGroupBuffer(const Buffer &buffer) {
return buffer.defined() && buffer.scope() == "metal.simdgroup";
}

inline bool IsRegisterBuffer(const tir::Buffer &buffer) {
inline bool IsRegisterBuffer(const Buffer &buffer) {
return IsFragmentBuffer(buffer) || IsSIMDGroupBuffer(buffer);
}

Expand Down
38 changes: 38 additions & 0 deletions src/runtime/metal/metal_module.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
#ifndef TVM_RUNTIME_METAL_METAL_MODULE_H_
#define TVM_RUNTIME_METAL_METAL_MODULE_H_

#include <tvm/ffi/container/map.h>
#include <tvm/ffi/extra/module.h>
#include <tvm/ffi/function.h>
#include <tvm/ir/module.h>
#include <tvm/runtime/logging.h>
#include <tvm/target/codegen.h>

#include <string>
#include <utility>

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) {
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}})
.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}})
.cast<ffi::Module>();
}
LOG(FATAL) << "Metal module factory not available.";
}

} // namespace codegen
} // namespace tvm

#endif // TVM_RUNTIME_METAL_METAL_MODULE_H_
Loading
Loading