Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
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
1 change: 1 addition & 0 deletions tests/cpp/operator/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ add_executable(test_operator
test_qdq.cu
test_cast_mxfp8.cu
test_cast_mxfp8_grouped.cu
test_cast_mxfp8_grouped_swiglu.cu
test_cast_nvfp4_transpose.cu
test_cast_float8blockwise.cu
test_cast_float8blockwise_grouped.cu
Expand Down
449 changes: 449 additions & 0 deletions tests/cpp/operator/test_cast_mxfp8_grouped_swiglu.cu

Large diffs are not rendered by default.

51 changes: 51 additions & 0 deletions tests/pytorch/test_grouped_tensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -637,6 +637,57 @@ def test_group_quantize_precomputed_offsets(self, output_dbias: bool) -> None:
assert torch.equal(grouped_output.rowwise_data, expected_output.rowwise_data)
assert torch.equal(grouped_output.scale_inv, expected_output.scale_inv)

@pytest.mark.parametrize("optimize_for_gemm", [False, True])
@pytest.mark.skipif(not mxfp8_available, reason=reason_for_no_mxfp8)
def test_group_swiglu_quantize_shapes(self, optimize_for_gemm: bool) -> None:
"""Test the grouped weighted-SwiGLU MXFP8 recompute binding plumbs shapes/dtypes.

Numerics live in tests/cpp/operator/test_cast_mxfp8_grouped_swiglu.cu; this only
covers the pybind layer: a [T, 2F] input plus a [T] prob must come back as a
columnwise-MXFP8 [T, F] grouped output.
"""
num_tensors = 3
last_dim = 256
split_sizes_list = [128, 384, 512]
total_tokens = sum(split_sizes_list)

input_2f = torch.randn(total_tokens, 2 * last_dim, dtype=torch.bfloat16, device="cuda")
# prob rides in the model dtype, matching TE's cuDNN fc1_prob_tensor convention.
prob = torch.rand(total_tokens, dtype=torch.bfloat16, device="cuda")
first_dims = torch.tensor(split_sizes_list, dtype=torch.int64, device="cuda")

quantizer = MXFP8Quantizer(fp8_dtype=tex.DType.kFloat8E4M3)
quantizer.set_usage(rowwise=False, columnwise=True)
quantizer.optimize_for_gemm = optimize_for_gemm

grouped_output = tex.group_swiglu_quantize(
input_2f, prob, quantizer, num_tensors, first_dims
)

outputs = grouped_output.split_into_quantized_tensors()
assert len(outputs) == num_tensors
for rows, output in zip(split_sizes_list, outputs):
assert output.shape == (rows, last_dim)
assert output._columnwise_data.numel() == rows * last_dim
# One e8m0 exponent per 32-row block of every column. Both the compact and the
# GEMM-swizzled layout need the same number of scales.
assert output._columnwise_scale_inv.numel() == (rows // 32) * last_dim

with pytest.raises(RuntimeError):
tex.group_swiglu_quantize(input_2f, prob.float(), quantizer, num_tensors, first_dims)

# Both operands reach the kernel as raw pointers over a densely packed range, so a
# strided view must be rejected instead of being read as if it were contiguous.
wide = torch.randn(total_tokens, 4 * last_dim, dtype=torch.bfloat16, device="cuda")
with pytest.raises(RuntimeError):
tex.group_swiglu_quantize(
wide[:, : 2 * last_dim], prob, quantizer, num_tensors, first_dims
)

strided_prob = torch.rand(2 * total_tokens, dtype=torch.bfloat16, device="cuda")[::2]
with pytest.raises(RuntimeError):
tex.group_swiglu_quantize(input_2f, strided_prob, quantizer, num_tensors, first_dims)

@pytest.mark.skipif(not mxfp8_available, reason=reason_for_no_mxfp8)
def test_bgrad_group_quantize_zero_size_tensor(self) -> None:
"""Test bgrad_group_quantize handles zero-row input without error."""
Expand Down
9 changes: 9 additions & 0 deletions transformer_engine/common/activation/swiglu_grouped.cu
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,15 @@ void nvte_group_silu(const NVTEGroupedTensor input, NVTEGroupedTensor output, cu
stream);
}

void nvte_group_swiglu_quantize(const NVTEGroupedTensor input, const NVTETensor prob,
NVTEGroupedTensor output, cudaStream_t stream) {
NVTE_API_CALL(nvte_group_swiglu_quantize);
using namespace transformer_engine;
// Weighted-SwiGLU recompute: (silu(act) * gate) * prob -> columnwise MXFP8.
dispatch::group_swiglu_quantize_fwd_helper<Empty, silu<fp32, fp32>>(input, prob, output, nullptr,
stream);
}

void nvte_group_dsilu(const NVTEGroupedTensor grad, const NVTEGroupedTensor input,
NVTEGroupedTensor output, cudaStream_t stream) {
NVTE_API_CALL(nvte_group_dsilu);
Expand Down
41 changes: 41 additions & 0 deletions transformer_engine/common/cast/dispatch/quantize.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
#include "../fp8/quantize_fp8.cuh"
#include "../fp8_blockwise/group_quantize_fp8_blockwise.cuh"
#include "../mxfp8/group_quantize_mxfp8.cuh"
#include "../mxfp8/group_swiglu_quantize_mxfp8.cuh"
#include "../mxfp8/quantize_mxfp8.cuh"
#include "../nvfp4/group_quantize_transpose_nvfp4.cuh"
#include "../nvfp4/quantize_4over6_nvfp4.cuh"
Expand Down Expand Up @@ -498,6 +499,46 @@ void group_quantize_fwd_helper(const NVTEGroupedTensor input, NVTEGroupedTensor
}
}

// Grouped weighted-SwiGLU recompute: input [T, 2F] ([act|gate]) + prob [T]
// -> columnwise MXFP8 of (silu(act) * gate) * prob.
template <typename ParamOP, float (*OP)(float, const ParamOP &)>
void group_swiglu_quantize_fwd_helper(const NVTEGroupedTensor input, const NVTETensor prob,
NVTEGroupedTensor output,
const NVTEQuantizationConfig quant_config,
cudaStream_t stream) {
using namespace detail;

NVTEScalingMode scaling_mode = nvte_grouped_tensor_scaling_mode(output);

const GroupedTensor *input_tensor = convertNVTEGroupedTensorCheck(input);
GroupedTensor *output_tensor = convertNVTEGroupedTensorCheck(output);
const Tensor *prob_tensor = convertNVTETensorCheck(prob);

// Quantization config
QuantizationConfig quant_config_cpp;
if (quant_config != nullptr) {
quant_config_cpp = *reinterpret_cast<QuantizationConfig *>(quant_config);
}

// Noop flag (graph-safe skip)
Tensor dummy_tensor;
Tensor *noop_tensor = &dummy_tensor;
if (quant_config_cpp.noop_tensor != nullptr) {
noop_tensor = convertNVTETensorCheck(quant_config_cpp.noop_tensor);
}

switch (scaling_mode) {
case NVTE_MXFP8_1D_SCALING: {
mxfp8::group_swiglu_quantize<ParamOP, OP>(input_tensor, prob_tensor, noop_tensor,
output_tensor, &quant_config_cpp, stream);
break;
}
default:
NVTE_ERROR("group_swiglu_quantize only supports NVTE_MXFP8_1D_SCALING, got: " +
to_string(scaling_mode) + ".");
}
}

template <bool IS_DBIAS, bool IS_DACT, typename ParamOP, float (*OP)(float, const ParamOP &)>
void group_quantize_bwd_helper(const NVTEGroupedTensor grad, const NVTEGroupedTensor input,
NVTEGroupedTensor output, NVTEGroupedTensor dbias,
Expand Down
Loading