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
10 changes: 8 additions & 2 deletions transformer_engine/common/cast/dispatch/gated.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,10 @@ void quantize_gated_fwd_helper(const NVTETensor nvte_input, NVTETensor nvte_outp

switch (output->scaling_mode) {
case NVTE_DELAYED_TENSOR_SCALING: {
const bool use_tma_kernels = (cols % 32 == 0) && is_supported_by_CC_100();
const bool use_tma_kernels =
(cols % 32 == 0) && is_supported_by_CC_100() &&
fp8::cast_gated_tma_fits_device</*IS_BWD=*/false, ParamOP, ActOP, nullptr>(
input.dtype(), output->dtype());
if (use_tma_kernels) {
Tensor dummy_grad_tensor;
fp8::cast_gated_tma</*IS_BWD=*/false, ParamOP, ActOP, nullptr>(input, dummy_grad_tensor,
Expand Down Expand Up @@ -137,7 +140,10 @@ void quantize_gated_bwd_helper(const NVTETensor nvte_grad, const NVTETensor nvte

switch (output->scaling_mode) {
case NVTE_DELAYED_TENSOR_SCALING: {
const bool use_tma_kernels = (cols % 32 == 0) && is_supported_by_CC_100();
const bool use_tma_kernels =
(cols % 32 == 0) && is_supported_by_CC_100() &&
fp8::cast_gated_tma_fits_device</*IS_BWD=*/true, ParamOP, ActOP, DActOP>(
gated_input.dtype(), output->dtype());
if (use_tma_kernels) {
fp8::cast_gated_tma</*IS_BWD=*/true, ParamOP, ActOP, DActOP>(gated_input, grad, output, p,
stream);
Expand Down
55 changes: 42 additions & 13 deletions transformer_engine/common/cast/fp8/gated_fp8.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
#include <transformer_engine/transformer_engine.h>

#include "../../common.h"
#include "../../util/cuda_runtime.h"
#include "../../util/math.h"
#include "../../util/ptx.cuh"
#include "../../util/vectorized_pointwise.h"
Expand Down Expand Up @@ -279,6 +280,45 @@ __global__ void __launch_bounds__(THREADS_PER_CHUNK)
}
} // namespace kernel

// Dynamic shared memory requested by cast_fp8_gated_kernel. Must match the
// buffer layout inside the kernel.
inline size_t cast_gated_tma_dynamic_shmem_size(const bool is_bwd, const DType itype,
const DType otype) {
using namespace kernel;
const size_t buff_elems_total = BUFFERS_NUM * SHMEM_DIM_Y * SHMEM_DIM_X;
const size_t buff_size_aligned_in =
DIVUP_TO_MULTIPLE(buff_elems_total * typeToNumBits(itype) / 8, TMA_SHMEM_ALIGNMENT);
const size_t buff_size_aligned_out =
DIVUP_TO_MULTIPLE(buff_elems_total * typeToNumBits(otype) / 8, TMA_SHMEM_ALIGNMENT);
const size_t grad_mem = (is_bwd ? buff_size_aligned_in : 0);
const size_t in_act_mem = buff_size_aligned_in;
const size_t in_gate_mem = buff_size_aligned_in;
const size_t out_act_mem = buff_size_aligned_out;
const size_t out_gate_mem = buff_size_aligned_out;
return grad_mem + (in_act_mem + in_gate_mem) + (out_act_mem + out_gate_mem) + TMA_SHMEM_ALIGNMENT;
}

// Whether cast_fp8_gated_kernel fits in the shared memory of the current device.
// Devices with compute capability 10.0+ differ in shared memory per block
// (e.g. SM 12.0 has much less than SM 10.0), so FP32 configurations may not fit.
// The static shared memory is read from the compiled kernel so that __shared__
// arrays in the device functions it calls (e.g. reduce_max) are included.
template <bool IS_BWD, typename ParamOP, float (*ActOP)(float, const ParamOP &),
float (*DActOP)(float, const ParamOP &)>
bool cast_gated_tma_fits_device(const DType itype, const DType otype) {
using namespace kernel;
size_t static_shmem_size = 0;
TRANSFORMER_ENGINE_TYPE_SWITCH_INPUT(
itype, IType,
TRANSFORMER_ENGINE_TYPE_SWITCH_OUTPUT(
otype, OType,
static_shmem_size = cuda::static_shared_memory_size(reinterpret_cast<const void *>(
&cast_fp8_gated_kernel<IS_BWD, ParamOP, ActOP, DActOP, IType, OType>));););
const size_t required =
cast_gated_tma_dynamic_shmem_size(IS_BWD, itype, otype) + static_shmem_size;
return required <= cuda::max_shared_memory_per_block_optin();
}

template <bool IS_BWD, typename ParamOP, float (*ActOP)(float, const ParamOP &),
float (*DActOP)(float, const ParamOP &)>
void cast_gated_tma(const Tensor &gated_input, const Tensor &grad, Tensor *output, ParamOP &p,
Expand Down Expand Up @@ -329,19 +369,8 @@ void cast_gated_tma(const Tensor &gated_input, const Tensor &grad, Tensor *outpu
SHMEM_DIM_X, tensor_stride_elems, cols,
typeToNumBits(output->dtype()));

const size_t buff_elems_total = BUFFERS_NUM * SHMEM_DIM_Y * SHMEM_DIM_X;
const size_t buff_size_aligned_in =
DIVUP_TO_MULTIPLE(buff_elems_total * sizeof(IType), TMA_SHMEM_ALIGNMENT);
const size_t buff_size_aligned_out =
DIVUP_TO_MULTIPLE(buff_elems_total * sizeof(OType), TMA_SHMEM_ALIGNMENT);
const size_t grad_mem = (IS_BWD ? buff_size_aligned_in : 0);
const size_t in_act_mem = buff_size_aligned_in;
const size_t in_gate_mem = buff_size_aligned_in;
const size_t out_act_mem = buff_size_aligned_out;
const size_t out_gate_mem = buff_size_aligned_out;

const size_t shmem_size = grad_mem + (in_act_mem + in_gate_mem) +
(out_act_mem + out_gate_mem) + TMA_SHMEM_ALIGNMENT;
const size_t shmem_size =
cast_gated_tma_dynamic_shmem_size(IS_BWD, gated_input.dtype(), output->dtype());

auto kernel = cast_fp8_gated_kernel<IS_BWD, ParamOP, ActOP, DActOP, IType, OType>;
NVTE_CHECK_CUDA(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
Expand Down
26 changes: 8 additions & 18 deletions transformer_engine/common/swizzle/swizzle.cu
Original file line number Diff line number Diff line change
Expand Up @@ -24,20 +24,6 @@ namespace {
constexpr int MXFP8_BLOCK_SIZE = 32;
constexpr int NVFP4_BLOCK_SIZE = 16;

int get_max_dynamic_smem(int device_id = -1) {
static std::vector<int> cache(cuda::num_devices(), -1);
static std::vector<std::once_flag> flags(cuda::num_devices());
if (device_id < 0) {
device_id = cuda::current_device();
}
auto init = [&]() {
NVTE_CHECK_CUDA(cudaDeviceGetAttribute(&cache[device_id],
cudaDevAttrMaxSharedMemoryPerBlockOptin, device_id));
};
std::call_once(flags[device_id], init);
return cache[device_id];
}

constexpr __device__ __host__ int TB_DIM = 32;
constexpr __device__ __host__ int NEW_SF_TILE_DIM_K = 16;
constexpr __device__ __host__ int N_SF_PER_TD_PER_TILE = 4;
Expand Down Expand Up @@ -1109,7 +1095,8 @@ void swizzle_scaling_factors(const Tensor* input, Tensor* output, cudaStream_t s

const int narrow_k_slm_size =
TB_DIM * num_tiles_k * SF_TILE_DIM_M * SF_TILE_DIM_K * static_cast<int>(sizeof(int8_t));
if (num_tiles_k < TB_DIM && narrow_k_slm_size <= get_max_dynamic_smem()) {
if (num_tiles_k < TB_DIM &&
static_cast<size_t>(narrow_k_slm_size) <= cuda::max_shared_memory_per_block_optin()) {
// Narrow-K: batch TB_DIM M-tiles per block, fully utilizing all threads.
dim3 num_blocks_narrow(DIVUP(num_tiles_m, TB_DIM));
NVTE_CHECK_CUDA(
Expand Down Expand Up @@ -1166,7 +1153,8 @@ void swizzle_scaling_factors(const Tensor* input, Tensor* output, cudaStream_t s

const int narrow_m_slm_size =
TB_DIM * num_tiles_m * SF_TILE_DIM_M * SF_TILE_DIM_K * static_cast<int>(sizeof(int8_t));
if (num_tiles_m < TB_DIM && narrow_m_slm_size <= get_max_dynamic_smem()) {
if (num_tiles_m < TB_DIM &&
static_cast<size_t>(narrow_m_slm_size) <= cuda::max_shared_memory_per_block_optin()) {
// Narrow-M: batch TB_DIM K-tiles per block, fully utilizing all threads.
dim3 num_blocks_narrow(DIVUP(num_tiles_k, TB_DIM));
NVTE_CHECK_CUDA(
Expand Down Expand Up @@ -1534,7 +1522,8 @@ void multi_tensor_swizzle_scaling_factors(const std::vector<Tensor*>& input,
const int narrow_k_slm =
TB_DIM * num_tiles_k * SF_TILE_DIM_M * SF_TILE_DIM_K * static_cast<int>(sizeof(int8_t));
all_narrow_k =
all_narrow_k && (num_tiles_k < TB_DIM) && (narrow_k_slm <= get_max_dynamic_smem());
all_narrow_k && (num_tiles_k < TB_DIM) &&
(static_cast<size_t>(narrow_k_slm) <= cuda::max_shared_memory_per_block_optin());
int vec_load_size_i = (num_tiles_k - 1) % 4 + 1;
// We use the minimum vec_load_size across all tensors.
// TODO(zhongbo): fix vec_load_size for NVFP4
Expand Down Expand Up @@ -1606,7 +1595,8 @@ void multi_tensor_swizzle_scaling_factors(const std::vector<Tensor*>& input,
const int narrow_m_slm =
TB_DIM * num_tiles_m * SF_TILE_DIM_M * SF_TILE_DIM_K * static_cast<int>(sizeof(int8_t));
all_narrow_m =
all_narrow_m && (num_tiles_m < TB_DIM) && (narrow_m_slm <= get_max_dynamic_smem());
all_narrow_m && (num_tiles_m < TB_DIM) &&
(static_cast<size_t>(narrow_m_slm) <= cuda::max_shared_memory_per_block_optin());
int vec_load_size_i = (num_tiles_k - 1) % 4 + 1;
// We use the minimum vec_load_size across all tensors.
vec_load_size = std::min(vec_load_size, vec_load_size_i);
Expand Down
33 changes: 33 additions & 0 deletions transformer_engine/common/util/cuda_runtime.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,9 @@

#include <filesystem>
#include <fstream>
#include <map>
#include <mutex>
#include <utility>

#include "../common.h"
#include "../util/cuda_driver.h"
Expand Down Expand Up @@ -143,6 +145,37 @@ int sm_count(int device_id) {
return cache[device_id];
}

size_t max_shared_memory_per_block_optin(int device_id) {
Comment thread
ravimajeti marked this conversation as resolved.
static std::vector<size_t> cache(num_devices(), 0);
static std::vector<std::once_flag> flags(num_devices());
if (device_id < 0) {
device_id = current_device();
}
NVTE_CHECK(0 <= device_id && device_id < num_devices(), "invalid CUDA device ID");
auto init = [&]() {
int value;
NVTE_CHECK_CUDA(
cudaDeviceGetAttribute(&value, cudaDevAttrMaxSharedMemoryPerBlockOptin, device_id));
cache[device_id] = static_cast<size_t>(value);
};
std::call_once(flags[device_id], init);
return cache[device_id];
}

size_t static_shared_memory_size(const void *kernel) {
static std::map<std::pair<const void *, int>, size_t> cache;
static std::mutex mutex;
const auto key = std::make_pair(kernel, current_device());
const std::lock_guard<std::mutex> lock(mutex);
auto it = cache.find(key);
if (it == cache.end()) {
cudaFuncAttributes attr;
NVTE_CHECK_CUDA(cudaFuncGetAttributes(&attr, kernel));
it = cache.emplace(key, attr.sharedSizeBytes).first;
}
return it->second;
}

void stream_priority_range(int *low_priority, int *high_priority, int device_id) {
static std::vector<std::pair<int, int>> cache(num_devices());
static std::vector<std::once_flag> flags(num_devices());
Expand Down
24 changes: 24 additions & 0 deletions transformer_engine/common/util/cuda_runtime.h
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,30 @@ int sm_arch(int device_id = -1);
*/
int sm_count(int device_id = -1);

/* \brief Maximum shared memory per block available with opt-in
*
* This is the upper bound for the static plus dynamic shared memory
* of a single thread block, after raising the kernel's
* cudaFuncAttributeMaxDynamicSharedMemorySize attribute.
*
* \param[in] device_id CUDA device (default is current device)
*
* \return Shared memory size in bytes
*/
size_t max_shared_memory_per_block_optin(int device_id = -1);

/* \brief Static shared memory used by a compiled kernel on the current device
*
* Size of the kernel's fixed-size __shared__ arrays, including those
* declared in device functions it calls, as reported by
* cudaFuncGetAttributes. The result is cached per kernel and device.
*
* \param[in] kernel Pointer to the __global__ function
*
* \return Shared memory size in bytes
*/
size_t static_shared_memory_size(const void *kernel);

/* \brief Minimum and maximum stream priorities supported on device
*
* \param[in] device_id CUDA device (default is current device)
Expand Down
Loading