Skip to content

Commit 944c612

Browse files
committed
unify with nvte_multi_tensor_quantize
1 parent 7d08697 commit 944c612

5 files changed

Lines changed: 29 additions & 44 deletions

File tree

benchmarks/cpp/cast/bench_group_quantize_mxfp8.cpp

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -303,7 +303,7 @@ static void BM_MultiQuantizeMXFP8(benchmark::State &state) {
303303

304304
for (auto _ : state) {
305305
HIP_CHECK(hipEventRecord(start, stream));
306-
nvte_multi_quantize_mxfp8(num_experts, nvte_inputs.data(), nvte_outputs.data(), stream);
306+
nvte_multi_tensor_quantize(nvte_inputs.data(), nvte_outputs.data(), nullptr, num_experts, stream);
307307
HIP_CHECK(hipEventRecord(stop, stream));
308308
HIP_CHECK(hipEventSynchronize(stop));
309309
float ms = 0;

tests/cpp/operator/test_multi_quantize_mxfp8.cu

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -47,8 +47,8 @@ void performTest(const std::vector<std::pair<size_t, size_t>> &tensor_dims,
4747
nvte_outputs_multi.push_back(t.data());
4848
}
4949

50-
nvte_multi_quantize_mxfp8(num_tensors, nvte_inputs.data(),
51-
nvte_outputs_multi.data(), 0);
50+
nvte_multi_tensor_quantize(nvte_inputs.data(), nvte_outputs_multi.data(),
51+
nullptr, num_tensors, 0);
5252

5353
for (size_t i = 0; i < num_tensors; i++) {
5454
if (tensor_dims[i].first > 0 && tensor_dims[i].second > 0)

transformer_engine/common/cast/cast.cu

Lines changed: 13 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -30,20 +30,6 @@ void nvte_quantize(const NVTETensor input, NVTETensor output, cudaStream_t strea
3030
dispatch::quantize_fwd_helper<IS_ACT, Empty, nullptr>(input, output, nullptr, stream);
3131
}
3232

33-
#ifdef __HIP_PLATFORM_AMD__
34-
void nvte_multi_quantize_mxfp8(size_t num_tensors, const NVTETensor *input_list,
35-
NVTETensor *output_list, cudaStream_t stream) {
36-
NVTE_API_CALL(nvte_multi_quantize_mxfp8);
37-
using namespace transformer_engine;
38-
std::vector<Tensor *> input_list_, output_list_;
39-
for (size_t i = 0; i < num_tensors; i++) {
40-
input_list_.push_back(convertNVTETensorCheck(input_list[i]));
41-
output_list_.push_back(convertNVTETensorCheck(output_list[i]));
42-
}
43-
dispatch::multi_quantize_mxfp8(input_list_, output_list_, stream);
44-
}
45-
#endif
46-
4733
void nvte_group_quantize(const NVTEGroupedTensor input, NVTEGroupedTensor output,
4834
cudaStream_t stream) {
4935
NVTE_API_CALL(nvte_group_quantize);
@@ -113,6 +99,19 @@ void nvte_multi_tensor_quantize(const NVTETensor *inputs, NVTETensor *outputs,
11399
NVTE_API_CALL(nvte_multi_tensor_quantize);
114100
using namespace transformer_engine;
115101

102+
#ifdef __HIP_PLATFORM_AMD__
103+
if (num_tensors > 0 &&
104+
convertNVTETensorCheck(outputs[0])->scaling_mode == NVTE_MXFP8_1D_SCALING) {
105+
std::vector<Tensor *> input_list, output_list;
106+
for (size_t i = 0; i < num_tensors; i++) {
107+
input_list.push_back(convertNVTETensorCheck(inputs[i]));
108+
output_list.push_back(convertNVTETensorCheck(outputs[i]));
109+
}
110+
dispatch::multi_quantize_mxfp8(input_list, output_list, stream);
111+
return;
112+
}
113+
#endif
114+
116115
constexpr bool IS_ACT = false;
117116

118117
const size_t num_streams = nvte_get_num_compute_streams();

transformer_engine/common/include/transformer_engine/cast.h

Lines changed: 0 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -101,19 +101,6 @@ void nvte_quantize(const NVTETensor input, NVTETensor output, cudaStream_t strea
101101
void nvte_group_quantize(const NVTEGroupedTensor input, NVTEGroupedTensor output,
102102
cudaStream_t stream);
103103

104-
#ifdef __HIP_PLATFORM_AMD__
105-
/*! \brief Fused multi-tensor MXFP8 quantize. Quantizes multiple tensors in a single kernel launch.
106-
* Each tensor can have different shapes. Output tensors are written to per-tensor pointers.
107-
*
108-
* \param[in] num_tensors Number of tensors to quantize.
109-
* \param[in] input_list Array of input tensors.
110-
* \param[in,out] output_list Array of output MXFP8 tensors.
111-
* \param[in] stream CUDA stream used for the operation.
112-
*/
113-
void nvte_multi_quantize_mxfp8(size_t num_tensors, const NVTETensor *input_list,
114-
NVTETensor *output_list, cudaStream_t stream);
115-
#endif
116-
117104
/*! \brief Casts input tensor to FP8/MXFP8/BlockwiseFP8, providing the option to immediately exit the kernel
118105
* based on the value of the 'noop' tensor.
119106
* The type of quantized tensor in the output depends on the scaling mode of the output

transformer_engine/pytorch/csrc/extensions/cast.cpp

Lines changed: 13 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -291,20 +291,6 @@ void multi_tensor_quantize_impl(const std::vector<TensorWrapper> &input_list,
291291
}
292292

293293
// Launch TE kernel
294-
#ifdef USE_ROCM
295-
if (num_tensors > 0 && detail::IsMXFP8Quantizers(quantizer_py_list[0].ptr())) {
296-
std::vector<NVTETensor> nvte_input_list, nvte_output_list;
297-
for (size_t i = 0; i < num_tensors; i++) {
298-
nvte_input_list.push_back(input_list[i].data());
299-
nvte_output_list.push_back(output_list[i].data());
300-
}
301-
NVTE_SCOPED_GIL_RELEASE({
302-
nvte_multi_quantize_mxfp8(nvte_input_list.size(), nvte_input_list.data(),
303-
nvte_output_list.data(), at::cuda::getCurrentCUDAStream());
304-
});
305-
return;
306-
}
307-
#endif
308294
if (with_fused_kernel) {
309295
// Fused kernel for multi-tensor quantize
310296
std::vector<NVTETensor> nvte_tensor_input_list;
@@ -317,6 +303,19 @@ void multi_tensor_quantize_impl(const std::vector<TensorWrapper> &input_list,
317303
nvte_multi_cast_transpose(nvte_tensor_input_list.size(), nvte_tensor_input_list.data(),
318304
nvte_tensor_output_list.data(), at::cuda::getCurrentCUDAStream());
319305
});
306+
#ifdef USE_ROCM
307+
} else if (num_tensors > 0 && detail::IsMXFP8Quantizers(quantizer_py_list[0].ptr())) {
308+
std::vector<NVTETensor> nvte_tensor_input_list;
309+
std::vector<NVTETensor> nvte_tensor_output_list;
310+
for (size_t i = 0; i < num_tensors; ++i) {
311+
nvte_tensor_input_list.push_back(input_list[i].data());
312+
nvte_tensor_output_list.push_back(output_list[i].data());
313+
}
314+
NVTE_SCOPED_GIL_RELEASE({
315+
nvte_multi_tensor_quantize(nvte_tensor_input_list.data(), nvte_tensor_output_list.data(),
316+
nullptr, num_tensors, at::cuda::getCurrentCUDAStream());
317+
});
318+
#endif
320319
} else {
321320
// Quantize kernels individually
322321
for (size_t i = 0; i < num_tensors; ++i) {

0 commit comments

Comments
 (0)