Skip to content
Merged
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
4 changes: 2 additions & 2 deletions transformer_engine/common/amd_detail/hip_float8.h
Original file line number Diff line number Diff line change
Expand Up @@ -49,8 +49,8 @@ enum class hip_f8_rounding_mode {
// => bias = 15 for 152, 7 for 143
// => NAN/INF are represented as per IEEE conventions

static __device__ bool hip_f8_bias_mode_bit_device;
static bool hip_f8_bias_mode_bit_host;
static __device__ bool hip_f8_bias_mode_bit_device = true;
static bool hip_f8_bias_mode_bit_host = true;

static __global__ void set_hip_f8_bias_mode_bit(bool v) {
hip_f8_bias_mode_bit_device = v;
Expand Down
157 changes: 102 additions & 55 deletions transformer_engine/common/gemm/cublaslt_gemm.cu
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,17 @@
* See LICENSE for license information.
************************************************************************/

#include <type_traits>
#include <transformer_engine/transformer_engine.h>
#include <transformer_engine/logging.h>
#include <transformer_engine/gemm.h>
#ifndef __HIP_PLATFORM_HCC__
#include <cublasLt.h>
#endif
#include <cublas_v2.h>
#else
#define ROCBLAS_BETA_FEATURES_API
#include <rocblas/rocblas.h>
#endif
#include "../common.h"
#include "../util/vectorized_pointwise.h"
#ifdef __HIP_PLATFORM_HCC__
Expand Down Expand Up @@ -43,10 +47,15 @@ namespace transformer_engine {
} \
break; \
case DType::kBFloat16: \
{ \
using type = bf16; \
{__VA_ARGS__} \
} \
break; \
case DType::kFloat8E5M2: \
case DType::kFloat8E4M3: \
{ \
NVTE_ERROR("Bfloat16 and FP8 type not instantiated"); \
NVTE_ERROR("FP8 type not instantiated"); \
} \
break; \
default: \
Expand Down Expand Up @@ -77,18 +86,20 @@ float gelu_forward(float x)
return x * cdf;
}


template <typename Tin, typename T>
__global__
void gelu_forward_kernel(const Tin* in, T* out, int m, int n) {
for(int id = blockIdx.x * blockDim.x + threadIdx.x; id < m * n; id += blockDim.x * gridDim.x)
{
Tin x = (Tin)(__ldg(&in[id]));
Tin x = in[id];
float y = gelu_forward((float)x);
out[id] = (T)(y);
}

}


template <typename Tin, typename T>
void gelu_forward_kernelLauncher(const Tin* in, T* out, int m, int n, hipStream_t stream) {
int blocks_per_row = ceil(float(n)/1024);
Expand Down Expand Up @@ -124,8 +135,8 @@ __global__
void gelu_backward_kernel(const Tin* dy, T* out, const T* __restrict pre_gelu_out, int m, int n) {
for(int id = blockIdx.x * blockDim.x + threadIdx.x; id < m * n; id += blockDim.x * gridDim.x)
{
Tin x = (Tin)(__ldg(&pre_gelu_out[id]));
Tin dx = gelu_backward((float)x, (float)dy[id]);
Tin x = (Tin)pre_gelu_out[id];
Tin dx = (Tin)gelu_backward((float)x, (float)dy[id]);
out[id] = (T)(dx);
}
}
Expand All @@ -138,47 +149,47 @@ void gelu_backward_kernelLauncher(const Tin* in, T* out, const T* pre_gelu_out,
hipLaunchKernelGGL(( gelu_backward_kernel<Tin, T>), dim3(grid), dim3(block), 0, stream, in, out, pre_gelu_out, m, n);
}

template <typename Tin, typename T>
template <typename Tin, typename T, typename Tb>
__global__
void add_bias_kernel(const Tin* in, T* out, const T* __restrict bias, int m, int n)
void add_bias_kernel(const Tin* in, T* out, const Tb* __restrict bias, int m, int n)
{
for(int id = blockIdx.x * blockDim.x + threadIdx.x; id < m * n; id += blockDim.x * gridDim.x)
{
Tin reg_bias = (Tin)(__ldg(&bias[id % n]));
Tin reg_bias = (Tin)bias[id % n];
Tin val = in[id] + reg_bias;
out[id] = (T)(val);
}
}


template <typename Tin, typename T>
void add_bias_kernelLauncher(const Tin* in, T* out, const T* __restrict bias, int m, int n, hipStream_t stream) {
template <typename Tin, typename T, typename Tb>
void add_bias_kernelLauncher(const Tin* in, T* out, const Tb* __restrict bias, int m, int n, hipStream_t stream) {
dim3 block, grid;
block.x = 1024;
grid.x = ceil(m * n / 1024.);
hipLaunchKernelGGL(( add_bias_kernel<Tin, T>), dim3(grid), dim3(block), 0, stream, in, out, bias, m, n);
hipLaunchKernelGGL(( add_bias_kernel<Tin, T, Tb>), dim3(grid), dim3(block), 0, stream, in, out, bias, m, n);

}

template <typename Tin, typename T>
template <typename Tin, typename T, typename Tb>
__global__
void add_bias_gelu_kernel(const Tin* in, T* out, T* pre_gelu_out, const T* __restrict bias, int m, int n)
void add_bias_gelu_kernel(const Tin* in, T* out, T* pre_gelu_out, const Tb* __restrict bias, int m, int n)
{
for(int id = blockIdx.x * blockDim.x + threadIdx.x; id < m * n; id += blockDim.x * gridDim.x)
{
Tin reg_bias = (Tin)(__ldg(&bias[id % n]));
Tin reg_bias = (Tin)bias[id % n];
Tin val = in[id] + reg_bias;
pre_gelu_out[id] = (T)(val);
out[id] = (T)(gelu_forward(val));
}
}

template <typename Tin, typename T>
void add_bias_gelu_kernelLauncher(const Tin* in, T* out, T* pre_gelu_out, const T* __restrict bias, int m, int n, hipStream_t stream) {
template <typename Tin, typename T, typename Tb>
void add_bias_gelu_kernelLauncher(const Tin* in, T* out, T* pre_gelu_out, const Tb* __restrict bias, int m, int n, hipStream_t stream) {
dim3 block, grid;
block.x = 1024;
grid.x = ceil(m * n / 1024.);
hipLaunchKernelGGL(( add_bias_gelu_kernel<Tin, T>), dim3(grid), dim3(block), 0, stream, in, out, pre_gelu_out, bias, m, n );
hipLaunchKernelGGL(( add_bias_gelu_kernel<Tin, T, Tb>), dim3(grid), dim3(block), 0, stream, in, out, pre_gelu_out, bias, m, n );

}

Expand Down Expand Up @@ -278,12 +289,10 @@ transformer_engine::DType get_transformer_engine_dtype(const rocblas_datatype t)
return DType::kFloat32;
case rocblas_datatype_bf16_r:
return DType::kBFloat16;
/* HIP-TODO: Add back after installing rocblas with FP8 types
case DType::kFloat8E4M3:
return CUDA_R_8F_E4M3;
case DType::kFloat8E5M2:
return CUDA_R_8F_E5M2;
*/
case rocblas_datatype_f8_r:
return DType::kFloat8E4M3;
case rocblas_datatype_bf8_r:
return DType::kFloat8E5M2;
default:
NVTE_ERROR("Invalid type");
}
Expand Down Expand Up @@ -326,24 +335,38 @@ void cublas_gemm(void* A,
float zero = 0.0;
float beta = (accumulate) ? one : zero;

float alpha = 1.0;
if (use_fp8) {
float A_scale_inv, B_scale_inv;
hipMemcpy(&A_scale_inv, A_scale_inverse, sizeof(float), hipMemcpyDeviceToHost);
hipMemcpy(&B_scale_inv, B_scale_inverse, sizeof(float), hipMemcpyDeviceToHost);
alpha = A_scale_inv * B_scale_inv;
}

rocblas_handle handle;
NVTE_CHECK_CUBLAS(rocblas_create_handle(&handle));
rocblas_datatype computeType = rocblas_datatype_f32_r;

int64_t ld_gelumat = (int64_t) ldd;


// We don't deal with fp8 for now
NVTE_CHECK(!use_fp8, "fp8 gemm is unavailable right now!");

NVTE_CHECK((A_type==rocblas_datatype_f16_r && B_type==rocblas_datatype_f16_r && D_type==rocblas_datatype_f16_r) ||
(A_type==rocblas_datatype_f32_r && B_type==rocblas_datatype_f32_r && D_type==rocblas_datatype_f32_r),
"Only fp32 and fp16 GEMMs are available now!");


(A_type==rocblas_datatype_f32_r && B_type==rocblas_datatype_f32_r && D_type==rocblas_datatype_f32_r) ||
(A_type==rocblas_datatype_f8_r && B_type==rocblas_datatype_f8_r && D_type==rocblas_datatype_f32_r) ||
(A_type==rocblas_datatype_f8_r && B_type==rocblas_datatype_bf8_r && D_type==rocblas_datatype_f32_r) ||
(A_type==rocblas_datatype_bf8_r && B_type==rocblas_datatype_f8_r && D_type==rocblas_datatype_f32_r),
/*
(A_type==rocblas_datatype_f8_r && B_type==rocblas_datatype_f8_r && D_type==rocblas_datatype_f8_r) ||
(A_type==rocblas_datatype_f8_r && B_type==rocblas_datatype_bf8_r && D_type==rocblas_datatype_bf8_r) ||
(A_type==rocblas_datatype_bf8_r && B_type==rocblas_datatype_f8_r && D_type==rocblas_datatype_bf8_r),
*/
//Currently does not support output of fp8 tensors
"Only the following combinations of data types are enabled now!\n 1. input: fp16, output: fp16.\n \
2. input: fp32, output: fp32.\n 3. input: fp8, output: fp32");


//If D is not fp32, then we need a temp buffer for GEMM result before applying epilogues. Otherwise, we can apply epilogues in-place.

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.

May want to consider to refactor this long logic expression (to a branch of subconditions) for better avg. performance

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Not sure what you mean by a branch of subconditions.

void* D_temp;
if ((bias || gelu) && (A_type==rocblas_datatype_f16_r && B_type==rocblas_datatype_f16_r && D_type==rocblas_datatype_f16_r)) {
if ((bias || gelu) && (D_type==rocblas_datatype_f16_r || D_type==rocblas_datatype_f8_r || D_type==rocblas_datatype_bf8_r)) {

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.

lines above NVTE_CHECK made sure D_type is not f8/bf8 (for the moment), so here the later 2 conditions seem to be stale, but it is minor.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yeah, I add these two conditions in advance as there would be support for fp8 output in the future. Right now, it is limited to fp32/fp16 output.

NVTE_CHECK_CUDA( hipMalloc(&D_temp, sizeof(float)*m*n) );
}
else {
Expand All @@ -356,14 +379,28 @@ void cublas_gemm(void* A,
D_temp_type = rocblas_datatype_f16_r;
}



// D = alpha * (A * B) + beta * C
// TODO: Can we search for rocblas_gemm_algo??
NVTE_CHECK_CUBLAS(rocblas_gemm_ex(handle, transa, transb, m, n, k, &one,
if (use_fp8) {
rocblas_computetype computeType = rocblas_compute_type_f32;

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.

Is it rocblas_compute_type_f32_r?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

No, it is rocblas_compute_type_f32

NVTE_CHECK_CUBLAS(rocblas_gemm_ex3(handle, transa, transb, m, n, k, &alpha,
A, A_type, lda,
B, B_type, ldb,
&beta, D_temp, D_temp_type, ldd, D_temp, D_temp_type, ldd,
computeType, rocblas_gemm_algo::rocblas_gemm_algo_standard,0,0));
}
else {
rocblas_datatype computeType = rocblas_datatype_f32_r;
NVTE_CHECK_CUBLAS(rocblas_gemm_ex(handle, transa, transb, m, n, k, &alpha,
A, A_type, lda,
B, B_type, ldb,
&beta, D_temp, D_temp_type, ldd, D_temp, D_temp_type, ldd,
computeType, rocblas_gemm_algo::rocblas_gemm_algo_standard,0,0));
}


NVTE_CHECK_CUBLAS(rocblas_destroy_handle(handle));

int batch_size, input_dim, output_dim;
Expand All @@ -389,6 +426,7 @@ void cublas_gemm(void* A,
output_dim = k;
DType input_dtype = get_transformer_engine_dtype(rocblas_datatype_f32_r);
DType output_dtype = get_transformer_engine_dtype(D_type);
DType bias_dtype = get_transformer_engine_dtype(bias_type);
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(input_dtype, IType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(output_dtype, OType,
detail::gelu_backward_kernelLauncher<IType, OType>(reinterpret_cast<const IType*>(D_temp),
Expand All @@ -401,7 +439,8 @@ void cublas_gemm(void* A,
);

void* bias_tmp;
if (D_type == rocblas_datatype_f16_r) {
//if (D_type == rocblas_datatype_f16_r) {
if (bias_type != rocblas_datatype_f32_r) {
NVTE_CHECK_CUDA( hipMalloc(&bias_tmp, sizeof(float)*input_dim) ); // The bias gradient is for the first linear layer
}
else {
Expand All @@ -416,8 +455,9 @@ void cublas_gemm(void* A,
0);
);

if (D_type == rocblas_datatype_f16_r) {
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(output_dtype, OType,
//if (D_type == rocblas_datatype_f16_r) {
if (bias_type != rocblas_datatype_f32_r) {
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(bias_dtype, OType,
detail::identity_kernelLauncher<float, OType>(reinterpret_cast<const float*>(bias_tmp),
reinterpret_cast<OType*>(bias_ptr),
input_dim,
Expand All @@ -437,15 +477,18 @@ void cublas_gemm(void* A,
output_dim = m;
DType input_dtype = get_transformer_engine_dtype(rocblas_datatype_f32_r);
DType output_dtype = get_transformer_engine_dtype(D_type);
DType bias_dtype = get_transformer_engine_dtype(bias_type);
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(input_dtype, IType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(output_dtype, OType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(bias_dtype, BType,
detail::add_bias_gelu_kernelLauncher<IType, OType>(reinterpret_cast<const IType*>(D_temp),
reinterpret_cast<OType*>(D),
reinterpret_cast<OType*>(pre_gelu_out),
reinterpret_cast<const OType*>(bias_ptr),
reinterpret_cast<const BType*>(bias_ptr),
batch_size,
output_dim,
0);
);
);
);
}
Expand All @@ -462,7 +505,8 @@ void cublas_gemm(void* A,
input_dim = m;
output_dim = n;
void * bias_tmp;
if (B_type == rocblas_datatype_f16_r) {
//if (B_type == rocblas_datatype_f16_r) {
if (bias_type != rocblas_datatype_f32_r) {
NVTE_CHECK_CUDA( hipMalloc(&bias_tmp, sizeof(float)*output_dim) );
}
else {
Expand All @@ -471,15 +515,16 @@ void cublas_gemm(void* A,

DType input_dtype = get_transformer_engine_dtype(B_type);
DType output_dtype = get_transformer_engine_dtype(D_type);
DType bias_dtype = get_transformer_engine_dtype(bias_type);
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(input_dtype, IType,
detail::bias_gradient_kernelLauncher<IType>(reinterpret_cast<const IType*>(B),
reinterpret_cast<float*>(bias_tmp),
batch_size,
output_dim,
0);
);
if (B_type == rocblas_datatype_f16_r) {
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(output_dtype, OType,
if (bias_type != rocblas_datatype_f32_r) {
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(bias_dtype, OType,
detail::identity_kernelLauncher<float, OType>(reinterpret_cast<const float*>(bias_tmp),
reinterpret_cast<OType*>(bias_ptr),
output_dim,
Expand All @@ -506,14 +551,17 @@ void cublas_gemm(void* A,
output_dim = m;
DType input_dtype = get_transformer_engine_dtype(rocblas_datatype_f32_r);
DType output_dtype = get_transformer_engine_dtype(D_type);
DType bias_dtype = get_transformer_engine_dtype(bias_type);
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(input_dtype, IType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(output_dtype, OType,
detail::add_bias_kernelLauncher<IType, OType>(reinterpret_cast<const IType*>(D_temp),
TRANSFORMER_ENGINE_TYPE_SWITCH_ROCM_SIM(bias_dtype, BType,
detail::add_bias_kernelLauncher<IType, OType, BType>(reinterpret_cast<const IType*>(D_temp),
reinterpret_cast<OType*>(D),
reinterpret_cast<const OType*>(bias_ptr),
reinterpret_cast<const BType*>(bias_ptr),
batch_size,
output_dim,
0);
);
);
);
}
Expand Down Expand Up @@ -567,16 +615,9 @@ void cublas_gemm(void* A,
);
}
}
if ((bias || gelu) && (A_type==rocblas_datatype_f16_r && B_type==rocblas_datatype_f16_r && D_type==rocblas_datatype_f16_r)) {
if ((bias || gelu) && (D_type==rocblas_datatype_f16_r || D_type==rocblas_datatype_f8_r || D_type==rocblas_datatype_bf8_r)) {

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.

same, D_type should not be f8/bf8 from earlier check, minor.

NVTE_CHECK_CUDA( hipFree(D_temp) );
}
/*
NVTE_CHECK_CUBLAS(cublasLtMatmulPreferenceDestroy(preference));
NVTE_CHECK_CUBLAS(cublasLtMatrixLayoutDestroy(Ddesc));
NVTE_CHECK_CUBLAS(cublasLtMatrixLayoutDestroy(Bdesc));
NVTE_CHECK_CUBLAS(cublasLtMatrixLayoutDestroy(Adesc));
NVTE_CHECK_CUBLAS(cublasLtMatmulDescDestroy(operationDesc));
*/
}
#else
void cublas_gemm(void* A,
Expand Down Expand Up @@ -767,12 +808,10 @@ rocblas_datatype get_cuda_dtype(const transformer_engine::DType t) {
return rocblas_datatype_f32_r;
case DType::kBFloat16:
return rocblas_datatype_bf16_r;
/* HIP-TODO: Add back after installing rocblas with FP8 types
case DType::kFloat8E4M3:
return CUDA_R_8F_E4M3;
return rocblas_datatype_f8_r;
case DType::kFloat8E5M2:
return CUDA_R_8F_E5M2;
*/
return rocblas_datatype_bf8_r;
default:
NVTE_ERROR("Invalid type");
}
Expand Down Expand Up @@ -853,18 +892,26 @@ void nvte_cublas_gemm(const NVTETensor A,
nvte_log_gemm_config = true;
}

if (nvte_log_gemm_config)
if (nvte_log_gemm_config) {
float A_scale_inv, B_scale_inv;
hipMemcpy(&A_scale_inv, Ainvscale->dptr, sizeof(float), hipMemcpyDeviceToHost);
hipMemcpy(&B_scale_inv, Binvscale->dptr, sizeof(float), hipMemcpyDeviceToHost);
std::cout << "m=" << m << " k=" << k << " n=" << n
<< " transa=" << (transa?"T":"N")
<< " transb=" << (transb?"T":"N")
<< " A_type=" << (int)inputA->dtype
<< " B_type=" << (int)inputB->dtype
<< " D_type=" << (int)outputD->dtype
<< " bias_type=" << (int)biasTensor->dtype
<< " grad=" << grad
<< " bias=" << (biasTensor->dptr != nullptr)
<< " gelu=" << (outputGelu->dptr != nullptr)
<< " use_fp8=" << ( is_fp8_dtype(inputA->dtype) || is_fp8_dtype(inputB->dtype) )
<< " A_scale_inverse = " << A_scale_inv
<< " B_scale_inverse = " << B_scale_inv
<< " accumulate=" << accumulate
<< std::endl;
}

cublas_gemm(inputA->dptr, Ainvscale->dptr,
inputB->dptr, Binvscale->dptr,
Expand Down
Loading