Skip to content
Open
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
60 changes: 59 additions & 1 deletion bench/gemm/gemm_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -1825,7 +1825,65 @@ def compute(self, x, wq, w_scale, w_zp):

@property
def supported_accelerators(self) -> set[Accelerator]:
return {Accelerator.AMD_GFX942}
return {Accelerator.AMD_GFX942, Accelerator.AMD_GFX950}


@register_gemm_op
class TritonBF16Int4Shuffled(TritonBF16Int4Rowwise):
"""ROCm Triton BF16xINT4 shuffled GEMM (routes to rowwise on AMD)."""

@property
def supported_accelerators(self) -> set[Accelerator]:
return {Accelerator.AMD_GFX942, Accelerator.AMD_GFX950}


@register_gemm_op
class TritonBF16Int4GroupedShuffled(CutlassFP8Int4Rowwise):
Comment thread
apicciau marked this conversation as resolved.
Outdated
"""ROCm Triton BF16xINT4 grouped shuffled GEMM."""

def preprocess(self, x, w):
assert isinstance(x, list) and isinstance(w, list)
m_values = [i.shape[0] for i in x]
m_sizes = torch.tensor(m_values).to(dtype=torch.int64, device=x[0].device)
wq_list, scale_list, zero_list = [], [], []
for wi in w:
wq_i, s_i, z_i = int4_row_quantize_zp(wi)
wq_list.append(pack_int4(wq_i))
scale_list.append(s_i)
zero_list.append(z_i)
wq = torch.stack(wq_list, dim=0).contiguous()
group_scale = torch.stack(scale_list, dim=0).contiguous()
group_zero = torch.stack(zero_list, dim=0).contiguous()
x = torch.concat(x, dim=0).contiguous()
return x, wq, group_scale, group_zero, m_sizes

def quantize(self, x, wq, group_scale, group_zero, m_sizes):
return x, wq, group_scale, group_zero, m_sizes

def compute(self, x, wq, group_scale, group_zero, m_sizes):
from mslk.gemm.triton.int4_grouped_gemm import matmul_bf16i4_rowwise_grouped

return matmul_bf16i4_rowwise_grouped(x, wq, group_scale, group_zero, m_sizes)

@property
def supported_accelerators(self) -> set[Accelerator]:
return {Accelerator.AMD_GFX942, Accelerator.AMD_GFX950}

@property
def supported_gemm_types(self) -> set[GemmType]:
return {GemmType.GROUPED}

@property
def compute_dtype(self) -> ComputeDtype:
return ComputeDtype.BF16

@property
def input_bytes_per_element(self) -> float:
return 2.0

@property
def weight_bytes_per_element(self) -> float:
return 0.5


@register_gemm_op
Expand Down
33 changes: 15 additions & 18 deletions csrc/gemm/gemm_ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,21 @@ TORCH_LIBRARY_FRAGMENT(mslk, m) {
// the Triton implementation registered by fp8_groupwise_grouped_gemm.py.
m.def(
"f8f8bf16_groupwise_grouped(Tensor XQ, Tensor WQ, Tensor x_scale, Tensor w_scale, Tensor M_sizes) -> Tensor");
// BF16xINT4 GEMMs: shared schema between CUDA and ROCm. On CUDA the
// implementations are CUTLASS-based; on ROCm they are registered by the
// Triton modules (int4_gemm.py, int4_grouped_gemm.py) via torch.library.impl
// at Python import time.
m.def(
"bf16i4bf16_rowwise(Tensor X, Tensor W, Tensor w_scale_group, Tensor w_zero_group) -> Tensor");
m.def(
"bf16i4bf16_rowwise_batched(Tensor X, Tensor WQ, Tensor w_scale, Tensor w_zp) -> Tensor");
m.def(
"bf16i4bf16_shuffled(Tensor X, Tensor W, Tensor w_scale_group, Tensor w_zero_group) -> Tensor");
m.def(
"bf16i4bf16_shuffled_grouped(Tensor X, Tensor WQ, Tensor w_scale_group, Tensor w_zero_group, Tensor M_sizes) -> Tensor");
m.def(
"bf16i4bf16_shuffled_batched(Tensor X, Tensor WQ, Tensor w_scale, Tensor w_zp) -> Tensor");
m.def("preshuffle_i4(Tensor WQ, Tensor w_scale) -> (Tensor, Tensor)");
#ifdef USE_ROCM
m.def(
"f8f8f16_rowwise(Tensor XQ, Tensor WQ, Tensor x_scale, Tensor w_scale, Tensor? bias=None, bool use_fast_accum=True) -> Tensor");
Expand All @@ -73,13 +88,6 @@ TORCH_LIBRARY_FRAGMENT(mslk, m) {
// Generic PyTorch grouped GEMM API is only available on AMD for now.
m.def(
"f8f8bf16_rowwise_grouped_mm(Tensor XQ, Tensor WQ, Tensor x_scale, Tensor w_scale, Tensor? offsets, Tensor(a!) output) -> Tensor");
// BF16xINT4 rowwise GEMMs: schema only on ROCm; implementations are
// registered by mslk.gemm.triton.int4_gemm via torch.library.impl at
// Python import time.
m.def(
"bf16i4bf16_rowwise(Tensor X, Tensor W, Tensor w_scale_group, Tensor w_zero_group) -> Tensor");
m.def(
"bf16i4bf16_rowwise_batched(Tensor X, Tensor WQ, Tensor w_scale, Tensor w_zp) -> Tensor");
// INT8 GEMM via Triton — static and dynamic scale variants.
m.def("i8i8bf16(Tensor XQ, Tensor WQ, float scale, int split_k=1) -> Tensor");
m.def(
Expand All @@ -103,21 +111,10 @@ TORCH_LIBRARY_FRAGMENT(mslk, m) {
"f8i4bf16_rowwise(Tensor XQ, Tensor WQ, Tensor x_scale, Tensor w_scale, Tensor w_zp) -> Tensor");
m.def(
"f8i4bf16_shuffled(Tensor XQ, Tensor WQ, Tensor x_scale, Tensor w_scale, Tensor w_scale_group) -> Tensor");
m.def(
"bf16i4bf16_shuffled(Tensor X, Tensor W, Tensor w_scale_group, Tensor w_zero_group) -> Tensor");
m.def(
"f8i4bf16_shuffled_grouped(Tensor XQ, Tensor WQ, Tensor x_scale, Tensor w_scale, Tensor w_scale_group, Tensor M_sizes) -> Tensor");
m.def(
"bf16i4bf16_shuffled_grouped(Tensor X, Tensor WQ, Tensor w_scale_group, Tensor w_zero_group, Tensor M_sizes) -> Tensor");
m.def(
"bf16i4bf16_rowwise(Tensor X, Tensor W, Tensor w_scale_group, Tensor w_zero_group) -> Tensor");
m.def(
"bf16i4bf16_shuffled_batched(Tensor X, Tensor WQ, Tensor w_scale, Tensor w_zp) -> Tensor");
m.def(
"bf16i4bf16_rowwise_batched(Tensor X, Tensor WQ, Tensor w_scale, Tensor w_zp) -> Tensor");
m.def(
"i8i8bf16_dynamic(Tensor XQ, Tensor WQ, Tensor scale, int split_k=1) -> Tensor");
m.def("preshuffle_i4(Tensor WQ, Tensor w_scale) -> (Tensor, Tensor)");
#endif
}

Expand Down
10 changes: 6 additions & 4 deletions mslk/gemm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,14 +29,16 @@
from . import _meta # noqa: F401, E402

if torch.version.hip is not None:
# Register the Triton ROCm implementations for mx8mx4bf16, mx8mx8bf16, and
# f8f8bf16_groupwise(_grouped). Each import triggers a
# @torch.library.impl(..., "CUDA") decoration that overrides the default
# (non-existent) CUDA impl so the ops dispatch to the Triton kernels on AMD.
# Register Triton implementations for ROCm. Each import triggers the
# @torch.library.impl("mslk::...", "CUDA") decoration in the respective
# module, which overrides the default (non-existent) CUDA impl so that
# torch.ops.mslk.* dispatches to the Triton kernel on AMD.
from .triton import ( # noqa: F401
fp8_groupwise_gemm,
fp8_groupwise_grouped_gemm,
grouped_gemm as _grouped_gemm,
int4_grouped_gemm as _int4_grouped_gemm,
int4_grouped_gemm_fused as _int4_grouped_gemm_fused,
mx8mx4_gemm,
mx8mx8_gemm,
)
15 changes: 15 additions & 0 deletions mslk/gemm/_meta.py
Original file line number Diff line number Diff line change
Expand Up @@ -448,6 +448,21 @@ def bf16i4bf16_shuffled_batched_meta(
return torch.empty((B, M, N), dtype=torch.bfloat16, device=X.device)


if hasattr(torch.ops.mslk, "bf16i4bf16_shuffled_grouped"):

@torch.library.register_fake("mslk::bf16i4bf16_shuffled_grouped")
def bf16i4bf16_shuffled_grouped_meta(
X: torch.Tensor,
WQ: torch.Tensor,
w_scale_group: torch.Tensor,
w_zero_group: torch.Tensor,
M_sizes: torch.Tensor,
) -> torch.Tensor:
M_total = X.shape[0]
N = WQ.shape[1]
return torch.empty((M_total, N), dtype=torch.bfloat16, device=X.device)


if hasattr(torch.ops.mslk, "bf16i4bf16_rowwise_batched"):

@torch.library.register_fake("mslk::bf16i4bf16_rowwise_batched")
Expand Down
Loading
Loading