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 @@ -15,6 +15,7 @@ add_executable(test_operator
test_cast_mxfp8_grouped.cu
test_cast_mxfp8_grouped_scaled_swiglu.cu
test_cast_nvfp4_transpose.cu
test_cast_nvfp4_transpose_grouped.cu
test_cast_float8blockwise.cu
test_cast_float8blockwise_grouped.cu
test_dequantize_mxfp8.cu
Expand Down
1,055 changes: 1,055 additions & 0 deletions tests/cpp/operator/test_cast_nvfp4_transpose_grouped.cu

Large diffs are not rendered by default.

100 changes: 100 additions & 0 deletions transformer_engine/common/cast/core/grouped_layout.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -236,6 +236,17 @@ __device__ __forceinline__ size_t get_tensor_rows_num(
return 0;
}

__device__ __forceinline__ size_t
get_tensor_base_offset(const size_t tensor_id, const ShapeRepresentation shape_rep,
const size_t first_logical_dim, const size_t last_logical_dim,
const size_t num_tensors, const int64_t *const __restrict__ offsets_ptr) {
if (shape_rep == ShapeRepresentation::SAME_BOTH_DIMS) {
const size_t rows_per_tensor = first_logical_dim / num_tensors;
return tensor_id * rows_per_tensor * last_logical_dim;
}
return static_cast<size_t>(offsets_ptr[tensor_id]);
}

template <ShapeRepresentation SHAPE_REP>
__device__ __forceinline__ size_t
get_tensor_cols_num(const size_t tensor_id, const size_t last_logical_dim,
Expand Down Expand Up @@ -319,6 +330,28 @@ struct BlockDescriptor {
block_offset_X(block_offset_X_) {}
};

// Per-tensor metadata prepared together with grouped TMA descriptors.
struct TensorMetadata {
size_t rows = 0;
size_t cols = 0;
size_t tensor_base = 0;
size_t rowwise_scale_base = 0;
size_t colwise_scale_base = 0;

__host__ __device__ __forceinline__ constexpr TensorMetadata() = default;

__host__ __device__ __forceinline__ constexpr TensorMetadata(const size_t rows_,
const size_t cols_,
const size_t tensor_base_,
const size_t rowwise_scale_base_,
const size_t colwise_scale_base_)
: rows(rows_),
cols(cols_),
tensor_base(tensor_base_),
rowwise_scale_base(rowwise_scale_base_),
colwise_scale_base(colwise_scale_base_) {}
};

template <ShapeRepresentation SHAPE_REP, size_t CHUNK_DIM_Y, size_t CHUNK_DIM_X>
__device__ __forceinline__ JobDescriptor decode_job(
const size_t num_tensors, const size_t first_logical_dim, const size_t last_logical_dim,
Expand All @@ -340,6 +373,28 @@ __device__ __forceinline__ JobDescriptor decode_job(
return JobDescriptor(block_id, block_global_offset, tensor_id, rows, cols);
}

template <ShapeRepresentation SHAPE_REP, size_t CHUNK_DIM_Y, size_t CHUNK_DIM_X>
__device__ __forceinline__ JobDescriptor
decode_job(const size_t num_tensors, const size_t first_logical_dim, const size_t last_logical_dim,
const size_t work_blocks_X, const int32_t ctaid_X, const int32_t ctaid_Y,
const int64_t *const __restrict__ offsets_ptr,
const TensorMetadata *const __restrict__ metadata_ptr) {
constexpr size_t ELTS_PER_CHUNK = CHUNK_DIM_Y * CHUNK_DIM_X;
constexpr bool is_single_tensor = (SHAPE_REP == ShapeRepresentation::SAME_BOTH_DIMS ||
SHAPE_REP == ShapeRepresentation::VARYING_FIRST_DIM);
const size_t ctaid_X_u = static_cast<size_t>(ctaid_X);
const size_t ctaid_Y_u = static_cast<size_t>(ctaid_Y);
const size_t block_id = ctaid_Y_u * work_blocks_X + ctaid_X_u;
const size_t block_global_offset =
is_single_tensor ? (ctaid_Y_u * CHUNK_DIM_Y * last_logical_dim + ctaid_X_u * CHUNK_DIM_X)
: (block_id * ELTS_PER_CHUNK);
const size_t tensor_id = get_current_tensor_id<SHAPE_REP, CHUNK_DIM_Y>(
num_tensors, block_global_offset, ctaid_Y_u, first_logical_dim, last_logical_dim,
offsets_ptr);
const TensorMetadata metadata = metadata_ptr[tensor_id];
return JobDescriptor(block_id, block_global_offset, tensor_id, metadata.rows, metadata.cols);
}

template <ShapeRepresentation SHAPE_REP>
__device__ __forceinline__ bool is_job_valid(const JobDescriptor &job,
const size_t total_work_blocks,
Expand Down Expand Up @@ -374,6 +429,13 @@ __device__ __forceinline__ bool job_has_work(const JobDescriptor &job) {
return job.rows != 0 && job.cols != 0;
}

template <ShapeRepresentation SHAPE_REP>
__device__ __forceinline__ bool is_job_valid_with_work(
const JobDescriptor &job, const size_t total_work_blocks,
const int64_t *const __restrict__ offsets_ptr) {
return is_job_valid<SHAPE_REP>(job, total_work_blocks, offsets_ptr) && job_has_work(job);
}

__device__ __forceinline__ void advance_to_next_job(bool &job_finished, int32_t &ctaid_X,
int32_t &ctaid_Y, size_t &static_next_block_id,
const size_t static_block_stride,
Expand All @@ -388,6 +450,13 @@ __device__ __forceinline__ void advance_to_next_job(bool &job_finished, int32_t
}
}

__device__ __forceinline__ void set_cta_coords_from_block_id(const size_t block_id,
const size_t work_blocks_X,
int32_t &ctaid_X, int32_t &ctaid_Y) {
ctaid_X = static_cast<int32_t>(block_id % work_blocks_X);
ctaid_Y = static_cast<int32_t>(block_id / work_blocks_X);
}

template <ShapeRepresentation SHAPE_REP, size_t CHUNK_DIM_Y, size_t CHUNK_DIM_X>
__device__ __forceinline__ BlockDescriptor
decode_block(const JobDescriptor &job, const int64_t *const __restrict__ offsets_ptr) {
Expand All @@ -406,6 +475,37 @@ decode_block(const JobDescriptor &job, const int64_t *const __restrict__ offsets
block_offset_Y, block_offset_X);
}

template <ShapeRepresentation SHAPE_REP, size_t CHUNK_DIM_Y, size_t CHUNK_DIM_X>
__device__ __forceinline__ BlockDescriptor decode_block(const JobDescriptor &job,
const size_t tensor_base,
const int32_t ctaid_X,
const int32_t ctaid_Y) {
constexpr bool is_single_tensor = (SHAPE_REP == ShapeRepresentation::SAME_BOTH_DIMS ||
SHAPE_REP == ShapeRepresentation::VARYING_FIRST_DIM);
constexpr size_t ELTS_PER_CHUNK = CHUNK_DIM_Y * CHUNK_DIM_X;
const size_t blocks_X_num_in_current_tensor = DIVUP(job.cols, CHUNK_DIM_X);
size_t block_id_in_current_tensor = 0;
size_t block_id_Y = 0;
size_t block_id_X = 0;
if constexpr (is_single_tensor) {
block_id_X = static_cast<size_t>(ctaid_X);
if constexpr (SHAPE_REP == ShapeRepresentation::SAME_BOTH_DIMS) {
const size_t blocks_Y_per_tensor = DIVUP(job.rows, static_cast<size_t>(CHUNK_DIM_Y));
block_id_Y = static_cast<size_t>(ctaid_Y) - job.tensor_id * blocks_Y_per_tensor;
} else {
const size_t tensor_base_row = tensor_base / job.cols;
block_id_Y = static_cast<size_t>(ctaid_Y) - tensor_base_row / CHUNK_DIM_Y;
}
block_id_in_current_tensor = block_id_Y * blocks_X_num_in_current_tensor + block_id_X;
} else {
block_id_in_current_tensor = job.block_id - tensor_base / ELTS_PER_CHUNK;
block_id_Y = block_id_in_current_tensor / blocks_X_num_in_current_tensor;
block_id_X = block_id_in_current_tensor % blocks_X_num_in_current_tensor;
}
return BlockDescriptor(tensor_base, block_id_in_current_tensor, block_id_Y, block_id_X,
block_id_Y * CHUNK_DIM_Y, block_id_X * CHUNK_DIM_X);
}

} // namespace common
} // namespace dispatch
} // namespace transformer_engine
Expand Down
143 changes: 98 additions & 45 deletions transformer_engine/common/cast/core/grouped_tma.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
#include "../../common.h"
#include "../../util/ptx.cuh"
#include "../../utils.cuh"
#include "../nvfp4/core_nvfp4.cuh"
#include "grouped_layout.cuh"

namespace transformer_engine {
Expand All @@ -38,13 +39,11 @@ struct alignas(128) TensorMapStorage {
alignas(128) CUtensorMap act_input[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
alignas(128) CUtensorMap output_rowwise[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
alignas(128) CUtensorMap output_colwise[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
size_t rows[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
size_t cols[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
size_t offsets[MAX_SUPPORTED_TENSOR_DESCRIPTORS];
};

// Internal linkage avoids device-link ODR issues when this header is included by multiple .cu TUs.
static __device__ TensorMapStorage g_tensor_maps;
alignas(128) static __device__ TensorMetadata g_tensor_metadata[MAX_SUPPORTED_TENSOR_DESCRIPTORS];

inline bool dimensions_supported_by_TMA(const Tensor *const t) {
const size_t cols = t->flat_last_dim();
Expand All @@ -59,18 +58,32 @@ __device__ __forceinline__ unsigned char *align_smem_ptr_per_TMA_requirements(un
return reinterpret_cast<unsigned char *>(addr);
}

__device__ __forceinline__ uintptr_t get_pointer_with_offset_bits(const uintptr_t base_ptr,
const size_t offset_elts,
const size_t data_type_bits) {
const size_t offset_bits = offset_elts * data_type_bits;
if (offset_bits % 8 != 0) {
NVTE_DEVICE_ERROR("Data offset is not byte-aligned.");
}
return base_ptr + offset_bits / 8;
}

// Copies the base tensor map to shmem, modifies the copy, stores the modified tensor map at index
__device__ __forceinline__ void modify_base_tensor_map(const CUtensorMap base_tensor_map,
CUtensorMap *global_tensor_map,
const uintptr_t global_data_ptr,
const size_t global_dim_Y,
const size_t global_dim_X,
const size_t data_type_size_bytes) {
__device__ __forceinline__ void modify_base_tensor_map_bits(const CUtensorMap base_tensor_map,
CUtensorMap *global_tensor_map,
const uintptr_t global_data_ptr,
const size_t global_dim_Y,
const size_t global_dim_X,
const size_t data_type_bits) {
__shared__ CUtensorMap shared_tensor_map;
shared_tensor_map = base_tensor_map; // Copy the base tensor map into shmem
constexpr bool is_blackwell = ARCH_BLACKWELL_FAMILY;
if constexpr (is_blackwell) {
const size_t global_stride_bytes = global_dim_X * data_type_size_bytes;
const size_t global_stride_bits = global_dim_X * data_type_bits;
if (global_stride_bits % 8 != 0) {
NVTE_DEVICE_ERROR("Shape not supported. Data stride must be byte-aligned.");
}
const size_t global_stride_bytes = global_stride_bits / 8;
if (global_stride_bytes % TMA_GMEM_ALIGNMENT != 0) {
NVTE_DEVICE_ERROR("Shape not supported. Data stride must be 16B aligned.");
}
Expand All @@ -96,7 +109,17 @@ __device__ __forceinline__ void modify_base_tensor_map(const CUtensorMap base_te
}
}

template <typename IType, typename OType>
__device__ __forceinline__ void modify_base_tensor_map(const CUtensorMap base_tensor_map,
CUtensorMap *global_tensor_map,
const uintptr_t global_data_ptr,
const size_t global_dim_Y,
const size_t global_dim_X,
const size_t data_type_size_bytes) {
modify_base_tensor_map_bits(base_tensor_map, global_tensor_map, global_data_ptr, global_dim_Y,
global_dim_X, data_type_size_bytes * 8);
}

template <typename IType, typename OType, bool NVFP4_CAST = false>
__global__ void __launch_bounds__(1)
update_tma_descriptors(const __grid_constant__ CUtensorMap base_tensor_map_input,
const __grid_constant__ CUtensorMap base_tensor_map_act_input,
Expand All @@ -113,50 +136,80 @@ __global__ void __launch_bounds__(1)
const int64_t *const __restrict__ last_dims_ptr, const bool rowwise,
const bool colwise, const bool compute_dactivations) {
const size_t tensor_id = blockIdx.x;
const size_t rows =
get_tensor_rows_num(tensor_id, shape_rep, first_logical_dim, first_dims_ptr, num_tensors);
const size_t cols = get_tensor_cols_num(tensor_id, shape_rep, last_logical_dim, last_dims_ptr);
if (tensor_id >= num_tensors) {
return;
}

const size_t offset_elts = offsets_ptr[tensor_id];
g_tensor_maps.rows[tensor_id] = rows;
g_tensor_maps.cols[tensor_id] = cols;
g_tensor_maps.offsets[tensor_id] = offset_elts;
const bool same_both_dims = shape_rep == ShapeRepresentation::SAME_BOTH_DIMS;
const size_t descriptor_first_logical_dim =
(NVFP4_CAST && same_both_dims) ? (first_logical_dim / num_tensors) : first_logical_dim;
const size_t descriptor_num_tensors = (NVFP4_CAST && same_both_dims) ? 1 : num_tensors;
const size_t rows = get_tensor_rows_num(tensor_id, shape_rep, descriptor_first_logical_dim,
first_dims_ptr, descriptor_num_tensors);
const size_t cols = get_tensor_cols_num(tensor_id, shape_rep, last_logical_dim, last_dims_ptr);
const size_t offset_elts =
NVFP4_CAST ? get_tensor_base_offset(tensor_id, shape_rep, first_logical_dim, last_logical_dim,
num_tensors, offsets_ptr)
: static_cast<size_t>(offsets_ptr[tensor_id]);

// Zero-sized groups: skip TMA descriptor update. The main kernel already returns
// early for rows==0 or cols==0, but creating a TMA descriptor with a zero dimension
// is invalid and causes CUDA_ERROR_ILLEGAL_ADDRESS.
if (rows == 0 || cols == 0) {
if constexpr (NVFP4_CAST) {
g_tensor_metadata[tensor_id] = TensorMetadata(rows, cols, offset_elts, 0, 0);
}
return;
}

if (tensor_id < num_tensors) {
{
CUtensorMap *modified_tensor_map_input = &g_tensor_maps.input[tensor_id];
const uintptr_t global_data_ptr = reinterpret_cast<uintptr_t>(input_data_ptr + offset_elts);
modify_base_tensor_map(base_tensor_map_input, modified_tensor_map_input, global_data_ptr,
rows, cols, sizeof(IType));
}
if (compute_dactivations) {
CUtensorMap *modified_tensor_map_act_input = &g_tensor_maps.act_input[tensor_id];
const uintptr_t global_data_ptr =
reinterpret_cast<uintptr_t>(act_input_data_ptr + offset_elts);
modify_base_tensor_map(base_tensor_map_act_input, modified_tensor_map_act_input,
global_data_ptr, rows, cols, sizeof(IType));
}
if (rowwise) {
CUtensorMap *modified_tensor_map_output_rowwise = &g_tensor_maps.output_rowwise[tensor_id];
const uintptr_t global_data_ptr =
reinterpret_cast<uintptr_t>(output_rowwise_data_ptr + offset_elts);
modify_base_tensor_map(base_tensor_map_output_rowwise, modified_tensor_map_output_rowwise,
global_data_ptr, rows, cols, sizeof(OType));
}
if (colwise) {
CUtensorMap *modified_tensor_map_output_colwise = &g_tensor_maps.output_colwise[tensor_id];
const uintptr_t global_data_ptr =
reinterpret_cast<uintptr_t>(output_colwise_data_ptr + offset_elts);
modify_base_tensor_map(base_tensor_map_output_colwise, modified_tensor_map_output_colwise,
global_data_ptr, rows, cols, sizeof(OType));
if constexpr (NVFP4_CAST) {
const size_t rowwise_scale_base =
nvfp4::core::get_rowwise_scale_base(shape_rep, tensor_id, offset_elts, rows, cols);
const size_t colwise_scale_base =
nvfp4::core::get_colwise_scale_base(shape_rep, tensor_id, offset_elts, rows, cols);
g_tensor_metadata[tensor_id] =
TensorMetadata(rows, cols, offset_elts, rowwise_scale_base, colwise_scale_base);
}

constexpr size_t input_type_bits = TypeInfo<IType>::size;
constexpr size_t output_type_bits = []() constexpr {
if constexpr (NVFP4_CAST) {
return size_t{4};
} else {
return TypeInfo<OType>::size;
}
}();
const size_t output_colwise_dim_Y = NVFP4_CAST ? cols : rows;
const size_t output_colwise_dim_X = NVFP4_CAST ? rows : cols;

{
CUtensorMap *modified_tensor_map_input = &g_tensor_maps.input[tensor_id];
const uintptr_t global_data_ptr = get_pointer_with_offset_bits(
reinterpret_cast<uintptr_t>(input_data_ptr), offset_elts, input_type_bits);
modify_base_tensor_map_bits(base_tensor_map_input, modified_tensor_map_input, global_data_ptr,
rows, cols, input_type_bits);
}
if (compute_dactivations) {
CUtensorMap *modified_tensor_map_act_input = &g_tensor_maps.act_input[tensor_id];
const uintptr_t global_data_ptr = get_pointer_with_offset_bits(
reinterpret_cast<uintptr_t>(act_input_data_ptr), offset_elts, input_type_bits);
modify_base_tensor_map_bits(base_tensor_map_act_input, modified_tensor_map_act_input,
global_data_ptr, rows, cols, input_type_bits);
}
if (rowwise) {
CUtensorMap *modified_tensor_map_output_rowwise = &g_tensor_maps.output_rowwise[tensor_id];
const uintptr_t global_data_ptr = get_pointer_with_offset_bits(
reinterpret_cast<uintptr_t>(output_rowwise_data_ptr), offset_elts, output_type_bits);
modify_base_tensor_map_bits(base_tensor_map_output_rowwise, modified_tensor_map_output_rowwise,
global_data_ptr, rows, cols, output_type_bits);
}
if (colwise) {
CUtensorMap *modified_tensor_map_output_colwise = &g_tensor_maps.output_colwise[tensor_id];
const uintptr_t global_data_ptr = get_pointer_with_offset_bits(
reinterpret_cast<uintptr_t>(output_colwise_data_ptr), offset_elts, output_type_bits);
modify_base_tensor_map_bits(base_tensor_map_output_colwise, modified_tensor_map_output_colwise,
global_data_ptr, output_colwise_dim_Y, output_colwise_dim_X,
output_type_bits);
}
}

Expand Down
Loading
Loading