Automated PR - 2026-07-07
This commit is contained in:
@@ -0,0 +1,424 @@
|
||||
/**
|
||||
* @file all2all.cpp
|
||||
* @brief Implementation of All2All communication primitives for multi-GPU tensor parallelism.
|
||||
*
|
||||
* This file implements the All2All class which provides efficient inter-GPU communication
|
||||
* using CUDA IPC (Inter-Process Communication). The implementation supports:
|
||||
* - Head redistribution for tensor-parallel attention (send_recv_heads, gather_heads)
|
||||
* - Sequence gathering for cross-rank aggregation (allgather)
|
||||
*
|
||||
* All operations use a barrier-based synchronization protocol where each GPU writes
|
||||
* directly to remote GPU memory via IPC, then signals completion through atomic
|
||||
* operations on barrier counters.
|
||||
*/
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <ATen/cuda/CUDADataType.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include <chrono>
|
||||
#include <cuda_runtime.h>
|
||||
#include <memory>
|
||||
#include <pybind11/functional.h>
|
||||
#include <torch/python.h>
|
||||
|
||||
#include "all2all.hpp"
|
||||
#include "cuda/api.cuh"
|
||||
#include "cuda/configs.cuh"
|
||||
|
||||
namespace ltx_kernels {
|
||||
namespace all2all {
|
||||
|
||||
/**
|
||||
* Constructs the All2All communication manager.
|
||||
*
|
||||
* Memory Allocation Strategy:
|
||||
* The constructor allocates a single contiguous GPU memory block that contains:
|
||||
* 1. Data buffer (tensor_bytes): Space for tensor data exchange
|
||||
* 2. Barrier signals (MAX_NUM_PEERS * sizeof(int)): Per-rank completion counters
|
||||
* 3. Buffer pointers (MAX_NUM_PEERS * sizeof(void*)): GPU-accessible pointer array
|
||||
* 4. Barrier pointer array (MAX_NUM_PEERS * sizeof(int*)): GPU-accessible signal pointers
|
||||
*
|
||||
* This layout minimizes memory allocations and allows the entire region to be
|
||||
* shared via a single IPC handle.
|
||||
*/
|
||||
All2All::All2All(int rank, int world_size, int num_tokens, int hidden_dim, int num_sms, at::ScalarType tensor_dtype,
|
||||
double timeout_seconds)
|
||||
: rank(rank), world_size(world_size), num_sms(num_sms), max_tokens(num_tokens), num_elems(0), tensor_bytes(0),
|
||||
tensor_dtype(tensor_dtype) {
|
||||
num_elems = int64_t(num_tokens) * int64_t(hidden_dim);
|
||||
tensor_bytes = num_elems * elementSize(tensor_dtype);
|
||||
|
||||
// Derive the barrier timeout from the device's peak SM clock so the wall-clock guard is
|
||||
// correct on any GPU (the kernel counts SM cycles via clock64). Use cudaDeviceGetAttribute,
|
||||
// not cudaDeviceProp::clockRate, which was removed in CUDA 13. The attribute is in kHz.
|
||||
int device = 0;
|
||||
CUDA_CHECK(cudaGetDevice(&device));
|
||||
int sm_clock_khz = 0;
|
||||
CUDA_CHECK(cudaDeviceGetAttribute(&sm_clock_khz, cudaDevAttrClockRate, device));
|
||||
sm_clock_hz_ = static_cast<double>(sm_clock_khz) * 1e3;
|
||||
set_timeout_seconds(timeout_seconds);
|
||||
|
||||
// Calculate sizes for each region of the shared memory block
|
||||
int64_t ptrs_bytes = MAX_NUM_PEERS * sizeof(void *);
|
||||
int64_t barrier_signal_bytes = MAX_NUM_PEERS * sizeof(int);
|
||||
int64_t barrier_signal_ptrs_bytes = MAX_NUM_PEERS * sizeof(int *);
|
||||
|
||||
// Allocate GPU memory for token count arrays (used by kernels)
|
||||
CUDA_CHECK(cudaMalloc(reinterpret_cast<void **>(&rank_tokens_gpu), sizeof(int) * MAX_NUM_PEERS));
|
||||
CUDA_CHECK(cudaMalloc(reinterpret_cast<void **>(&prefix_rank_tokens_gpu), sizeof(int) * MAX_NUM_PEERS));
|
||||
|
||||
// Allocate the main shared memory block and create IPC handle
|
||||
// Layout: [data_buffer | barrier_signals | buffer_ptrs | barrier_signal_ptrs]
|
||||
CUDA_CHECK(
|
||||
cudaMalloc(&buffer_ptrs[rank], tensor_bytes + barrier_signal_bytes + ptrs_bytes + barrier_signal_ptrs_bytes));
|
||||
CUDA_CHECK(cudaIpcGetMemHandle(&ipc_handlers[rank], buffer_ptrs[rank]));
|
||||
|
||||
// Set up pointers to each region within the allocated block
|
||||
buffer_ptrs_gpu =
|
||||
reinterpret_cast<void **>(static_cast<uint8_t *>(buffer_ptrs[rank]) + tensor_bytes + barrier_signal_bytes);
|
||||
barrier_signal_ptrs[rank] = reinterpret_cast<int *>(static_cast<uint8_t *>(buffer_ptrs[rank]) + tensor_bytes);
|
||||
barrier_signal_ptrs_gpu = reinterpret_cast<int **>(static_cast<uint8_t *>(buffer_ptrs[rank]) + tensor_bytes +
|
||||
barrier_signal_bytes + ptrs_bytes);
|
||||
|
||||
// Initialize barrier signals to zero
|
||||
CUDA_CHECK(cudaMemset(barrier_signal_ptrs[rank], 0, barrier_signal_bytes));
|
||||
}
|
||||
|
||||
All2All::~All2All() noexcept(false) {
|
||||
if (!destroyed) {
|
||||
printf("WARNING: destroy() was not called, which can leak resources.\n");
|
||||
fflush(stdout);
|
||||
destroy();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Releases all allocated resources.
|
||||
*
|
||||
* This must be called explicitly before destruction to ensure proper cleanup of:
|
||||
* - IPC memory mappings to remote GPUs
|
||||
* - Local GPU memory allocations
|
||||
*
|
||||
* The method synchronizes the device to ensure all pending operations complete
|
||||
* before releasing resources.
|
||||
*/
|
||||
void All2All::destroy() {
|
||||
if (destroyed) {
|
||||
return;
|
||||
}
|
||||
CUDA_CHECK(cudaDeviceSynchronize());
|
||||
|
||||
// Close IPC mappings to remote GPU memory (skip our own rank)
|
||||
// Only close handles that were actually opened via sync()
|
||||
for (int i = 0; i < world_size; i++) {
|
||||
if (i != rank && buffer_ptrs[i] != nullptr) {
|
||||
CUDA_CHECK(cudaIpcCloseMemHandle(buffer_ptrs[i]));
|
||||
}
|
||||
}
|
||||
|
||||
// Free local GPU memory allocations
|
||||
CUDA_CHECK(cudaFree(buffer_ptrs[rank]));
|
||||
CUDA_CHECK(cudaFree(rank_tokens_gpu));
|
||||
CUDA_CHECK(cudaFree(prefix_rank_tokens_gpu));
|
||||
destroyed = true;
|
||||
}
|
||||
|
||||
/**
|
||||
* Opens IPC memory mappings to all peer GPUs.
|
||||
*
|
||||
* This method processes IPC handles gathered from all ranks and opens memory
|
||||
* mappings to enable direct GPU-to-GPU memory access. After calling this method,
|
||||
* each GPU can read/write directly to any other GPU's buffer via buffer_ptrs.
|
||||
*
|
||||
* The barrier_signal_ptrs are also set up to point to the correct offset within
|
||||
* each peer's shared memory block.
|
||||
*/
|
||||
void All2All::sync(const std::vector<std::optional<pybind11::bytearray>> &all_gathered_handles) {
|
||||
for (int i = 0; i < world_size; i++) {
|
||||
auto handle_str = std::string(all_gathered_handles[i].value());
|
||||
EP_HOST_ASSERT(handle_str.size() == CUDA_IPC_HANDLE_SIZE);
|
||||
|
||||
if (i != rank) {
|
||||
// Open IPC mapping to remote GPU's memory
|
||||
std::memcpy(ipc_handlers[i].reserved, handle_str.c_str(), CUDA_IPC_HANDLE_SIZE);
|
||||
CUDA_CHECK(cudaIpcOpenMemHandle(&buffer_ptrs[i], ipc_handlers[i], cudaIpcMemLazyEnablePeerAccess));
|
||||
// Calculate offset to barrier signals in remote buffer
|
||||
barrier_signal_ptrs[i] = reinterpret_cast<int *>(static_cast<uint8_t *>(buffer_ptrs[i]) + tensor_bytes);
|
||||
} else {
|
||||
// Verify our own handle matches what we sent
|
||||
EP_HOST_ASSERT(std::memcmp(ipc_handlers[i].reserved, handle_str.c_str(), CUDA_IPC_HANDLE_SIZE) == 0);
|
||||
}
|
||||
}
|
||||
|
||||
// Copy pointer arrays to GPU for kernel access
|
||||
CUDA_CHECK(cudaMemcpy(buffer_ptrs_gpu, buffer_ptrs, sizeof(void *) * world_size, cudaMemcpyHostToDevice));
|
||||
CUDA_CHECK(
|
||||
cudaMemcpy(barrier_signal_ptrs_gpu, barrier_signal_ptrs, sizeof(int *) * world_size, cudaMemcpyHostToDevice));
|
||||
CUDA_CHECK(cudaDeviceSynchronize());
|
||||
}
|
||||
|
||||
pybind11::bytearray All2All::get_local_ipc_handle() const {
|
||||
return {ipc_handlers[rank].reserved, CUDA_IPC_HANDLE_SIZE};
|
||||
}
|
||||
|
||||
/**
|
||||
* Configures token distribution across ranks for the current batch.
|
||||
*
|
||||
* This method computes prefix sums needed by the kernels to calculate source
|
||||
* and destination offsets. It must be called before any communication operation
|
||||
* when the token distribution changes between batches.
|
||||
*
|
||||
* Example: For rank_num_tokens = {128, 96, 128, 64}
|
||||
* - rank_tokens = {128, 96, 128, 64}
|
||||
* - prefix_rank_tokens = {0, 128, 224, 352}
|
||||
* - total_tokens = 416
|
||||
*/
|
||||
void All2All::set_rank_tokens(const std::vector<int> &rank_num_tokens) {
|
||||
EP_HOST_ASSERT(static_cast<int>(rank_num_tokens.size()) == world_size);
|
||||
|
||||
// Initialize prefix sums to zero
|
||||
for (int i = 0; i < world_size; i++) {
|
||||
prefix_rank_tokens[i] = 0;
|
||||
}
|
||||
|
||||
// Compute prefix sums (exclusive scan)
|
||||
for (int i = 0; i < world_size; i++) {
|
||||
rank_tokens[i] = rank_num_tokens[i];
|
||||
if (i > 0) {
|
||||
prefix_rank_tokens[i] = prefix_rank_tokens[i - 1] + rank_tokens[i - 1];
|
||||
}
|
||||
}
|
||||
|
||||
// Total tokens is the sum of all rank tokens
|
||||
total_tokens = prefix_rank_tokens[world_size - 1] + rank_tokens[world_size - 1];
|
||||
|
||||
// Copy to GPU for kernel access
|
||||
CUDA_CHECK(cudaMemcpy(rank_tokens_gpu, rank_tokens, sizeof(int) * MAX_NUM_PEERS, cudaMemcpyHostToDevice));
|
||||
CUDA_CHECK(
|
||||
cudaMemcpy(prefix_rank_tokens_gpu, prefix_rank_tokens, sizeof(int) * MAX_NUM_PEERS, cudaMemcpyHostToDevice));
|
||||
CUDA_CHECK(cudaDeviceSynchronize());
|
||||
}
|
||||
|
||||
/**
|
||||
* Creates a tensor from the local IPC buffer.
|
||||
*
|
||||
* This helper method returns either a zero-copy view of the IPC buffer or
|
||||
* a newly allocated tensor with the data copied. The zero-copy mode is more
|
||||
* efficient but the tensor lifetime is tied to the All2All instance.
|
||||
*
|
||||
* @note The buffer pointer is cast to the template type T for proper interpretation.
|
||||
*/
|
||||
at::Tensor All2All::get_local_buffer_tensor(at::Tensor &x, int batch_size, int out_tokens, int out_heads, int head_size,
|
||||
bool should_copy, cudaStream_t stream) {
|
||||
auto ptr = buffer_ptrs[rank];
|
||||
if (should_copy) {
|
||||
// Allocate new tensor and copy data from IPC buffer
|
||||
auto out_tensor = torch::empty({batch_size, out_tokens, out_heads, head_size}, x.options());
|
||||
CUDA_CHECK(cudaMemcpyAsync(out_tensor.data_ptr(), ptr,
|
||||
int64_t(batch_size) * int64_t(out_tokens) * int64_t(out_heads) * int64_t(head_size) *
|
||||
int64_t(elementSize(x.scalar_type())),
|
||||
cudaMemcpyDeviceToDevice, stream));
|
||||
return out_tensor;
|
||||
} else {
|
||||
// Return a view directly into the IPC buffer (zero-copy)
|
||||
auto out_tensor = torch::from_blob(ptr, {batch_size, out_tokens, out_heads, head_size}, x.options());
|
||||
return out_tensor;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* All2All communication to redistribute attention heads across GPUs.
|
||||
*
|
||||
* This operation is used in tensor-parallel transformers to exchange attention heads:
|
||||
* - Before: Each GPU has all tokens but only a subset of heads
|
||||
* - After: Each GPU has all tokens with heads redistributed
|
||||
*
|
||||
* Tensor Layout Transformation:
|
||||
* Input: [batch, local_tokens, all_heads, head_size] per GPU
|
||||
* Output: [batch, all_tokens, heads_per_rank, head_size] per GPU
|
||||
*
|
||||
* The operation partitions heads evenly: heads_per_rank = all_heads / world_size
|
||||
* GPU i receives heads [i*heads_per_rank : (i+1)*heads_per_rank] from all GPUs.
|
||||
*/
|
||||
at::Tensor All2All::send_recv_heads(at::Tensor &x, bool copy_output) {
|
||||
// Validate input tensor properties
|
||||
EP_HOST_ASSERT(x.dim() == 4 and x.is_contiguous());
|
||||
EP_HOST_ASSERT(x.dtype() == tensor_dtype);
|
||||
EP_HOST_ASSERT(x.device().is_cuda());
|
||||
EP_HOST_ASSERT(x.device().index() == rank);
|
||||
|
||||
int batch_size = x.size(0);
|
||||
int num_tokens = x.size(1);
|
||||
int num_heads = x.size(2);
|
||||
int head_size = x.size(3);
|
||||
|
||||
// Output dimensions after redistribution
|
||||
int out_tokens = total_tokens; // All tokens from all ranks
|
||||
int out_heads = num_heads / world_size; // Each rank gets 1/world_size of heads
|
||||
|
||||
EP_HOST_ASSERT(int64_t(batch_size) * int64_t(out_tokens) * int64_t(out_heads) * int64_t(head_size) *
|
||||
int64_t(elementSize(x.scalar_type())) <=
|
||||
tensor_bytes);
|
||||
|
||||
at::cuda::CUDAGuard device_guard{x.device()};
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
// Launch the All2All kernel
|
||||
all2all_cuda::all2all_head_launch(buffer_ptrs_gpu, barrier_signal_ptrs_gpu, x.data_ptr(), prefix_rank_tokens_gpu,
|
||||
rank, world_size, batch_size, total_tokens, num_tokens, num_heads, head_size,
|
||||
stream, num_sms, tensor_dtype, timeout_cycles_);
|
||||
|
||||
return get_local_buffer_tensor(x, batch_size, out_tokens, out_heads, head_size, copy_output, stream);
|
||||
}
|
||||
|
||||
/**
|
||||
* Inverse All2All to gather heads back to original distribution.
|
||||
*
|
||||
* This is the inverse operation of send_recv_heads(). It redistributes data
|
||||
* so each GPU gets back its original tokens with all attention heads.
|
||||
*
|
||||
* Tensor Layout Transformation:
|
||||
* Input: [batch, all_tokens, heads_per_rank, head_size] per GPU
|
||||
* Output: [batch, local_tokens, all_heads, head_size] per GPU
|
||||
*
|
||||
* Each GPU sends its portion of tokens to the originating rank, reconstructing
|
||||
* the original head distribution.
|
||||
*/
|
||||
at::Tensor All2All::gather_heads(at::Tensor &x, bool copy_output) {
|
||||
// Validate input tensor properties
|
||||
EP_HOST_ASSERT(x.dim() == 4 and x.is_contiguous());
|
||||
EP_HOST_ASSERT(x.dtype() == tensor_dtype);
|
||||
EP_HOST_ASSERT(x.device().is_cuda());
|
||||
EP_HOST_ASSERT(x.device().index() == rank);
|
||||
|
||||
at::cuda::CUDAGuard device_guard{x.device()};
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
int batch_size = x.size(0);
|
||||
int num_heads = x.size(2) * world_size; // Reconstruct total head count
|
||||
int head_size = x.size(3);
|
||||
|
||||
// Output dimensions: this rank's tokens with all heads
|
||||
int out_tokens = rank_tokens[rank];
|
||||
int out_heads = num_heads;
|
||||
|
||||
EP_HOST_ASSERT(int64_t(batch_size) * int64_t(out_tokens) * int64_t(out_heads) * int64_t(head_size) *
|
||||
int64_t(elementSize(x.scalar_type())) <=
|
||||
tensor_bytes);
|
||||
|
||||
// Launch the gather kernel
|
||||
all2all_cuda::all2all_head_gather_launch(buffer_ptrs_gpu, barrier_signal_ptrs_gpu, x.data_ptr(), rank_tokens_gpu,
|
||||
prefix_rank_tokens_gpu, rank, world_size, batch_size, total_tokens,
|
||||
num_heads, head_size, stream, num_sms, tensor_dtype, timeout_cycles_);
|
||||
|
||||
return get_local_buffer_tensor(x, batch_size, out_tokens, out_heads, head_size, copy_output, stream);
|
||||
}
|
||||
|
||||
/**
|
||||
* AllGather operation to collect sequence tokens from all ranks.
|
||||
*
|
||||
* Each GPU contributes its local sequence tokens, which are gathered into
|
||||
* a complete sequence replicated on all GPUs. This is typically used after
|
||||
* tensor-parallel operations to reconstruct the full sequence.
|
||||
*
|
||||
* Tensor Layout Transformation:
|
||||
* Input: [batch, local_seqlen, heads, head_size] per GPU
|
||||
* Output: [batch, total_seqlen, heads, head_size] per GPU (identical on all GPUs)
|
||||
*
|
||||
* Each GPU's tokens are placed at offset prefix_rank_tokens[rank] in the output.
|
||||
*/
|
||||
at::Tensor All2All::allgather(at::Tensor &x, bool copy_output) {
|
||||
// Validate input tensor properties
|
||||
EP_HOST_ASSERT(x.dim() == 4 and x.is_contiguous());
|
||||
EP_HOST_ASSERT(x.dtype() == tensor_dtype);
|
||||
EP_HOST_ASSERT(x.device().is_cuda());
|
||||
EP_HOST_ASSERT(x.device().index() == rank);
|
||||
|
||||
at::cuda::CUDAGuard device_guard{x.device()};
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
int batch_size = x.size(0);
|
||||
int seqlen = x.size(1);
|
||||
int num_heads = x.size(2);
|
||||
int head_size = x.size(3);
|
||||
|
||||
// Output contains all tokens from all ranks
|
||||
int out_tokens = total_tokens;
|
||||
int out_heads = num_heads;
|
||||
int hidden_dim = num_heads * head_size;
|
||||
|
||||
EP_HOST_ASSERT(int64_t(batch_size) * int64_t(out_tokens) * int64_t(out_heads) * int64_t(head_size) *
|
||||
int64_t(elementSize(x.scalar_type())) <=
|
||||
tensor_bytes);
|
||||
|
||||
// Launch the allgather kernel
|
||||
all2all_cuda::allgather_launch(buffer_ptrs_gpu, barrier_signal_ptrs_gpu, x.data_ptr(), prefix_rank_tokens_gpu, rank,
|
||||
world_size, batch_size, seqlen, hidden_dim, total_tokens, stream, num_sms,
|
||||
tensor_dtype, timeout_cycles_);
|
||||
|
||||
return get_local_buffer_tensor(x, batch_size, out_tokens, out_heads, head_size, copy_output, stream);
|
||||
}
|
||||
|
||||
} // namespace all2all
|
||||
} // namespace ltx_kernels
|
||||
|
||||
/**
|
||||
* Python bindings for the All2All communication library.
|
||||
*
|
||||
* Usage from Python:
|
||||
* import all2all_cpp
|
||||
*
|
||||
* # Create instance (one per GPU)
|
||||
* comm = all2all_cpp.All2All(rank, world_size, max_tokens, hidden_dim, num_sms, dtype)
|
||||
*
|
||||
* # Exchange IPC handles and synchronize
|
||||
* handle = comm.get_local_ipc_handle()
|
||||
* # ... gather handles via NCCL ...
|
||||
* comm.sync(all_handles)
|
||||
*
|
||||
* # Set token distribution
|
||||
* comm.set_rank_tokens([128, 128, 128, 128])
|
||||
*
|
||||
* # Perform operations
|
||||
* output = comm.send_recv_heads(input_tensor, copy_output=False)
|
||||
*
|
||||
* # Cleanup
|
||||
* comm.destroy()
|
||||
*/
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "High-performance All2All communication library for multi-GPU tensor parallelism.\n\n"
|
||||
"This library provides IPC-based All2All operations optimized for transformer models.\n"
|
||||
"Supported operations:\n"
|
||||
" - send_recv_heads: Redistribute attention heads across GPUs\n"
|
||||
" - gather_heads: Inverse of send_recv_heads\n"
|
||||
" - allgather: Gather sequence tokens from all ranks\n";
|
||||
|
||||
pybind11::class_<ltx_kernels::all2all::All2All>(
|
||||
m, "All2All",
|
||||
"Manages All2All communication state for multi-GPU operations.\n\n"
|
||||
"Args:\n"
|
||||
" rank: This GPU's rank (0 to world_size-1)\n"
|
||||
" world_size: Total number of GPUs\n"
|
||||
" num_tokens: Maximum tokens per rank\n"
|
||||
" hidden_dim: Hidden dimension (heads * head_size)\n"
|
||||
" num_sms: Number of SMs for kernel launches\n"
|
||||
" tensor_dtype: Tensor data type (torch.bfloat16 or torch.float8_e4m3fn)\n"
|
||||
" timeout_seconds: Optional initial barrier timeout in seconds (defaults to the kernel default)")
|
||||
.def(pybind11::init<int, int, int, int, int, at::ScalarType>())
|
||||
.def(pybind11::init<int, int, int, int, int, at::ScalarType, double>())
|
||||
.def("get_local_ipc_handle", <x_kernels::all2all::All2All::get_local_ipc_handle,
|
||||
"Returns the IPC handle for this rank's buffer.")
|
||||
.def("sync", <x_kernels::all2all::All2All::sync, "Opens IPC mappings to all peer GPUs using gathered handles.")
|
||||
.def("destroy", <x_kernels::all2all::All2All::destroy,
|
||||
"Releases all GPU resources. Must be called before destruction.")
|
||||
.def("send_recv_heads", <x_kernels::all2all::All2All::send_recv_heads,
|
||||
"All2All operation to redistribute attention heads.")
|
||||
.def("gather_heads", <x_kernels::all2all::All2All::gather_heads,
|
||||
"Inverse All2All to gather heads back to original distribution.")
|
||||
.def("allgather", <x_kernels::all2all::All2All::allgather, "Gathers sequence tokens from all ranks.")
|
||||
.def("set_rank_tokens", <x_kernels::all2all::All2All::set_rank_tokens,
|
||||
"Sets token counts per rank for the current batch.")
|
||||
.def("set_timeout_seconds", <x_kernels::all2all::All2All::set_timeout_seconds,
|
||||
"Sets the barrier timeout in seconds (converted to cycles via the device peak SM clock).");
|
||||
}
|
||||
@@ -0,0 +1,265 @@
|
||||
/**
|
||||
* @file all2all.hpp
|
||||
* @brief High-performance All2All communication primitives for multi-GPU tensor parallelism.
|
||||
*
|
||||
* This library provides efficient All2All communication operations optimized for transformer
|
||||
* models using tensor parallelism. It uses CUDA IPC (Inter-Process Communication) for
|
||||
* zero-copy data transfer between GPUs in the same node.
|
||||
*
|
||||
* ## Architecture Overview
|
||||
*
|
||||
* The All2All class manages shared memory buffers accessible by all GPUs via IPC handles.
|
||||
* Each GPU allocates a contiguous memory region containing:
|
||||
* - Data buffer: Stores tensor data for exchange
|
||||
* - Barrier signals: Synchronization counters for coordination
|
||||
* - GPU pointer arrays: Device-accessible pointers to all peer buffers
|
||||
*
|
||||
* Memory Layout (per GPU):
|
||||
* ```
|
||||
* |<---- tensor_bytes ---->|<-- barrier signals -->|<-- buffer_ptrs_gpu -->|<-- barrier_signal_ptrs_gpu -->|
|
||||
* | Data Buffer | MAX_PEERS * int | MAX_PEERS * void* | MAX_PEERS * int* |
|
||||
* ```
|
||||
*
|
||||
* ## Supported Operations
|
||||
*
|
||||
* 1. **send_recv_heads**: Redistributes attention heads across GPUs (All2All)
|
||||
* - Input: [batch, tokens, heads, head_size] on each GPU
|
||||
* - Output: [batch, total_tokens, heads/world_size, head_size] on each GPU
|
||||
*
|
||||
* 2. **gather_heads**: Inverse of send_recv_heads
|
||||
* - Gathers distributed heads back to original distribution
|
||||
*
|
||||
* 3. **allgather**: Gathers sequence data from all ranks
|
||||
* - Each GPU contributes its local tokens to form the complete sequence
|
||||
*
|
||||
* ## Thread Safety
|
||||
*
|
||||
* - The class is NOT thread-safe. Each thread/process should have its own instance.
|
||||
* - Multiple CUDA streams may use the same instance sequentially.
|
||||
* - The `destroy()` method MUST be called before destruction to properly release IPC handles.
|
||||
*
|
||||
* ## Usage Example
|
||||
*
|
||||
* ```cpp
|
||||
* // Initialize on each GPU
|
||||
* auto comm = All2All(rank, world_size, max_tokens, hidden_dim, num_sms, dtype);
|
||||
*
|
||||
* // Exchange IPC handles (via NCCL or other collective)
|
||||
* auto my_handle = comm.get_local_ipc_handle();
|
||||
* // ... gather all handles ...
|
||||
* comm.sync(all_handles);
|
||||
*
|
||||
* // Set token distribution for current batch
|
||||
* comm.set_rank_tokens({128, 128, 128, 128}); // tokens per rank
|
||||
*
|
||||
* // Perform All2All on attention heads
|
||||
* auto result = comm.send_recv_heads(input_tensor, copy_output=false);
|
||||
*
|
||||
* // Clean up
|
||||
* comm.destroy();
|
||||
* ```
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cuda/configs.cuh"
|
||||
#include "event.hpp"
|
||||
#include <cmath>
|
||||
#include <limits>
|
||||
#include <pybind11/pybind11.h>
|
||||
#include <pybind11/pytypes.h>
|
||||
#include <stdexcept>
|
||||
#include <torch/types.h>
|
||||
#include <tuple>
|
||||
#include <vector>
|
||||
|
||||
namespace ltx_kernels {
|
||||
namespace all2all {
|
||||
|
||||
/**
|
||||
* @class All2All
|
||||
* @brief Manages All2All communication state and operations for multi-GPU tensor parallelism.
|
||||
*
|
||||
* This class encapsulates the IPC-based communication infrastructure needed for
|
||||
* efficient All2All operations. It maintains shared memory buffers, barrier signals,
|
||||
* and provides methods for head-parallel tensor redistribution.
|
||||
*/
|
||||
struct All2All {
|
||||
private:
|
||||
int rank; ///< This GPU's rank (0 to world_size-1)
|
||||
int world_size; ///< Total number of GPUs in the communication group
|
||||
int num_sms; ///< Number of SMs to use for kernel launches
|
||||
int max_tokens; ///< Maximum number of tokens the buffer was allocated for
|
||||
int64_t num_elems; ///< Number of elements in the data buffer (tokens * hidden_dim)
|
||||
int64_t tensor_bytes; ///< Size of the data buffer in bytes
|
||||
|
||||
/// Host array of pointers to each rank's data buffer (GPU memory)
|
||||
void *buffer_ptrs[MAX_NUM_PEERS] = {nullptr};
|
||||
/// Device-accessible array of buffer pointers (copied to GPU)
|
||||
void **buffer_ptrs_gpu = nullptr;
|
||||
|
||||
/// Host array of pointers to each rank's barrier signal buffer
|
||||
int *barrier_signal_ptrs[MAX_NUM_PEERS] = {nullptr};
|
||||
/// Device-accessible array of barrier signal pointers
|
||||
int **barrier_signal_ptrs_gpu = nullptr;
|
||||
|
||||
/// IPC handles for sharing memory between processes
|
||||
cudaIpcMemHandle_t ipc_handlers[MAX_NUM_PEERS];
|
||||
|
||||
at::ScalarType tensor_dtype; ///< Data type of tensors (BFloat16 or Float8_e4m3fn)
|
||||
bool destroyed = false; ///< Flag to track if resources have been released
|
||||
|
||||
int total_tokens; ///< Sum of tokens across all ranks for current batch
|
||||
int rank_tokens[MAX_NUM_PEERS]; ///< Number of tokens on each rank
|
||||
int prefix_rank_tokens[MAX_NUM_PEERS]; ///< Cumulative sum of tokens (for offset calculation)
|
||||
int *rank_tokens_gpu = nullptr; ///< Device copy of rank_tokens
|
||||
int *prefix_rank_tokens_gpu = nullptr; ///< Device copy of prefix_rank_tokens
|
||||
|
||||
/// Device peak SM clock in Hz (from cudaDeviceGetAttribute(cudaDevAttrClockRate)), queried
|
||||
/// once at construction. Used to convert a wall-clock timeout in seconds to barrier cycles.
|
||||
double sm_clock_hz_ = 0.0;
|
||||
|
||||
/// All2All barrier timeout in GPU clock cycles. The constructor sets it from
|
||||
/// DEFAULT_BARRIER_TIMEOUT_SECONDS and the queried SM clock; raise it (set_timeout_seconds)
|
||||
/// to tolerate large cross-rank kernel-launch skew during the first torch.compile forward,
|
||||
/// where one rank's recompile can delay its launch past the steady-state timeout.
|
||||
uint64_t timeout_cycles_ = 0;
|
||||
|
||||
public:
|
||||
/**
|
||||
* @brief Constructs an All2All communication manager.
|
||||
*
|
||||
* Allocates GPU memory for the local data buffer, barrier signals, and pointer arrays.
|
||||
* The IPC handle for the local buffer is created and can be retrieved via get_local_ipc_handle().
|
||||
*
|
||||
* @param rank This GPU's rank in the communication group (0-indexed)
|
||||
* @param world_size Total number of GPUs/ranks
|
||||
* @param num_tokens Maximum number of tokens this rank will handle
|
||||
* @param hidden_dim Hidden dimension size (heads * head_size)
|
||||
* @param num_sms Number of CUDA SMs to use for kernel execution
|
||||
* @param tensor_dtype Data type for tensors (BFloat16 or Float8_e4m3fn)
|
||||
* @param timeout_seconds Initial barrier timeout in seconds (see set_timeout_seconds); may be
|
||||
* raised/reset at runtime for the first torch.compile forward
|
||||
*/
|
||||
All2All(int rank, int world_size, int num_tokens, int hidden_dim, int num_sms, at::ScalarType tensor_dtype,
|
||||
double timeout_seconds = DEFAULT_BARRIER_TIMEOUT_SECONDS);
|
||||
|
||||
/**
|
||||
* @brief Destructor - warns if destroy() was not called.
|
||||
*
|
||||
* @warning Always call destroy() explicitly before the destructor to properly
|
||||
* release IPC handles. Failing to do so may leak resources.
|
||||
*/
|
||||
~All2All() noexcept(false);
|
||||
|
||||
/**
|
||||
* @brief Synchronizes IPC handles from all ranks and opens remote memory mappings.
|
||||
*
|
||||
* This method must be called after all ranks have created their All2All instances
|
||||
* and exchanged IPC handles via an external collective (e.g., NCCL allgather).
|
||||
*
|
||||
* @param all_gathered_handles Vector of IPC handles from all ranks (indexed by rank)
|
||||
*/
|
||||
void sync(const std::vector<std::optional<pybind11::bytearray>> &all_gathered_handles);
|
||||
|
||||
/**
|
||||
* @brief Returns the IPC handle for this rank's shared buffer.
|
||||
*
|
||||
* The returned handle should be gathered across all ranks and passed to sync().
|
||||
*
|
||||
* @return pybind11::bytearray containing the CUDA IPC handle (CUDA_IPC_HANDLE_SIZE bytes)
|
||||
*/
|
||||
pybind11::bytearray get_local_ipc_handle() const;
|
||||
|
||||
/**
|
||||
* @brief Creates a tensor view or copy of the local output buffer.
|
||||
*
|
||||
* @param x Reference tensor for options (dtype, device)
|
||||
* @param batch_size Batch dimension size
|
||||
* @param out_tokens Output token dimension size
|
||||
* @param out_heads Output heads dimension size
|
||||
* @param head_size Head dimension size
|
||||
* @param should_copy If true, copies data to a new tensor; if false, returns a view
|
||||
* @param stream CUDA stream for async copy
|
||||
* @return Tensor with shape [batch_size, out_tokens, out_heads, head_size]
|
||||
*/
|
||||
at::Tensor get_local_buffer_tensor(at::Tensor &x, int batch_size, int out_tokens, int out_heads, int head_size,
|
||||
bool should_copy, cudaStream_t stream);
|
||||
|
||||
/**
|
||||
* @brief Releases all GPU resources and closes IPC handles.
|
||||
*
|
||||
* This method MUST be called before the object is destroyed. It synchronizes
|
||||
* the device, closes remote IPC mappings, and frees local GPU memory.
|
||||
*/
|
||||
void destroy();
|
||||
|
||||
/**
|
||||
* @brief Performs All2All communication to redistribute attention heads.
|
||||
*
|
||||
* Redistributes tensor from [batch, local_tokens, all_heads, head_size] to
|
||||
* [batch, all_tokens, local_heads, head_size]. Each rank sends its portion
|
||||
* of heads to the corresponding target rank.
|
||||
*
|
||||
* @param x Input tensor with shape [batch, num_tokens, num_heads, head_size]
|
||||
* @param copy_output If true, returns a copy; if false, returns a view of the IPC buffer
|
||||
* @return Tensor with shape [batch, total_tokens, num_heads/world_size, head_size]
|
||||
*/
|
||||
at::Tensor send_recv_heads(at::Tensor &x, bool copy_output);
|
||||
|
||||
/**
|
||||
* @brief Performs inverse All2All to gather heads back to original distribution.
|
||||
*
|
||||
* Inverse of send_recv_heads(). Redistributes from [batch, all_tokens, local_heads, head_size]
|
||||
* back to [batch, local_tokens, all_heads, head_size].
|
||||
*
|
||||
* @param x Input tensor with shape [batch, total_tokens, heads_per_rank, head_size]
|
||||
* @param copy_output If true, returns a copy; if false, returns a view of the IPC buffer
|
||||
* @return Tensor with shape [batch, rank_tokens[rank], num_heads, head_size]
|
||||
*/
|
||||
at::Tensor gather_heads(at::Tensor &x, bool copy_output);
|
||||
|
||||
/**
|
||||
* @brief Gathers sequence tokens from all ranks.
|
||||
*
|
||||
* Each rank contributes its local sequence tokens, which are gathered into
|
||||
* a complete sequence on all ranks.
|
||||
*
|
||||
* @param x Input tensor with shape [batch, seqlen, num_heads, head_size]
|
||||
* @param copy_output If true, returns a copy; if false, returns a view of the IPC buffer
|
||||
* @return Tensor with shape [batch, total_tokens, num_heads, head_size]
|
||||
*/
|
||||
at::Tensor allgather(at::Tensor &x, bool copy_output);
|
||||
|
||||
/**
|
||||
* @brief Sets the token count for each rank in the current batch.
|
||||
*
|
||||
* Must be called before send_recv_heads(), gather_heads(), or allgather()
|
||||
* to configure the token distribution. This allows variable-length sequences
|
||||
* across ranks.
|
||||
*
|
||||
* @param rank_num_tokens Vector of token counts, one per rank (must have world_size elements)
|
||||
*/
|
||||
void set_rank_tokens(const std::vector<int> &rank_num_tokens);
|
||||
|
||||
/**
|
||||
* @brief Sets the all2all barrier timeout in seconds.
|
||||
*
|
||||
* Converted to GPU clock cycles using the device's peak SM clock (queried at construction).
|
||||
* Relaxes deadlock detection during the first torch.compile forward, where asymmetric
|
||||
* per-rank recompilation can delay a rank's kernel launch beyond the steady-state timeout.
|
||||
* Reset to the default for steady-state replay.
|
||||
*/
|
||||
void set_timeout_seconds(double seconds) {
|
||||
if (!std::isfinite(seconds) || seconds < 0.0) {
|
||||
throw std::invalid_argument("All2All timeout (seconds) must be finite and non-negative");
|
||||
}
|
||||
// Saturate rather than overflow the float->uint64 cast (out-of-range conversion is UB).
|
||||
const double cycles = seconds * sm_clock_hz_;
|
||||
const double max_cycles = static_cast<double>(std::numeric_limits<uint64_t>::max());
|
||||
timeout_cycles_ = cycles >= max_cycles ? std::numeric_limits<uint64_t>::max() : static_cast<uint64_t>(cycles);
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace all2all
|
||||
} // namespace ltx_kernels
|
||||
@@ -0,0 +1,372 @@
|
||||
/**
|
||||
* @file all2all_heads.cu
|
||||
* @brief CUDA kernels for All2All attention head redistribution.
|
||||
*
|
||||
* This file implements the GPU kernels for redistributing attention heads across
|
||||
* multiple GPUs using IPC-based direct memory access. The kernels are designed
|
||||
* for tensor-parallel transformer models where attention heads need to be
|
||||
* exchanged between GPUs.
|
||||
*
|
||||
* ## Algorithm Overview
|
||||
*
|
||||
* The kernels use a direct-write approach where each GPU writes its data directly
|
||||
* to the target GPU's memory buffer via IPC. This avoids intermediate copies and
|
||||
* achieves near-peak memory bandwidth utilization.
|
||||
*
|
||||
* ## SM Work Distribution (Round-Robin)
|
||||
*
|
||||
* SMs are distributed round-robin among target ranks to handle non-divisible SM counts:
|
||||
* - SM i writes to rank (i % world_size)
|
||||
* - With 132 SMs and 8 GPUs: ranks 0-3 get 17 SMs, ranks 4-7 get 16 SMs
|
||||
* - Each SM group processes all tokens for its assigned target rank
|
||||
* - Within each group, SMs cooperate to cover all tokens in strided fashion
|
||||
*
|
||||
* ## Synchronization Protocol
|
||||
*
|
||||
* After data transfer, a barrier synchronization ensures all ranks have completed:
|
||||
* 1. Each SM atomically increments the target rank's barrier counter for this rank
|
||||
* 2. SM 0 waits until it has received signals from all ranks
|
||||
* 3. Barrier counters are reset for the next operation
|
||||
*/
|
||||
|
||||
#include "cuda/configs.cuh"
|
||||
#include "cuda/exceptions.cuh"
|
||||
#include "cuda/utils.cuh"
|
||||
#include <ATen/cuda/CUDADataType.h>
|
||||
|
||||
namespace ltx_kernels {
|
||||
namespace all2all {
|
||||
namespace all2all_cuda {
|
||||
|
||||
/**
|
||||
* @brief All2All kernel for redistributing attention heads across GPUs.
|
||||
*
|
||||
* This kernel performs the "send" phase of All2All: each GPU writes its assigned
|
||||
* subset of attention heads to all other GPUs. The data layout transformation is:
|
||||
*
|
||||
* Source: [batch, num_tokens, num_heads, head_size]
|
||||
* Dest: [batch, total_tokens, heads_per_rank, head_size]
|
||||
*
|
||||
* Each GPU writes heads [target_rank * heads_per_rank : (target_rank+1) * heads_per_rank]
|
||||
* to target_rank's buffer at token offset prefix_rank_tokens[rank].
|
||||
*
|
||||
* ## Memory Layout
|
||||
*
|
||||
* Input tensor x (row-major, contiguous):
|
||||
* - Batch dimension: outermost
|
||||
* - Token dimension: batch_stride = num_tokens * num_heads * head_size
|
||||
* - Head dimension: token_stride = num_heads * head_size
|
||||
* - Head element: head_stride = head_size
|
||||
*
|
||||
* Output buffer (per target rank):
|
||||
* - Similar layout but with heads_per_rank instead of num_heads
|
||||
* - Tokens from this rank placed at offset prefix_rank_tokens[rank]
|
||||
*
|
||||
* ## Thread Block Organization
|
||||
*
|
||||
* Each thread block handles multiple tokens cooperatively:
|
||||
* - Threads are organized in a 2D logical grid (rows=tokens, cols=elements)
|
||||
* - Each thread copies 16 bytes (int4) per iteration
|
||||
* - num_threads_per_token = (heads_per_rank * head_size) / elements_per_thread
|
||||
* - num_tokens_per_copy = num_threads / num_threads_per_token
|
||||
*
|
||||
* @tparam ELEM_T Element type (at::BFloat16 or at::Float8_e4m3fn)
|
||||
* @param buffer_ptrs Device array of pointers to each rank's data buffer
|
||||
* @param barrier_signal_ptrs Device array of pointers to each rank's barrier signals
|
||||
* @param x Source tensor data pointer
|
||||
* @param rank This GPU's rank
|
||||
* @param world_size Total number of GPUs
|
||||
* @param batch_size Number of batches
|
||||
* @param num_tokens Number of tokens on this rank
|
||||
* @param num_heads Total number of attention heads
|
||||
* @param head_size Size of each attention head
|
||||
* @param total_tokens Sum of tokens across all ranks
|
||||
* @param prefix_rank_tokens Cumulative token counts for offset calculation
|
||||
*/
|
||||
template <typename ELEM_T>
|
||||
__global__ void send_recv_all2all(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int rank, int world_size,
|
||||
int batch_size, int num_tokens, int num_heads, int head_size, int total_tokens,
|
||||
int *prefix_rank_tokens, uint64_t timeout_cycles) {
|
||||
// Grid dimensions
|
||||
int num_sms = gridDim.x;
|
||||
int sm_id = blockIdx.x;
|
||||
int num_threads = blockDim.x;
|
||||
|
||||
// === SM Work Distribution (Round-Robin) ===
|
||||
// Use modular assignment to handle num_sms not divisible by world_size.
|
||||
// This ensures all SMs are utilized: some ranks get ceil(num_sms/world_size)
|
||||
// SMs, others get floor(num_sms/world_size) SMs.
|
||||
int64_t target_rank = get_target_rank(sm_id, world_size);
|
||||
int64_t rank_local_sm_id = get_rank_local_sm_id(sm_id, world_size);
|
||||
int64_t num_sms_for_this_rank = get_num_sms_for_rank(target_rank, num_sms, world_size);
|
||||
|
||||
// === Head Assignment ===
|
||||
// Heads are partitioned evenly: rank i gets heads [i*hpr : (i+1)*hpr]
|
||||
int64_t heads_per_rank = num_heads / world_size;
|
||||
int64_t head_id = target_rank * heads_per_rank; // Starting head for target rank
|
||||
|
||||
// === Thread Mapping ===
|
||||
// Each thread copies an int4 (16 bytes) per memory operation
|
||||
// Threads form a 2D grid: (tokens_per_copy, threads_per_token)
|
||||
int64_t num_elems_per_thread = sizeof(int4) / sizeof(ELEM_T);
|
||||
int64_t num_threads_per_token = heads_per_rank * head_size / num_elems_per_thread;
|
||||
int64_t num_tokens_per_copy = num_threads / num_threads_per_token;
|
||||
|
||||
// 2D thread coordinates within the logical grid
|
||||
int64_t copy_thr_col_idx = threadIdx.x % num_threads_per_token; // Element offset
|
||||
int64_t copy_thr_row_idx = threadIdx.x / num_threads_per_token; // Token offset
|
||||
|
||||
// Get target rank's buffer pointer
|
||||
auto ptr = reinterpret_cast<void *>(static_cast<int8_t *>(buffer_ptrs[target_rank]));
|
||||
|
||||
// Use 64-bit arithmetic to avoid overflow for large tensors
|
||||
int64_t num_tokens_64b = int64_t(num_tokens);
|
||||
int64_t num_heads_64b = int64_t(num_heads);
|
||||
int64_t head_size_64b = int64_t(head_size);
|
||||
|
||||
// === Main Copy Loop ===
|
||||
// Iterate over batches and tokens, with SMs in the same group
|
||||
// working on different token ranges in strided fashion
|
||||
for (int64_t batch_ind = 0; batch_ind < batch_size; batch_ind++) {
|
||||
// Strided token iteration: each SM in the group handles different token ranges
|
||||
for (int64_t token_idx = rank_local_sm_id * num_tokens_per_copy; token_idx < num_tokens;
|
||||
token_idx += num_tokens_per_copy * num_sms_for_this_rank) {
|
||||
int64_t copy_token_idx = token_idx + copy_thr_row_idx;
|
||||
// Destination token index accounts for this rank's offset in the global sequence
|
||||
int64_t dst_token_idx = prefix_rank_tokens[rank] + copy_token_idx;
|
||||
|
||||
if (copy_token_idx >= num_tokens)
|
||||
break;
|
||||
|
||||
// === Pointer Arithmetic ===
|
||||
// Source: Read from this rank's input tensor at [batch, token, head_id:head_id+hpr, :]
|
||||
// Note: We read a contiguous chunk of heads starting at head_id
|
||||
int4 *shuffled_x_ptr =
|
||||
reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(x) +
|
||||
batch_ind * num_tokens_64b * num_heads_64b * head_size_64b * sizeof(ELEM_T) +
|
||||
copy_token_idx * num_heads_64b * head_size_64b * sizeof(ELEM_T) +
|
||||
head_id * head_size_64b * sizeof(ELEM_T)) +
|
||||
copy_thr_col_idx;
|
||||
|
||||
// Destination: Write to target rank's buffer at [batch, dst_token, :, :]
|
||||
// The buffer has layout [batch, total_tokens, heads_per_rank, head_size]
|
||||
int4 *shuffled_buffer_ptr =
|
||||
reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(ptr) +
|
||||
batch_ind * total_tokens * heads_per_rank * head_size_64b * sizeof(ELEM_T) +
|
||||
dst_token_idx * heads_per_rank * head_size_64b * sizeof(ELEM_T)) +
|
||||
copy_thr_col_idx;
|
||||
|
||||
// Non-allocating store to avoid polluting L1 cache
|
||||
st_na_global(shuffled_buffer_ptr, __ldg(shuffled_x_ptr));
|
||||
}
|
||||
}
|
||||
|
||||
// === Barrier Synchronization ===
|
||||
// Signal completion to target rank and wait for all ranks to finish
|
||||
barrier_wait_and_reset_roundrobin(barrier_signal_ptrs, target_rank, rank, world_size, num_sms, sm_id, threadIdx.x,
|
||||
timeout_cycles);
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief All2All kernel for gathering attention heads back to original distribution.
|
||||
*
|
||||
* This kernel performs the inverse of send_recv_all2all: it gathers heads from
|
||||
* all ranks back to reconstruct the original tensor layout. Each GPU reads from
|
||||
* its local buffer and writes its portion of heads to all target ranks.
|
||||
*
|
||||
* Data layout transformation:
|
||||
* Source: [batch, total_tokens, heads_per_rank, head_size] (per GPU)
|
||||
* Dest: [batch, rank_tokens[target], num_heads, head_size] (per target GPU)
|
||||
*
|
||||
* ## Memory Layout
|
||||
*
|
||||
* Input tensor x (this rank's portion after send_recv_all2all):
|
||||
* - Contains all tokens but only heads_per_rank heads
|
||||
* - Layout: [batch, total_tokens, heads_per_rank, head_size]
|
||||
*
|
||||
* Output buffer (per target rank):
|
||||
* - Contains only that rank's tokens but all heads
|
||||
* - Layout: [batch, rank_tokens[target], num_heads, head_size]
|
||||
* - This rank writes heads [rank * heads_per_rank : (rank+1) * heads_per_rank]
|
||||
*
|
||||
* @tparam ELEM_T Element type (at::BFloat16 or at::Float8_e4m3fn)
|
||||
* @param buffer_ptrs Device array of pointers to each rank's data buffer
|
||||
* @param barrier_signal_ptrs Device array of pointers to barrier signals
|
||||
* @param x Source tensor data (this rank's buffer after send_recv)
|
||||
* @param rank This GPU's rank
|
||||
* @param world_size Total number of GPUs
|
||||
* @param batch_size Number of batches
|
||||
* @param num_heads Total number of heads (reconstructed)
|
||||
* @param head_size Size of each attention head
|
||||
* @param rank_tokens Number of tokens for each rank
|
||||
* @param total_tokens Sum of tokens across all ranks
|
||||
* @param prefix_rank_tokens Cumulative token counts for offset calculation
|
||||
*/
|
||||
template <typename ELEM_T>
|
||||
__global__ void gather_heads(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int rank, int world_size,
|
||||
int batch_size, int num_heads, int head_size, const int *__restrict__ rank_tokens,
|
||||
int total_tokens, int *prefix_rank_tokens, uint64_t timeout_cycles) {
|
||||
// Grid dimensions
|
||||
int num_sms = gridDim.x;
|
||||
int sm_id = blockIdx.x;
|
||||
int num_threads = blockDim.x;
|
||||
|
||||
// === SM Work Distribution (Round-Robin) ===
|
||||
// Same partitioning as send_recv_all2all
|
||||
int64_t target_rank = get_target_rank(sm_id, world_size);
|
||||
int64_t rank_local_sm_id = get_rank_local_sm_id(sm_id, world_size);
|
||||
int64_t num_sms_for_this_rank = get_num_sms_for_rank(target_rank, num_sms, world_size);
|
||||
int64_t heads_per_rank = num_heads / world_size;
|
||||
|
||||
// === Thread Mapping ===
|
||||
int64_t num_elems_per_thread = sizeof(int4) / sizeof(ELEM_T);
|
||||
int64_t num_threads_per_token = heads_per_rank * head_size / num_elems_per_thread;
|
||||
int64_t num_tokens_per_copy = num_threads / num_threads_per_token;
|
||||
|
||||
int64_t copy_thr_col_idx = threadIdx.x % num_threads_per_token;
|
||||
int64_t copy_thr_row_idx = threadIdx.x / num_threads_per_token;
|
||||
|
||||
// Number of tokens owned by target rank
|
||||
const int64_t tgt_tokens = int64_t(rank_tokens[target_rank]);
|
||||
|
||||
// This rank writes its heads at offset [rank * heads_per_rank] in the output
|
||||
int64_t head_idx = rank * heads_per_rank;
|
||||
int64_t num_heads_64b = int64_t(num_heads);
|
||||
int64_t head_size_64b = int64_t(head_size);
|
||||
int64_t total_tokens_64b = int64_t(total_tokens);
|
||||
|
||||
// Get target rank's buffer pointer
|
||||
auto ptr = reinterpret_cast<void *>(static_cast<int8_t *>(buffer_ptrs[target_rank]));
|
||||
|
||||
// === Main Copy Loop ===
|
||||
// Process target rank's tokens: read from global position, write to local position
|
||||
for (int64_t batch_idx = 0; batch_idx < batch_size; batch_idx++) {
|
||||
for (int64_t token_idx = rank_local_sm_id * num_tokens_per_copy; token_idx < tgt_tokens;
|
||||
token_idx += num_tokens_per_copy * num_sms_for_this_rank) {
|
||||
int64_t copy_token = token_idx + copy_thr_row_idx;
|
||||
if (copy_token >= tgt_tokens)
|
||||
break;
|
||||
|
||||
// Source: Read from global token position (target rank's tokens in our buffer)
|
||||
int64_t src_token_idx = prefix_rank_tokens[target_rank] + copy_token;
|
||||
// Destination: Write to local token position in target's buffer
|
||||
int64_t dst_token_idx = copy_token;
|
||||
|
||||
// Source pointer: our input tensor at [batch, src_token, :, :]
|
||||
int4 *shuffled_x_ptr =
|
||||
reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(x) +
|
||||
batch_idx * total_tokens_64b * heads_per_rank * head_size_64b * sizeof(ELEM_T) +
|
||||
src_token_idx * heads_per_rank * head_size_64b * sizeof(ELEM_T)) +
|
||||
copy_thr_col_idx;
|
||||
|
||||
// Destination pointer: target's buffer at [batch, dst_token, head_idx:head_idx+hpr, :]
|
||||
int4 *shuffled_buffer_ptr =
|
||||
reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(ptr) +
|
||||
batch_idx * tgt_tokens * num_heads_64b * head_size_64b * sizeof(ELEM_T) +
|
||||
dst_token_idx * num_heads_64b * head_size_64b * sizeof(ELEM_T) +
|
||||
head_idx * head_size_64b * sizeof(ELEM_T)) +
|
||||
copy_thr_col_idx;
|
||||
|
||||
st_na_global(shuffled_buffer_ptr, __ldg(shuffled_x_ptr));
|
||||
}
|
||||
}
|
||||
|
||||
// === Barrier Synchronization ===
|
||||
barrier_wait_and_reset_roundrobin(barrier_signal_ptrs, target_rank, rank, world_size, num_sms, sm_id, threadIdx.x,
|
||||
timeout_cycles);
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Host function to launch the gather_heads kernel.
|
||||
*
|
||||
* Selects the appropriate template instantiation based on tensor data type
|
||||
* and launches the kernel with the specified number of SMs.
|
||||
*
|
||||
* @param buffer_ptrs Device array of buffer pointers
|
||||
* @param barrier_signal_ptrs Device array of barrier signal pointers
|
||||
* @param x Input tensor data pointer
|
||||
* @param rank_tokens Token count per rank (device memory)
|
||||
* @param prefix_rank_tokens Cumulative token counts (device memory)
|
||||
* @param rank This GPU's rank
|
||||
* @param world_size Total number of GPUs
|
||||
* @param batch_size Number of batches
|
||||
* @param total_tokens Sum of tokens across all ranks
|
||||
* @param num_heads Total number of attention heads
|
||||
* @param head_size Size of each attention head
|
||||
* @param stream CUDA stream for async execution
|
||||
* @param num_sms Number of SMs to launch
|
||||
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
|
||||
*/
|
||||
void all2all_head_gather_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, const int *rank_tokens,
|
||||
int *prefix_rank_tokens, int rank, int world_size, int batch_size, int total_tokens,
|
||||
int num_heads, int head_size, cudaStream_t stream, int num_sms,
|
||||
at::ScalarType tensor_dtype, uint64_t timeout_cycles) {
|
||||
do {
|
||||
if (tensor_dtype == at::ScalarType::BFloat16) {
|
||||
gather_heads<at::BFloat16><<<num_sms, DEFAULT_KERNEL_THREADS, 0, stream>>>(
|
||||
buffer_ptrs, barrier_signal_ptrs, x, rank, world_size, batch_size, num_heads, head_size, rank_tokens,
|
||||
total_tokens, prefix_rank_tokens, timeout_cycles);
|
||||
} else if (tensor_dtype == at::ScalarType::Float8_e4m3fn) {
|
||||
gather_heads<at::Float8_e4m3fn><<<num_sms, DEFAULT_KERNEL_THREADS, 0, stream>>>(
|
||||
buffer_ptrs, barrier_signal_ptrs, x, rank, world_size, batch_size, num_heads, head_size, rank_tokens,
|
||||
total_tokens, prefix_rank_tokens, timeout_cycles);
|
||||
}
|
||||
|
||||
// Check for kernel launch errors
|
||||
cudaError_t e = cudaGetLastError();
|
||||
if (e != cudaSuccess) {
|
||||
EPException cuda_exception("CUDA", __FILE__, __LINE__, cudaGetErrorString(e));
|
||||
fprintf(stderr, "%s\n", cuda_exception.what());
|
||||
throw cuda_exception;
|
||||
}
|
||||
} while (0);
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Host function to launch the send_recv_all2all kernel.
|
||||
*
|
||||
* Selects the appropriate template instantiation based on tensor data type
|
||||
* and launches the kernel with the specified number of SMs.
|
||||
*
|
||||
* @param buffer_ptrs Device array of buffer pointers
|
||||
* @param barrier_signal_ptrs Device array of barrier signal pointers
|
||||
* @param x Input tensor data pointer
|
||||
* @param prefix_rank_tokens Cumulative token counts (device memory)
|
||||
* @param rank This GPU's rank
|
||||
* @param world_size Total number of GPUs
|
||||
* @param batch_size Number of batches
|
||||
* @param total_tokens Sum of tokens across all ranks
|
||||
* @param num_tokens Number of tokens on this rank
|
||||
* @param num_heads Total number of attention heads
|
||||
* @param head_size Size of each attention head
|
||||
* @param stream CUDA stream for async execution
|
||||
* @param num_sms Number of SMs to launch
|
||||
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
|
||||
*/
|
||||
void all2all_head_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int *prefix_rank_tokens, int rank,
|
||||
int world_size, int batch_size, int total_tokens, int num_tokens, int num_heads, int head_size,
|
||||
cudaStream_t stream, int num_sms, at::ScalarType tensor_dtype, uint64_t timeout_cycles) {
|
||||
do {
|
||||
if (tensor_dtype == at::ScalarType::BFloat16) {
|
||||
send_recv_all2all<at::BFloat16><<<num_sms, DEFAULT_KERNEL_THREADS, 0, stream>>>(
|
||||
buffer_ptrs, barrier_signal_ptrs, x, rank, world_size, batch_size, num_tokens, num_heads, head_size,
|
||||
total_tokens, prefix_rank_tokens, timeout_cycles);
|
||||
} else if (tensor_dtype == at::ScalarType::Float8_e4m3fn) {
|
||||
send_recv_all2all<at::Float8_e4m3fn><<<num_sms, DEFAULT_KERNEL_THREADS, 0, stream>>>(
|
||||
buffer_ptrs, barrier_signal_ptrs, x, rank, world_size, batch_size, num_tokens, num_heads, head_size,
|
||||
total_tokens, prefix_rank_tokens, timeout_cycles);
|
||||
}
|
||||
|
||||
// Check for kernel launch errors
|
||||
cudaError_t e = cudaGetLastError();
|
||||
if (e != cudaSuccess) {
|
||||
EPException cuda_exception("CUDA", __FILE__, __LINE__, cudaGetErrorString(e));
|
||||
fprintf(stderr, "%s\n", cuda_exception.what());
|
||||
throw cuda_exception;
|
||||
}
|
||||
} while (0);
|
||||
}
|
||||
|
||||
} // namespace all2all_cuda
|
||||
} // namespace all2all
|
||||
} // namespace ltx_kernels
|
||||
@@ -0,0 +1,198 @@
|
||||
/**
|
||||
* @file allgather.cu
|
||||
* @brief CUDA kernel for AllGather operation using IPC-based direct memory access.
|
||||
*
|
||||
* This file implements the GPU kernel for gathering sequence tokens from all GPUs
|
||||
* into a complete sequence on each GPU. Unlike the head redistribution kernels,
|
||||
* this kernel preserves the head dimension and only gathers across the token
|
||||
* (sequence) dimension.
|
||||
*
|
||||
* ## Algorithm Overview
|
||||
*
|
||||
* Each GPU broadcasts its local tokens to all other GPUs' buffers:
|
||||
* - GPU i writes its tokens to position [prefix_rank_tokens[i]] in each buffer
|
||||
* - After completion, all buffers contain the full sequence [0:total_tokens]
|
||||
*
|
||||
* ## Use Case
|
||||
*
|
||||
* This is typically used after tensor-parallel computation to reconstruct the
|
||||
* full sequence for operations that require global context (e.g., output projection).
|
||||
*/
|
||||
|
||||
#include "cuda/configs.cuh"
|
||||
#include "cuda/exceptions.cuh"
|
||||
#include "cuda/utils.cuh"
|
||||
#include <ATen/cuda/CUDADataType.h>
|
||||
|
||||
namespace ltx_kernels {
|
||||
namespace all2all {
|
||||
namespace all2all_cuda {
|
||||
|
||||
/**
|
||||
* @brief AllGather kernel to collect sequence tokens from all ranks.
|
||||
*
|
||||
* Each GPU writes its local sequence tokens to all other GPUs' buffers at the
|
||||
* appropriate offset. After synchronization, all GPUs have the complete sequence.
|
||||
*
|
||||
* Data layout transformation:
|
||||
* Input per GPU: [batch, seqlen, hidden_dim]
|
||||
* Output per GPU: [batch, total_tokens, hidden_dim] (identical on all GPUs)
|
||||
*
|
||||
* ## Memory Layout
|
||||
*
|
||||
* Input tensor x (contiguous):
|
||||
* - Shape: [batch, seqlen, hidden_dim]
|
||||
* - hidden_dim = num_heads * head_size (flattened)
|
||||
*
|
||||
* Output buffer (per target rank, after gather):
|
||||
* - Shape: [batch, total_tokens, hidden_dim]
|
||||
* - This rank's tokens placed at offset rank_tokens_prefix[rank]
|
||||
*
|
||||
* ## Thread Mapping
|
||||
*
|
||||
* Similar to all2all_heads, threads cooperate to copy tokens:
|
||||
* - Each thread copies 16 bytes (int4)
|
||||
* - Threads per token = hidden_dim * sizeof(ELEM_T) / sizeof(int4)
|
||||
* - Multiple tokens processed per thread block
|
||||
*
|
||||
* @tparam ELEM_T Element type (__nv_bfloat16 or at::Float8_e4m3fn)
|
||||
* @param x Source tensor data pointer (this rank's tokens)
|
||||
* @param buffer_ptrs Device array of pointers to each rank's data buffer
|
||||
* @param barrier_signal_ptrs Device array of pointers to barrier signals
|
||||
* @param batch_size Number of batches
|
||||
* @param seqlen Number of tokens on this rank
|
||||
* @param hidden_dim Hidden dimension size (num_heads * head_size)
|
||||
* @param world_size Total number of GPUs
|
||||
* @param rank This GPU's rank
|
||||
* @param total_tokens Sum of tokens across all ranks
|
||||
* @param rank_tokens_prefix Cumulative token counts (device memory)
|
||||
*/
|
||||
template <typename ELEM_T>
|
||||
__global__ void allgather(void *x, void **buffer_ptrs, int **barrier_signal_ptrs, int batch_size, int seqlen,
|
||||
int hidden_dim, int world_size, int rank, int total_tokens, int *rank_tokens_prefix,
|
||||
uint64_t timeout_cycles) {
|
||||
|
||||
// Grid dimensions
|
||||
int num_sms = gridDim.x;
|
||||
int sm_id = blockIdx.x;
|
||||
int num_threads = blockDim.x;
|
||||
|
||||
// === SM Work Distribution (Round-Robin) ===
|
||||
// Use modular assignment to handle num_sms not divisible by world_size.
|
||||
// This ensures all SMs are utilized: some ranks get ceil(num_sms/world_size)
|
||||
// SMs, others get floor(num_sms/world_size) SMs.
|
||||
int tgt_rank = get_target_rank(sm_id, world_size);
|
||||
int rank_local_sm_id = get_rank_local_sm_id(sm_id, world_size);
|
||||
int num_sms_for_this_rank = get_num_sms_for_rank(tgt_rank, num_sms, world_size);
|
||||
|
||||
// Get target rank's buffer pointer
|
||||
auto ptr = reinterpret_cast<void *>(static_cast<int8_t *>(buffer_ptrs[tgt_rank]));
|
||||
|
||||
// === Thread Mapping ===
|
||||
// Each thread copies one int4 (16 bytes)
|
||||
int64_t num_elems_per_thread = sizeof(int4) / sizeof(ELEM_T);
|
||||
int64_t num_threads_per_token = hidden_dim / num_elems_per_thread;
|
||||
int64_t num_tokens_per_copy = num_threads / num_threads_per_token;
|
||||
|
||||
// 2D thread coordinates
|
||||
int64_t copy_thr_col_idx = threadIdx.x % num_threads_per_token; // Element offset
|
||||
int64_t copy_thr_row_idx = threadIdx.x / num_threads_per_token; // Token offset
|
||||
|
||||
// Use 64-bit arithmetic to avoid overflow
|
||||
int64_t hidden_dim_64b = int64_t(hidden_dim);
|
||||
int64_t total_tokens_64b = int64_t(total_tokens);
|
||||
int64_t seqlen_64b = int64_t(seqlen);
|
||||
|
||||
// === Main Copy Loop ===
|
||||
// Broadcast this rank's tokens to all target ranks' buffers
|
||||
for (int64_t batch_idx = 0; batch_idx < batch_size; batch_idx++) {
|
||||
// Strided token iteration within SM group for this target rank
|
||||
for (int64_t token_idx = rank_local_sm_id * num_tokens_per_copy; token_idx < seqlen;
|
||||
token_idx += num_tokens_per_copy * num_sms_for_this_rank) {
|
||||
int64_t copy_token = token_idx + copy_thr_row_idx;
|
||||
if (copy_token >= seqlen)
|
||||
break;
|
||||
|
||||
// Source: local token index in input tensor
|
||||
int64_t src_token_idx = copy_token;
|
||||
// Destination: global token index in output buffer
|
||||
// This rank's tokens start at prefix_rank_tokens[rank]
|
||||
int64_t dst_token_idx = copy_token + rank_tokens_prefix[rank];
|
||||
|
||||
// Source pointer: input tensor at [batch, src_token, :]
|
||||
int4 *shuffled_x_ptr = reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(x) +
|
||||
batch_idx * seqlen_64b * hidden_dim_64b * sizeof(ELEM_T) +
|
||||
src_token_idx * hidden_dim_64b * sizeof(ELEM_T)) +
|
||||
copy_thr_col_idx;
|
||||
|
||||
// Destination pointer: target buffer at [batch, dst_token, :]
|
||||
int4 *shuffled_buffer_ptr =
|
||||
reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(ptr) +
|
||||
batch_idx * total_tokens_64b * hidden_dim_64b * sizeof(ELEM_T) +
|
||||
dst_token_idx * hidden_dim_64b * sizeof(ELEM_T)) +
|
||||
copy_thr_col_idx;
|
||||
|
||||
// Non-allocating store for better cache behavior
|
||||
st_na_global(shuffled_buffer_ptr, __ldg(shuffled_x_ptr));
|
||||
}
|
||||
}
|
||||
|
||||
// === Barrier Synchronization ===
|
||||
// Signal completion to target rank and wait for all ranks
|
||||
// Use round-robin variant since SM counts per rank may differ
|
||||
barrier_wait_and_reset_roundrobin(barrier_signal_ptrs, tgt_rank, rank, world_size, num_sms, sm_id, threadIdx.x,
|
||||
timeout_cycles);
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Host function to launch the allgather kernel.
|
||||
*
|
||||
* Launches the AllGather kernel with the specified configuration.
|
||||
* Uses ALLGATHER_KERNEL_THREADS (1024) threads per block for higher
|
||||
* occupancy than the All2All kernels.
|
||||
*
|
||||
* @param buffer_ptrs Device array of buffer pointers
|
||||
* @param barrier_signal_ptrs Device array of barrier signal pointers
|
||||
* @param x Input tensor data pointer
|
||||
* @param prefix_rank_tokens Cumulative token counts (device memory)
|
||||
* @param rank This GPU's rank
|
||||
* @param world_size Total number of GPUs
|
||||
* @param batch_size Number of batches
|
||||
* @param seqlen Number of tokens on this rank
|
||||
* @param hidden_dim Hidden dimension size
|
||||
* @param total_tokens Sum of tokens across all ranks
|
||||
* @param stream CUDA stream for async execution
|
||||
* @param num_sms Number of SMs to launch
|
||||
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
|
||||
*/
|
||||
void allgather_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int *prefix_rank_tokens, int rank,
|
||||
int world_size, int batch_size, int seqlen, int hidden_dim, int total_tokens, cudaStream_t stream,
|
||||
int num_sms, at::ScalarType tensor_dtype, uint64_t timeout_cycles) {
|
||||
do {
|
||||
if (tensor_dtype == at::ScalarType::BFloat16) {
|
||||
allgather<at::BFloat16><<<num_sms, ALLGATHER_KERNEL_THREADS, 0, stream>>>(
|
||||
x, buffer_ptrs, barrier_signal_ptrs, batch_size, seqlen, hidden_dim, world_size, rank, total_tokens,
|
||||
prefix_rank_tokens, timeout_cycles);
|
||||
} else if (tensor_dtype == at::ScalarType::Float8_e4m3fn) {
|
||||
allgather<at::Float8_e4m3fn><<<num_sms, ALLGATHER_KERNEL_THREADS, 0, stream>>>(
|
||||
x, buffer_ptrs, barrier_signal_ptrs, batch_size, seqlen, hidden_dim, world_size, rank, total_tokens,
|
||||
prefix_rank_tokens, timeout_cycles);
|
||||
} else {
|
||||
EPException dtype_exception("allgather_launch", __FILE__, __LINE__, "Unsupported dtype");
|
||||
fprintf(stderr, "%s\n", dtype_exception.what());
|
||||
throw dtype_exception;
|
||||
}
|
||||
|
||||
// Check for kernel launch errors
|
||||
cudaError_t e = cudaGetLastError();
|
||||
if (e != cudaSuccess) {
|
||||
EPException cuda_exception("CUDA", __FILE__, __LINE__, cudaGetErrorString(e));
|
||||
fprintf(stderr, "%s\n", cuda_exception.what());
|
||||
throw cuda_exception;
|
||||
}
|
||||
} while (0);
|
||||
}
|
||||
|
||||
} // namespace all2all_cuda
|
||||
} // namespace all2all
|
||||
} // namespace ltx_kernels
|
||||
@@ -0,0 +1,99 @@
|
||||
/**
|
||||
* @file api.cuh
|
||||
* @brief CUDA kernel launch function declarations for All2All operations.
|
||||
*
|
||||
* This header provides the host-callable interface for launching the All2All
|
||||
* CUDA kernels. These functions handle template instantiation and kernel
|
||||
* configuration based on the tensor data type.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <ATen/cuda/CUDADataType.h>
|
||||
#include <vector>
|
||||
|
||||
namespace ltx_kernels {
|
||||
namespace all2all {
|
||||
namespace all2all_cuda {
|
||||
|
||||
/**
|
||||
* @brief Launches the All2All head redistribution kernel.
|
||||
*
|
||||
* Redistributes attention heads across GPUs:
|
||||
* Input: [batch, num_tokens, num_heads, head_size] per GPU
|
||||
* Output: [batch, total_tokens, num_heads/world_size, head_size] per GPU
|
||||
*
|
||||
* @param buffer_ptrs Device array of pointers to each rank's data buffer
|
||||
* @param barrier_signal_ptrs Device array of pointers to barrier signals
|
||||
* @param x Source tensor data pointer
|
||||
* @param prefix_rank_tokens Cumulative token counts per rank (device memory)
|
||||
* @param rank This GPU's rank (0 to world_size-1)
|
||||
* @param world_size Total number of GPUs
|
||||
* @param batch_size Batch dimension size
|
||||
* @param total_tokens Sum of tokens across all ranks
|
||||
* @param num_tokens Number of tokens on this rank
|
||||
* @param num_heads Total number of attention heads
|
||||
* @param head_size Size of each attention head
|
||||
* @param stream CUDA stream for async execution
|
||||
* @param num_sms Number of SMs to use for the kernel
|
||||
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
|
||||
*/
|
||||
void all2all_head_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int *prefix_rank_tokens, int rank,
|
||||
int world_size, int batch_size, int total_tokens, int num_tokens, int num_heads, int head_size,
|
||||
cudaStream_t stream, int num_sms, at::ScalarType tensor_dtype, uint64_t timeout_cycles);
|
||||
|
||||
/**
|
||||
* @brief Launches the gather heads kernel (inverse of all2all_head_launch).
|
||||
*
|
||||
* Redistributes tokens back to original head distribution:
|
||||
* Input: [batch, total_tokens, heads_per_rank, head_size] per GPU
|
||||
* Output: [batch, rank_tokens[rank], num_heads, head_size] per GPU
|
||||
*
|
||||
* @param buffer_ptrs Device array of pointers to each rank's data buffer
|
||||
* @param barrier_signal_ptrs Device array of pointers to barrier signals
|
||||
* @param x Source tensor data pointer
|
||||
* @param rank_tokens Token count for each rank (device memory)
|
||||
* @param prefix_rank_tokens Cumulative token counts (device memory)
|
||||
* @param rank This GPU's rank
|
||||
* @param world_size Total number of GPUs
|
||||
* @param batch_size Batch dimension size
|
||||
* @param total_tokens Sum of tokens across all ranks
|
||||
* @param num_heads Total number of attention heads (reconstructed)
|
||||
* @param head_size Size of each attention head
|
||||
* @param stream CUDA stream for async execution
|
||||
* @param num_sms Number of SMs to use for the kernel
|
||||
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
|
||||
*/
|
||||
void all2all_head_gather_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, const int *rank_tokens,
|
||||
int *prefix_rank_tokens, int rank, int world_size, int batch_size, int total_tokens,
|
||||
int num_heads, int head_size, cudaStream_t stream, int num_sms,
|
||||
at::ScalarType tensor_dtype, uint64_t timeout_cycles);
|
||||
|
||||
/**
|
||||
* @brief Launches the AllGather kernel for sequence tokens.
|
||||
*
|
||||
* Gathers sequence tokens from all ranks:
|
||||
* Input: [batch, seqlen, hidden_dim] per GPU
|
||||
* Output: [batch, total_tokens, hidden_dim] per GPU (identical on all)
|
||||
*
|
||||
* @param buffer_ptrs Device array of pointers to each rank's data buffer
|
||||
* @param barrier_signal_ptrs Device array of pointers to barrier signals
|
||||
* @param x Source tensor data pointer
|
||||
* @param prefix_rank_tokens Cumulative token counts (device memory)
|
||||
* @param rank This GPU's rank
|
||||
* @param world_size Total number of GPUs
|
||||
* @param batch_size Batch dimension size
|
||||
* @param seqlen Number of tokens on this rank
|
||||
* @param hidden_dim Hidden dimension size (num_heads * head_size)
|
||||
* @param total_tokens Sum of tokens across all ranks
|
||||
* @param stream CUDA stream for async execution
|
||||
* @param num_sms Number of SMs to use for the kernel
|
||||
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
|
||||
*/
|
||||
void allgather_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int *prefix_rank_tokens, int rank,
|
||||
int world_size, int batch_size, int seqlen, int hidden_dim, int total_tokens, cudaStream_t stream,
|
||||
int num_sms, at::ScalarType tensor_dtype, uint64_t timeout_cycles);
|
||||
|
||||
} // namespace all2all_cuda
|
||||
} // namespace all2all
|
||||
} // namespace ltx_kernels
|
||||
@@ -0,0 +1,89 @@
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <torch/extension.h>
|
||||
#include <vector>
|
||||
#include <stdio.h>
|
||||
|
||||
#ifdef __SM90__
|
||||
#include "sm90_fp8_gemm_1d2d_bias.hpp"
|
||||
#endif
|
||||
|
||||
#include "sm89_fp8_gemm_1d2d.hpp"
|
||||
|
||||
namespace blockwise{
|
||||
template <int N>
|
||||
static auto get_shape(const torch::Tensor& t) {
|
||||
return [&t] <size_t... Is> (std::index_sequence<Is...>) {
|
||||
return std::make_tuple(static_cast<int>(t.sizes()[Is])...);
|
||||
}(std::make_index_sequence<N>());
|
||||
}
|
||||
|
||||
#ifdef __SM90__
|
||||
static void fp8_gemm_nt_sm90(const std::pair<torch::Tensor, torch::Tensor>& a,
|
||||
const std::pair<torch::Tensor, torch::Tensor>& b,
|
||||
const torch::Tensor& d,
|
||||
const std::optional<torch::Tensor>& bias,
|
||||
const std::optional<torch::Tensor>& c, const int num_sms) {
|
||||
|
||||
// Type and shape checks
|
||||
const auto& [m , k ] = get_shape<2>(a.first);
|
||||
const auto& [n , k_] = get_shape<2>(b.first);
|
||||
const auto& [m_, n_] = get_shape<2>(d);
|
||||
|
||||
// The SM90 kernel always adds bias; synthesize a zero bias when the layer is
|
||||
// bias-less (e.g. the no-bias video FFN of v3 checkpoints), mirroring SM89 below.
|
||||
torch::Tensor bias_tensor = bias.has_value()
|
||||
? bias.value()
|
||||
: torch::zeros({n}, d.options().dtype(torch::kFloat32));
|
||||
sm90_fp8_gemm_1d2d_bias(a.first, a.second, b.first, b.second, bias_tensor, c, d, m, n, k, num_sms);
|
||||
}
|
||||
#endif
|
||||
|
||||
static void fp8_gemm_nt_sm89(const std::pair<torch::Tensor, torch::Tensor>& a,
|
||||
const std::pair<torch::Tensor, torch::Tensor>& b,
|
||||
const torch::Tensor& d,
|
||||
const std::optional<torch::Tensor>& bias,
|
||||
const bool use_fast_accum = true) {
|
||||
|
||||
const auto& [m, k] = get_shape<2>(a.first);
|
||||
const auto& [n, k_] = get_shape<2>(b.first);
|
||||
const auto& [m_, n_] = get_shape<2>(d);
|
||||
|
||||
// The SM89 kernel always adds bias; synthesize a zero bias when the layer is
|
||||
// bias-less so we add 0 rather than uninitialized memory (mirrors SM90 above).
|
||||
torch::Tensor bias_tensor = bias.has_value()
|
||||
? bias.value()
|
||||
: torch::zeros({n}, d.options().dtype(torch::kFloat32));
|
||||
|
||||
blockwise::sm89_fp8_gemm_1d2d_bias(
|
||||
a.first, a.second, // a data, sfa scales
|
||||
b.first, b.second, // b data, sfb scales
|
||||
bias_tensor, // bias (or empty tensor)
|
||||
d, // output
|
||||
m, n, k,
|
||||
use_fast_accum); // pass through accumulation mode
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
// m.def("package_name", &function_name, "function_docstring"")
|
||||
#ifdef __SM90__
|
||||
m.def("fp8_gemm_nt_sm90", &fp8_gemm_nt_sm90,
|
||||
py::arg("a"), py::arg("b"), py::arg("d"),
|
||||
py::arg("bias") = std::nullopt,
|
||||
py::arg("c") = std::nullopt,
|
||||
py::arg("num_sms") = 132
|
||||
);
|
||||
#endif
|
||||
m.def("fp8_gemm_nt_sm89", &fp8_gemm_nt_sm89,
|
||||
py::arg("a"), py::arg("b"), py::arg("d"),
|
||||
py::arg("bias") = std::nullopt,
|
||||
py::arg("use_fast_accum") = true
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
#pragma once
|
||||
#include <torch/python.h>
|
||||
#include <cute/arch/mma_sm100_umma.hpp>
|
||||
#include "utils.hpp"
|
||||
#include "exceptions.hpp"
|
||||
|
||||
namespace blockwise{
|
||||
struct MulticastConfig {
|
||||
int num_multicast;
|
||||
bool is_multicast_on_a;
|
||||
|
||||
MulticastConfig(const int& num_multicast, const bool& is_multicast_on_a):
|
||||
num_multicast(num_multicast), is_multicast_on_a(is_multicast_on_a) {
|
||||
DG_HOST_ASSERT(1 <= num_multicast and num_multicast <= 2);
|
||||
}
|
||||
};
|
||||
|
||||
struct SharedMemoryConfig {
|
||||
int smem_size;
|
||||
int swizzle_a_mode;
|
||||
int swizzle_b_mode;
|
||||
int swizzle_cd_mode;
|
||||
};
|
||||
|
||||
struct ThreadConfig {
|
||||
int num_threads;
|
||||
|
||||
// SM90
|
||||
int num_tma_threads;
|
||||
int num_math_threads;
|
||||
|
||||
// SM100
|
||||
int num_non_epilogue_threads;
|
||||
int num_epilogue_threads;
|
||||
|
||||
static ThreadConfig sm90(const int& num_tma_threads,
|
||||
const int& num_math_threads) {
|
||||
auto config = ThreadConfig();
|
||||
config.num_threads = num_tma_threads + num_math_threads;
|
||||
config.num_tma_threads = num_tma_threads;
|
||||
config.num_math_threads = num_math_threads;
|
||||
return config;
|
||||
}
|
||||
|
||||
static ThreadConfig sm100(const int& num_non_epilogue_threads,
|
||||
const int& num_epilogue_threads) {
|
||||
auto config = ThreadConfig();
|
||||
config.num_threads = num_non_epilogue_threads + num_epilogue_threads;
|
||||
config.num_non_epilogue_threads = num_non_epilogue_threads;
|
||||
config.num_epilogue_threads = num_epilogue_threads;
|
||||
return config;
|
||||
}
|
||||
};
|
||||
|
||||
template<int SM>
|
||||
struct GemmConfig{};
|
||||
// {
|
||||
// // Templated configs
|
||||
|
||||
// at::ScalarType ab_dtype, cd_dtype;
|
||||
// bool with_accumulation;
|
||||
// int block_m, block_n, block_k;
|
||||
// int num_stages, num_last_stages;
|
||||
|
||||
// // Templated device configs
|
||||
// int num_sms;
|
||||
|
||||
// // Structured configs
|
||||
// MulticastConfig multicast_config;
|
||||
// SharedMemoryConfig smem_config;
|
||||
// ThreadConfig thread_config;
|
||||
// };
|
||||
|
||||
|
||||
template <>
|
||||
struct GemmConfig<90>
|
||||
{
|
||||
at::ScalarType ab_dtype = torch::kFloat8_e4m3fn;
|
||||
at::ScalarType cd_dtype = torch::kBFloat16;
|
||||
bool with_accumulation = false;
|
||||
int block_m = 256;
|
||||
int block_n = 128;
|
||||
int block_k = 128;
|
||||
int num_stages = 3;
|
||||
int num_last_stages = 2;
|
||||
int num_sms = 132;
|
||||
MulticastConfig multicast_config{2, true};
|
||||
SharedMemoryConfig smem_config{216240, 128, 128, 128};
|
||||
ThreadConfig thread_config = ThreadConfig::sm90(128, 256);
|
||||
};
|
||||
|
||||
};
|
||||
@@ -0,0 +1,65 @@
|
||||
#pragma once
|
||||
|
||||
#include <exception>
|
||||
#include <string>
|
||||
#include <sstream>
|
||||
|
||||
namespace blockwise {
|
||||
|
||||
class DGException final : public std::exception {
|
||||
std::string message = {};
|
||||
|
||||
public:
|
||||
explicit DGException(const char *name, const char* file, const int line, const std::string& error) {
|
||||
message = std::string(name) + " error (" + file + ":" + std::to_string(line) + "): " + error;
|
||||
}
|
||||
|
||||
const char *what() const noexcept override {
|
||||
return message.c_str();
|
||||
}
|
||||
};
|
||||
|
||||
#ifndef DG_STATIC_ASSERT
|
||||
#define DG_STATIC_ASSERT(cond, ...) static_assert(cond, __VA_ARGS__)
|
||||
#endif
|
||||
|
||||
#ifndef DG_HOST_ASSERT
|
||||
#define DG_HOST_ASSERT(cond) \
|
||||
do { \
|
||||
if (not (cond)) { \
|
||||
throw DGException("Assertion", __FILE__, __LINE__, #cond); \
|
||||
} \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
#ifndef DG_HOST_UNREACHABLE
|
||||
#define DG_HOST_UNREACHABLE(reason) (throw DGException("Assertion", __FILE__, __LINE__, reason))
|
||||
#endif
|
||||
|
||||
// #ifndef DG_CUDA_DRIVER_CHECK
|
||||
// #define DG_CUDA_DRIVER_CHECK(cmd) \
|
||||
// do { \
|
||||
// const auto& e = (cmd); \
|
||||
// if (e != CUDA_SUCCESS) { \
|
||||
// std::stringstream ss; \
|
||||
// const char *name, *info; \
|
||||
// cuGetErrorName(e, &name), cuGetErrorString(e, &info); \
|
||||
// ss << static_cast<int>(e) << " (" << name << ", " << info << ")"; \
|
||||
// throw DGException("CUDA driver", __FILE__, __LINE__, ss.str()); \
|
||||
// } \
|
||||
// } while (0)
|
||||
// #endif
|
||||
|
||||
#ifndef DG_CUDA_RUNTIME_CHECK
|
||||
#define DG_CUDA_RUNTIME_CHECK(cmd) \
|
||||
do { \
|
||||
const auto& e = (cmd); \
|
||||
if (e != cudaSuccess) { \
|
||||
std::stringstream ss; \
|
||||
ss << static_cast<int>(e) << " (" << cudaGetErrorName(e) << ", " << cudaGetErrorString(e) << ")"; \
|
||||
throw DGException("CUDA runtime", __FILE__, __LINE__, ss.str()); \
|
||||
} \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
} // namespace deep_gemm
|
||||
+48
@@ -0,0 +1,48 @@
|
||||
#pragma once
|
||||
|
||||
namespace cute {
|
||||
|
||||
struct ignore_t {
|
||||
template <typename T>
|
||||
constexpr const ignore_t& operator=(T&&) const noexcept {
|
||||
return *this;
|
||||
}
|
||||
};
|
||||
|
||||
inline constexpr ignore_t ignore{};
|
||||
|
||||
} // namespace cute
|
||||
|
||||
#define CUTE_TIE_CONCAT_IMPL(A, B) A##B
|
||||
#define CUTE_TIE_CONCAT(A, B) CUTE_TIE_CONCAT_IMPL(A, B)
|
||||
|
||||
#define CUTE_TIE_GET_NTH_ARG(_1, _2, _3, _4, _5, _6, _7, _8, _9, _10, N, ...) N
|
||||
#define CUTE_TIE_COUNT_ARGS(...) \
|
||||
CUTE_TIE_GET_NTH_ARG(__VA_ARGS__, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0)
|
||||
|
||||
#define CUTE_TIE_OP_DECL(I, TUPLE, VAR) auto VAR = ::cute::get<I>(TUPLE)
|
||||
#define CUTE_TIE_OP_ASSIGN(I, TUPLE, VAR) VAR = ::cute::get<I>(TUPLE)
|
||||
|
||||
#define CUTE_TIE_APPLY_OP_1(OP, T, V1) OP(0, T, V1);
|
||||
#define CUTE_TIE_APPLY_OP_2(OP, T, V1, V2) OP(0, T, V1); OP(1, T, V2);
|
||||
#define CUTE_TIE_APPLY_OP_3(OP, T, V1, V2, V3) OP(0, T, V1); OP(1, T, V2); OP(2, T, V3);
|
||||
#define CUTE_TIE_APPLY_OP_4(OP, T, V1, V2, V3, V4) OP(0, T, V1); OP(1, T, V2); OP(2, T, V3); OP(3, T, V4);
|
||||
#define CUTE_TIE_APPLY_OP_5(OP, T, V1, V2, V3, V4, V5) OP(0, T, V1); OP(1, T, V2); OP(2, T, V3); OP(3, T, V4); OP(4, T, V5);
|
||||
|
||||
#define CUTE_TIE_DECL(TUPLE_EXPR, ...) \
|
||||
auto&& CUTE_TIE_CONCAT(cute_tie__temp_tuple_, __LINE__) = (TUPLE_EXPR); \
|
||||
CUTE_TIE_CONCAT(CUTE_TIE_APPLY_OP_, CUTE_TIE_COUNT_ARGS(__VA_ARGS__)) ( \
|
||||
CUTE_TIE_OP_DECL, \
|
||||
CUTE_TIE_CONCAT(cute_tie__temp_tuple_, __LINE__), \
|
||||
__VA_ARGS__ \
|
||||
)
|
||||
|
||||
#define CUTE_TIE(TUPLE_EXPR, ...) \
|
||||
do { \
|
||||
auto&& CUTE_TIE_CONCAT(cute_tie__temp_tuple_, __LINE__) = (TUPLE_EXPR); \
|
||||
CUTE_TIE_CONCAT(CUTE_TIE_APPLY_OP_, CUTE_TIE_COUNT_ARGS(__VA_ARGS__)) ( \
|
||||
CUTE_TIE_OP_ASSIGN, \
|
||||
CUTE_TIE_CONCAT(cute_tie__temp_tuple_, __LINE__), \
|
||||
__VA_ARGS__ \
|
||||
); \
|
||||
} while (0)
|
||||
+27
@@ -0,0 +1,27 @@
|
||||
#pragma once
|
||||
|
||||
#include <deep_gemm/common/types.hpp>
|
||||
#include <deep_gemm/common/utils.cuh>
|
||||
|
||||
namespace deep_gemm {
|
||||
|
||||
struct EpilogueIdentity {
|
||||
template <uint32_t STORE_BLOCK_N>
|
||||
__device__ __forceinline__ static uint32_t apply_index_n(const uint32_t &n_idx) {
|
||||
return n_idx;
|
||||
}
|
||||
};
|
||||
|
||||
template <uint32_t kLeft, uint32_t kMid, uint32_t kRight>
|
||||
struct EpilogueHeadSplits: EpilogueIdentity {
|
||||
template <uint32_t STORE_BLOCK_N>
|
||||
__device__ __forceinline__ static uint32_t apply_index_n(const uint32_t &n_idx) {
|
||||
DG_STATIC_ASSERT(kLeft % STORE_BLOCK_N == 0 and kMid % STORE_BLOCK_N == 0
|
||||
and kRight % STORE_BLOCK_N == 0, "Invalid head splits config");
|
||||
return n_idx + (n_idx + kRight) / (kLeft + kRight) * kMid;
|
||||
}
|
||||
};
|
||||
|
||||
#pragma clang diagnostic pop
|
||||
|
||||
} // namespace deep_gemm
|
||||
+44
@@ -0,0 +1,44 @@
|
||||
#pragma once
|
||||
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp8.h>
|
||||
#include <cuda/std/cstdint>
|
||||
#include <cuda/std/utility>
|
||||
|
||||
#include <deep_gemm/common/utils.cuh>
|
||||
|
||||
// Operation functors
|
||||
template <typename T> struct ReduceSum { __device__ T operator()(T a, T b) const { return a + b; } };
|
||||
template <typename T> struct ReduceMax { __device__ T operator()(T a, T b) const { return a > b ? a : b; } };
|
||||
template <typename T> struct ReduceMin { __device__ T operator()(T a, T b) const { return a < b ? a : b; } };
|
||||
template <typename T> struct ReduceAnd { __device__ T operator()(T a, T b) const { return a & b; } };
|
||||
template <typename T> struct ReduceOr { __device__ T operator()(T a, T b) const { return a | b; } };
|
||||
|
||||
// Unified reduction function
|
||||
template <int kNumLanesPerGroup, bool kIntergroupReduce, typename T, typename Op>
|
||||
__forceinline__ __device__ T warp_reduce(T value, Op op) {
|
||||
DG_STATIC_ASSERT(kNumLanesPerGroup == 32 or kNumLanesPerGroup == 16 or kNumLanesPerGroup == 8 or
|
||||
kNumLanesPerGroup == 4 or kNumLanesPerGroup == 2 or kNumLanesPerGroup == 1,
|
||||
"Invalid number of lanes");
|
||||
constexpr uint32_t mask = 0xffffffff;
|
||||
if constexpr (kIntergroupReduce) {
|
||||
if constexpr (kNumLanesPerGroup <= 1) value = op(value, __shfl_xor_sync(mask, value, 1));
|
||||
if constexpr (kNumLanesPerGroup <= 2) value = op(value, __shfl_xor_sync(mask, value, 2));
|
||||
if constexpr (kNumLanesPerGroup <= 4) value = op(value, __shfl_xor_sync(mask, value, 4));
|
||||
if constexpr (kNumLanesPerGroup <= 8) value = op(value, __shfl_xor_sync(mask, value, 8));
|
||||
if constexpr (kNumLanesPerGroup <= 16) value = op(value, __shfl_xor_sync(mask, value, 16));
|
||||
} else {
|
||||
if constexpr (kNumLanesPerGroup >= 32) value = op(value, __shfl_xor_sync(mask, value, 16));
|
||||
if constexpr (kNumLanesPerGroup >= 16) value = op(value, __shfl_xor_sync(mask, value, 8));
|
||||
if constexpr (kNumLanesPerGroup >= 8) value = op(value, __shfl_xor_sync(mask, value, 4));
|
||||
if constexpr (kNumLanesPerGroup >= 4) value = op(value, __shfl_xor_sync(mask, value, 2));
|
||||
if constexpr (kNumLanesPerGroup >= 2) value = op(value, __shfl_xor_sync(mask, value, 1));
|
||||
}
|
||||
return value;
|
||||
}
|
||||
|
||||
// Convenience aliases
|
||||
template <int kNumLanesPerGroup = 32, bool kIntergroupReduce = false, typename T>
|
||||
__forceinline__ __device__ T warp_reduce_sum(T value) {
|
||||
return warp_reduce<kNumLanesPerGroup, kIntergroupReduce, T>(value, ReduceSum<T>{});
|
||||
}
|
||||
+239
@@ -0,0 +1,239 @@
|
||||
#pragma once
|
||||
|
||||
#include <deep_gemm/common/types.hpp>
|
||||
#include <deep_gemm/common/utils.cuh>
|
||||
|
||||
namespace deep_gemm {
|
||||
|
||||
enum class KGroupedIndexType {
|
||||
MN,
|
||||
K,
|
||||
SF_K,
|
||||
};
|
||||
|
||||
template <GemmType kGemmType, uint32_t BLOCK_M, uint32_t BLOCK_N, uint32_t kNumSMs, bool kIsMulticastOnA>
|
||||
static constexpr uint32_t get_num_1d_blocks_per_group() {
|
||||
// Select the best from candidates
|
||||
uint32_t num_best_blocks = 0, min_usage = cute::numeric_limits<uint32_t>::max();
|
||||
for (const auto& candidate: {8u, 16u}) {
|
||||
const auto& usage = kIsMulticastOnA ?
|
||||
candidate * BLOCK_N + constexpr_ceil_div(kNumSMs, candidate) * BLOCK_M: // Grouping on N
|
||||
candidate * BLOCK_M + constexpr_ceil_div(kNumSMs, candidate) * BLOCK_N; // Grouping on M
|
||||
if (usage < min_usage)
|
||||
min_usage = usage, num_best_blocks = candidate;
|
||||
}
|
||||
return num_best_blocks;
|
||||
}
|
||||
|
||||
#pragma clang diagnostic push
|
||||
#pragma ide diagnostic ignored "cppcoreguidelines-pro-type-member-init"
|
||||
template <GemmType kGemmType,
|
||||
uint32_t BLOCK_M, uint32_t BLOCK_N,
|
||||
uint32_t kNumGroups,
|
||||
uint32_t kNumMulticast, bool kIsMulticastOnA,
|
||||
uint32_t kNumSMs,
|
||||
uint32_t SF_K_ALIGNMENT = 512u, // for k-grouped GEMM only: 128 (SM90 float SF) or 512 (SM100 UE8M0 SF)
|
||||
uint32_t kNum1DBlocksPerGroup = get_num_1d_blocks_per_group<kGemmType, BLOCK_M, BLOCK_N, kNumSMs, kIsMulticastOnA>()>
|
||||
struct Scheduler {
|
||||
int current_iter = -1;
|
||||
|
||||
// Block configs
|
||||
uint32_t num_blocks;
|
||||
uint32_t num_m_blocks;
|
||||
uint32_t num_n_blocks;
|
||||
|
||||
// For SM90 multicast checks
|
||||
uint32_t num_blocks_in_group;
|
||||
bool is_peer_cta_alive = true;
|
||||
|
||||
// For grouped GEMM
|
||||
int* grouped_layout;
|
||||
uint32_t current_group_idx = 0;
|
||||
// Only used for masked layout
|
||||
uint32_t current_m_cumsum = 0;
|
||||
// Only used for k-grouped layout
|
||||
uint32_t current_shape_k, current_num_valid_groups = 0, current_k_cumsum = 0, current_sf_k_cumsum = 0;
|
||||
uint32_t next_group_idx, next_shape_k;
|
||||
|
||||
// Only used for k-grouped gemm
|
||||
__device__ __forceinline__ void get_next_k_group(uint32_t &group_idx, uint32_t &shape_k) const {
|
||||
for (; group_idx < kNumGroups; ++ group_idx) {
|
||||
shape_k = __ldg(grouped_layout + group_idx);
|
||||
if (shape_k > 0)
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// ReSharper disable once CppPossiblyUninitializedMember
|
||||
__device__ __forceinline__ explicit Scheduler(const uint32_t& shape_m, const uint32_t& shape_n, const uint32_t& shape_k,
|
||||
int* grouped_layout = nullptr) {
|
||||
num_m_blocks = ceil_div(shape_m, BLOCK_M);
|
||||
num_n_blocks = ceil_div(shape_n, BLOCK_N);
|
||||
current_shape_k = shape_k;
|
||||
if constexpr (kGemmType == GemmType::Normal) {
|
||||
num_blocks = num_m_blocks * num_n_blocks;
|
||||
} else if (kGemmType == GemmType::MGroupedContiguous) {
|
||||
num_blocks = num_m_blocks * num_n_blocks;
|
||||
this->grouped_layout = grouped_layout;
|
||||
} else if (kGemmType == GemmType::MGroupedMasked) {
|
||||
this->grouped_layout = grouped_layout;
|
||||
} else if (kGemmType == GemmType::KGroupedContiguous) {
|
||||
this->grouped_layout = grouped_layout;
|
||||
get_next_k_group(current_group_idx, current_shape_k);
|
||||
next_group_idx = current_group_idx + 1;
|
||||
get_next_k_group(next_group_idx, next_shape_k);
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void get_swizzled_block_idx(const uint32_t& block_idx, uint32_t& m_block_idx, uint32_t& n_block_idx) {
|
||||
DG_STATIC_ASSERT(kNum1DBlocksPerGroup % kNumMulticast == 0, "Invalid group size");
|
||||
|
||||
// Swizzle for better L2 usages
|
||||
const auto& primary_num_blocks = kIsMulticastOnA ? num_n_blocks : num_m_blocks;
|
||||
const auto& secondary_num_blocks = kIsMulticastOnA ? num_m_blocks : num_n_blocks;
|
||||
const auto& num_blocks_per_group = secondary_num_blocks * kNum1DBlocksPerGroup;
|
||||
const auto& group_idx = block_idx / num_blocks_per_group;
|
||||
auto first_block_idx = group_idx * kNum1DBlocksPerGroup;
|
||||
auto in_group_idx = block_idx % num_blocks_per_group;
|
||||
num_blocks_in_group = min(kNum1DBlocksPerGroup, primary_num_blocks - first_block_idx);
|
||||
|
||||
// Fix unaligned TMA multicast
|
||||
// NOTES: for SM90 only, as SM90 can dynamically disable TMA multicast
|
||||
// while SM100 uses 2-CTA, which can not be dynamically disabled
|
||||
#if __CUDA_ARCH__ < 1000
|
||||
if (kNumMulticast > 1 and num_blocks_in_group % 2 != 0) {
|
||||
if (in_group_idx < (num_blocks_in_group ^ 1) * secondary_num_blocks) {
|
||||
num_blocks_in_group = num_blocks_in_group ^ 1;
|
||||
} else {
|
||||
in_group_idx = in_group_idx - (num_blocks_in_group ^ 1) * secondary_num_blocks;
|
||||
first_block_idx += num_blocks_in_group ^ 1;
|
||||
num_blocks_in_group = 1;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
// Convert to final M/N block indices
|
||||
// `kIsMulticastOnA == true` leads to groups on N
|
||||
if constexpr (kIsMulticastOnA) {
|
||||
m_block_idx = in_group_idx / num_blocks_in_group;
|
||||
n_block_idx = first_block_idx + in_group_idx % num_blocks_in_group;
|
||||
} else {
|
||||
m_block_idx = first_block_idx + in_group_idx % num_blocks_in_group;
|
||||
n_block_idx = in_group_idx / num_blocks_in_group;
|
||||
}
|
||||
}
|
||||
|
||||
template <bool kWithGroupOffset, KGroupedIndexType kIndexType = KGroupedIndexType::MN>
|
||||
__device__ __forceinline__ uint32_t get_global_idx(const uint32_t shape_dim, const uint32_t block_size,
|
||||
const uint32_t& block_idx, const uint32_t& m_block_idx = 0) {
|
||||
if constexpr (kGemmType == GemmType::Normal) {
|
||||
return block_idx * block_size;
|
||||
} else if constexpr (kGemmType == GemmType::MGroupedContiguous) {
|
||||
const auto offset = kWithGroupOffset ? cute::max(0, __ldg(grouped_layout + m_block_idx * BLOCK_M)) : 0;
|
||||
return offset * shape_dim + block_idx * block_size;
|
||||
} else if constexpr (kGemmType == GemmType::MGroupedMasked) {
|
||||
const auto offset = kWithGroupOffset ? current_group_idx : 0;
|
||||
return offset * shape_dim + block_idx * block_size;
|
||||
} else if constexpr (kGemmType == GemmType::KGroupedContiguous) {
|
||||
auto offset = 0;
|
||||
if constexpr (kWithGroupOffset) {
|
||||
if constexpr (kIndexType == KGroupedIndexType::MN)
|
||||
offset = current_group_idx * shape_dim;
|
||||
else if constexpr (kIndexType == KGroupedIndexType::K)
|
||||
offset = current_k_cumsum;
|
||||
else if constexpr (kIndexType == KGroupedIndexType::SF_K)
|
||||
offset = current_sf_k_cumsum;
|
||||
}
|
||||
return offset + block_idx * block_size;
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__ bool get_next_block(uint32_t& m_block_idx, uint32_t& n_block_idx) {
|
||||
const auto next_block_idx = (++ current_iter) * kNumSMs + blockIdx.x;
|
||||
|
||||
if constexpr (kGemmType == GemmType::MGroupedMasked) {
|
||||
while (true) {
|
||||
// End of the task
|
||||
if (current_group_idx == kNumGroups)
|
||||
return false;
|
||||
|
||||
// Within current group
|
||||
num_m_blocks = ceil_div(static_cast<uint32_t>(__ldg(grouped_layout + current_group_idx)), BLOCK_M);
|
||||
const auto current_m_block_cumsum = current_m_cumsum + num_m_blocks;
|
||||
if (next_block_idx < current_m_block_cumsum * num_n_blocks)
|
||||
break;
|
||||
|
||||
// Move to check the next group
|
||||
current_group_idx ++, current_m_cumsum = current_m_block_cumsum;
|
||||
}
|
||||
|
||||
get_swizzled_block_idx(next_block_idx - current_m_cumsum * num_n_blocks, m_block_idx, n_block_idx);
|
||||
} else if (kGemmType == GemmType::KGroupedContiguous) {
|
||||
while (true) {
|
||||
// End of the task
|
||||
if (current_group_idx == kNumGroups)
|
||||
return false;
|
||||
|
||||
// Within current group
|
||||
if (next_block_idx < (current_num_valid_groups + 1) * num_m_blocks * num_n_blocks)
|
||||
break;
|
||||
|
||||
// Move to check the next group
|
||||
current_k_cumsum += current_shape_k;
|
||||
current_sf_k_cumsum += ceil_div(current_shape_k, SF_K_ALIGNMENT);
|
||||
current_num_valid_groups ++;
|
||||
|
||||
current_group_idx = next_group_idx ++;
|
||||
current_shape_k = next_shape_k;
|
||||
get_next_k_group(next_group_idx, next_shape_k);
|
||||
}
|
||||
|
||||
get_swizzled_block_idx(next_block_idx - current_num_valid_groups * num_m_blocks * num_n_blocks, m_block_idx, n_block_idx);
|
||||
} else {
|
||||
if (next_block_idx >= num_blocks)
|
||||
return false;
|
||||
|
||||
// For SM90 only
|
||||
// NOTES: we don't have to set `is_peer_cta_alive` for masked grouped GEMM, as it must be aligned
|
||||
is_peer_cta_alive = num_n_blocks % kNumMulticast == 0 or // Always aligned on N (constant bypass)
|
||||
num_m_blocks % kNumMulticast == 0 or // Always aligned on M (constant bypass)
|
||||
(next_block_idx ^ 1) < num_blocks; // Peer CTA in bound
|
||||
get_swizzled_block_idx(next_block_idx, m_block_idx, n_block_idx);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
// For SM90 only
|
||||
__device__ __forceinline__ bool is_tma_multicast_valid(const uint32_t& m_block_idx) const {
|
||||
if (num_blocks_in_group == 1)
|
||||
return false;
|
||||
if constexpr (kGemmType == GemmType::Normal or kGemmType == GemmType::MGroupedMasked or kGemmType == GemmType::KGroupedContiguous) {
|
||||
return true;
|
||||
} else {
|
||||
DG_STATIC_ASSERT(kGemmType == GemmType::MGroupedContiguous, "Invalid Gemm type");
|
||||
if constexpr (kIsMulticastOnA) {
|
||||
return true;
|
||||
} else {
|
||||
const auto& group_idx = __ldg(grouped_layout + m_block_idx * BLOCK_M);
|
||||
const auto& peer_group_idx = __ldg(grouped_layout + (m_block_idx ^ 1) * BLOCK_M);
|
||||
return group_idx == peer_group_idx;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// For SM90 only
|
||||
// ReSharper disable once CppNotAllPathsReturnValue
|
||||
__device__ __forceinline__ bool is_computation_valid(const uint32_t& m_block_idx, const uint32_t& m_offset) const {
|
||||
if constexpr (kGemmType == GemmType::Normal) {
|
||||
return true;
|
||||
} else if constexpr (kGemmType == GemmType::MGroupedContiguous) {
|
||||
return __ldg(grouped_layout + m_offset + m_block_idx * BLOCK_M) >= 0;
|
||||
} else if constexpr (kGemmType == GemmType::MGroupedMasked) {
|
||||
return m_offset + m_block_idx * BLOCK_M < __ldg(grouped_layout + current_group_idx);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
#pragma clang diagnostic pop
|
||||
|
||||
} // namespace deep_gemm
|
||||
+260
@@ -0,0 +1,260 @@
|
||||
#pragma once
|
||||
|
||||
#include <cute/atom/mma_traits_sm100.hpp>
|
||||
#include <cute/arch/mma_sm100_umma.hpp>
|
||||
#include <cute/arch/tmem_allocator_sm100.hpp>
|
||||
|
||||
#include <deep_gemm/common/utils.cuh>
|
||||
|
||||
namespace deep_gemm::sm100 {
|
||||
|
||||
template <uint32_t BLOCK_INNER, uint32_t kSwizzleMode, typename dtype_t>
|
||||
constexpr uint32_t get_inner_block_atom_size() {
|
||||
return kSwizzleMode == 0 ? BLOCK_INNER : kSwizzleMode / sizeof(dtype_t);
|
||||
}
|
||||
|
||||
template <uint32_t BLOCK_INNER, uint32_t BLOCK_OUTER,
|
||||
uint32_t kSwizzleMode, uint32_t kNumMulticast,
|
||||
typename dtype_t>
|
||||
__device__ __forceinline__ void
|
||||
tma_copy(void const* desc_ptr, cutlass::arch::ClusterTransactionBarrier* barrier_ptr,
|
||||
dtype_t* smem_ptr, const uint32_t& inner_idx, const int32_t& outer_idx) {
|
||||
DG_STATIC_ASSERT(1 <= kNumMulticast and kNumMulticast <= 2, "Invalid multicast config");
|
||||
DG_STATIC_ASSERT(static_cast<uint64_t>(cute::TMA::CacheHintSm90::EVICT_NORMAL) ==
|
||||
static_cast<uint64_t>(cute::TMA::CacheHintSm100::EVICT_NORMAL), "Invalid cache hint");
|
||||
|
||||
// 2-CTA function will send signals to the leader CTA only
|
||||
const auto copy_func = kNumMulticast == 1 ? cute::SM90_TMA_LOAD_2D::copy : cute::SM100_TMA_2SM_LOAD_2D::copy;
|
||||
|
||||
// Issue multiple TMAs
|
||||
constexpr uint32_t BLOCK_INNER_ATOM = get_inner_block_atom_size<BLOCK_INNER, kSwizzleMode, dtype_t>();
|
||||
#pragma unroll
|
||||
for (uint32_t i = 0; i < BLOCK_INNER / BLOCK_INNER_ATOM; ++ i) {
|
||||
copy_func(desc_ptr, reinterpret_cast<uint64_t*>(barrier_ptr),
|
||||
static_cast<uint64_t>(cute::TMA::CacheHintSm100::EVICT_NORMAL),
|
||||
smem_ptr + i * BLOCK_OUTER * BLOCK_INNER_ATOM, inner_idx + i * BLOCK_INNER_ATOM, outer_idx);
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
cute::UMMA::SmemDescriptor make_smem_desc(cute::UMMA::LayoutType layout, void* smem_ptr,
|
||||
uint32_t stride_byte_offset, uint32_t leading_byte_offset) {
|
||||
cute::UMMA::SmemDescriptor desc;
|
||||
|
||||
// Set the version for SM100
|
||||
desc.version_ = 1;
|
||||
|
||||
// Legacy mode
|
||||
desc.lbo_mode_ = 0;
|
||||
|
||||
// Layout
|
||||
desc.layout_type_ = static_cast<uint8_t>(layout);
|
||||
|
||||
// Start address
|
||||
const auto uint_ptr = cute::cast_smem_ptr_to_uint(smem_ptr);
|
||||
desc.start_address_ = static_cast<uint16_t>(uint_ptr >> 4);
|
||||
|
||||
// Base offset
|
||||
desc.base_offset_ = 0;
|
||||
|
||||
// SBO and LBO
|
||||
desc.stride_byte_offset_ = stride_byte_offset >> 4;
|
||||
desc.leading_byte_offset_ = leading_byte_offset >> 4;
|
||||
|
||||
return desc;
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
cute::UMMA::SmemDescriptor make_sf_desc(void* smem_ptr) {
|
||||
// NOTES: the UTCCP layout is K-major by default
|
||||
// Atom size: 8 x 128 bits
|
||||
// {SBO, LBO} means the byte stride between atoms on {MN, K}
|
||||
// Since the UTCCP we used is 128b-wide (only 1 atom on K), so LBO can be zero
|
||||
return make_smem_desc(cute::UMMA::LayoutType::SWIZZLE_NONE, smem_ptr, 8 * 16, 0);
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void replace_smem_desc_addr(cute::UMMA::SmemDescriptor& desc, const void* smem_ptr) {
|
||||
const auto uint_ptr = cute::cast_smem_ptr_to_uint(smem_ptr);
|
||||
desc.start_address_ = static_cast<uint16_t>(uint_ptr >> 4);
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
static uint32_t get_atom_base(const cute::UMMA::LayoutType& layout_type) {
|
||||
return layout_type == cute::UMMA::LayoutType::SWIZZLE_128B_BASE32B ? 32 : 16;
|
||||
}
|
||||
|
||||
// ReSharper disable once CppNotAllPathsReturnValue
|
||||
template <cute::UMMA::Major kMajorMode, uint32_t kSwizzleMode, bool kUseBase32, typename dtype_t>
|
||||
constexpr static cute::UMMA::LayoutType to_umma_layout_type() {
|
||||
DG_STATIC_ASSERT(kSwizzleMode == 0 or kSwizzleMode == 16 or
|
||||
kSwizzleMode == 32 or kSwizzleMode == 64 or
|
||||
kSwizzleMode == 128, "Invalid swizzling mode");
|
||||
// A special case
|
||||
if constexpr ((cute::is_same_v<dtype_t, float> and kMajorMode == cute::UMMA::Major::MN) or kUseBase32) {
|
||||
DG_STATIC_ASSERT(kUseBase32, "Invalid swizzling base");
|
||||
return cute::UMMA::LayoutType::SWIZZLE_128B_BASE32B;
|
||||
}
|
||||
|
||||
// Normal cases
|
||||
if constexpr (kSwizzleMode == 0) return cute::UMMA::LayoutType::SWIZZLE_NONE;
|
||||
if constexpr (kSwizzleMode == 16) return cute::UMMA::LayoutType::SWIZZLE_NONE;
|
||||
if constexpr (kSwizzleMode == 32) return cute::UMMA::LayoutType::SWIZZLE_32B;
|
||||
if constexpr (kSwizzleMode == 64) return cute::UMMA::LayoutType::SWIZZLE_64B;
|
||||
if constexpr (kSwizzleMode == 128) return cute::UMMA::LayoutType::SWIZZLE_128B;
|
||||
}
|
||||
|
||||
template <cute::UMMA::Major kMajorMode, uint32_t BLOCK_MN, uint32_t kSwizzleMode, typename dtype_t>
|
||||
__device__ __forceinline__
|
||||
constexpr uint32_t get_umma_desc_stride_k() {
|
||||
return kMajorMode == cute::UMMA::Major::K ? 1 : get_inner_block_atom_size<BLOCK_MN, kSwizzleMode, dtype_t>();
|
||||
}
|
||||
|
||||
template <cute::UMMA::Major kMajorMode, uint32_t BLOCK_MN, uint32_t kSwizzleMode, typename dtype_t>
|
||||
__device__ __forceinline__
|
||||
uint32_t advance_umma_desc_lo(const uint32_t& base, const uint32_t& offset, const uint32_t& k_idx) {
|
||||
return base + (((offset + k_idx * get_umma_desc_stride_k<kMajorMode, BLOCK_MN, kSwizzleMode, dtype_t>()) * static_cast<uint32_t>(sizeof(dtype_t))) >> 4u);
|
||||
}
|
||||
|
||||
template <cute::UMMA::Major kMajorMode, uint32_t BLOCK_MN, uint32_t BLOCK_K, uint32_t kSwizzleMode, bool kUseBase32 = false, typename dtype_t>
|
||||
__device__ __forceinline__
|
||||
cute::UMMA::SmemDescriptor make_umma_desc(dtype_t* base_smem_ptr, uint32_t mn_idx, uint32_t k_idx) {
|
||||
const uint32_t stride_k = get_umma_desc_stride_k<kMajorMode, BLOCK_MN, kSwizzleMode, dtype_t>();
|
||||
const auto& layout_type = to_umma_layout_type<kMajorMode, kSwizzleMode, kUseBase32, dtype_t>();
|
||||
const auto& num_non_contiguous = 128 / get_atom_base(layout_type);
|
||||
if constexpr (kMajorMode == cute::UMMA::Major::K) {
|
||||
// NOTES: for K-major layout, the swizzle must be 128B (also, atom index must be 0), as `BLOCK_K` is always 128
|
||||
DG_STATIC_ASSERT(kSwizzleMode == BLOCK_K * sizeof(dtype_t), "Unexpected value");
|
||||
|
||||
// Atom size: 8 x `kSwizzleMode` (in bytes, on K)
|
||||
// {SBO, LBO} means the byte stride between atoms on {MN, K}
|
||||
// NOTES: on K, there is only 1 atom as asserted previously, so LBO can be 0
|
||||
const uint32_t stride_byte_offset = num_non_contiguous * BLOCK_K * sizeof(dtype_t);
|
||||
const uint32_t leading_byte_offset = 0;
|
||||
return make_smem_desc(layout_type,
|
||||
base_smem_ptr + mn_idx * BLOCK_K + k_idx * stride_k,
|
||||
stride_byte_offset, leading_byte_offset);
|
||||
} else {
|
||||
constexpr uint32_t BLOCK_MN_ATOM = get_inner_block_atom_size<BLOCK_MN, kSwizzleMode, dtype_t>();
|
||||
|
||||
// Must have no in-atom MN-idx
|
||||
// NOTES: no worries for the runtime assert, the `mn_idx` are constants at compilation time
|
||||
DG_DEVICE_ASSERT(mn_idx % BLOCK_MN_ATOM == 0);
|
||||
DG_STATIC_ASSERT(kSwizzleMode > 0, "Invalid swizzling");
|
||||
|
||||
// Atom size: `kSwizzleMode` (in bytes, on MN) x 8
|
||||
// NOTES: `kSwizzleMode == 16` mean non-swizzling but interleaving
|
||||
// {SBO, LBO} means the byte stride between atoms on {K, MN} for swizzling
|
||||
// {SBO, LBO} means the byte stride between atoms on {MN, K} for non-swizzling
|
||||
uint32_t stride_byte_offset = num_non_contiguous * BLOCK_MN_ATOM * sizeof(dtype_t);
|
||||
uint32_t leading_byte_offset = BLOCK_K * BLOCK_MN_ATOM * sizeof(dtype_t);
|
||||
if constexpr (kSwizzleMode == 16)
|
||||
swap(stride_byte_offset, leading_byte_offset);
|
||||
return make_smem_desc(layout_type,
|
||||
base_smem_ptr + mn_idx * BLOCK_K + k_idx * stride_k,
|
||||
stride_byte_offset, leading_byte_offset);
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
uint64_t make_runtime_instr_desc_with_sf_id(cute::UMMA::InstrDescriptorBlockScaled desc, const uint32_t& sf_id) {
|
||||
desc.a_sf_id_ = sf_id, desc.b_sf_id_ = sf_id;
|
||||
return static_cast<uint64_t>(static_cast<uint32_t>(desc)) << 32;
|
||||
}
|
||||
|
||||
template <uint32_t kNumCols>
|
||||
__device__ constexpr uint32_t get_num_aligned_tmem_cols() {
|
||||
DG_STATIC_ASSERT(kNumCols <= 512, "Too many tensor memory columns");
|
||||
if (kNumCols <= 32) return 32;
|
||||
if (kNumCols <= 64) return 64;
|
||||
if (kNumCols <= 128) return 128;
|
||||
if (kNumCols <= 256) return 256;
|
||||
return 512;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_before_thread_sync() {
|
||||
asm volatile("tcgen05.fence::before_thread_sync;");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_after_thread_sync() {
|
||||
asm volatile("tcgen05.fence::after_thread_sync;");
|
||||
}
|
||||
|
||||
// UMMA versions with relaxed assertions
|
||||
struct SM100_MMA_F16BF16_SS {
|
||||
__device__ static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scale_c,
|
||||
uint64_t const& desc) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.cta_group::1.kind::f16 [%0], %1, %2, %3, p; \n\t"
|
||||
"}\n"
|
||||
:: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c));
|
||||
}
|
||||
};
|
||||
|
||||
struct SM100_MMA_F16BF16_2x1SM_SS {
|
||||
__device__ static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scale_c,
|
||||
uint64_t const& desc) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.cta_group::2.kind::f16 [%0], %1, %2, %3, p; \n\t"
|
||||
"}\n"
|
||||
:: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c));
|
||||
}
|
||||
};
|
||||
|
||||
struct SM100_MMA_MXF8F6F4_SS {
|
||||
__device__ static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scale_c,
|
||||
uint64_t const& desc,
|
||||
uint32_t const& tmem_sfa,
|
||||
uint32_t const& tmem_sfb) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.cta_group::1.kind::mxf8f6f4.block_scale [%0], %1, %2, %3, [%5], [%6], p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c),
|
||||
"r"(tmem_sfa), "r"(tmem_sfb));
|
||||
}
|
||||
};
|
||||
|
||||
struct SM100_MMA_MXF8F6F4_2x1SM_SS {
|
||||
__device__ static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scale_c,
|
||||
uint64_t const& desc,
|
||||
uint32_t const& tmem_sfa,
|
||||
uint32_t const& tmem_sfb) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.cta_group::2.kind::mxf8f6f4.block_scale [%0], %1, %2, %3, [%5], [%6], p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c),
|
||||
"r"(tmem_sfa), "r"(tmem_sfb));
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace `deep_gemm::sm100`
|
||||
+283
@@ -0,0 +1,283 @@
|
||||
#pragma once
|
||||
|
||||
#include <cute/arch/copy_sm90_tma.hpp>
|
||||
#include <cute/arch/cluster_sm90.hpp>
|
||||
#include <cute/arch/mma_sm90_gmma.hpp>
|
||||
#include <cute/arch/mma_sm90_gmma_ext.hpp>
|
||||
|
||||
#include <deep_gemm/common/utils.cuh>
|
||||
|
||||
namespace deep_gemm::sm90 {
|
||||
|
||||
template <int N_, typename MMA>
|
||||
struct FP8MMA {
|
||||
|
||||
template <size_t ...Idx>
|
||||
__forceinline__ __device__ static void call_fma_impl(uint64_t const& desc_a, uint64_t const& desc_b, float* d, bool scale_d, cute::index_sequence<Idx...>) {
|
||||
using namespace cute::SM90::GMMA;
|
||||
MMA::fma(desc_a, desc_b, d[Idx]..., (scale_d ? ScaleOut::One : ScaleOut::Zero));
|
||||
}
|
||||
|
||||
__forceinline__ __device__ static void wgmma(uint64_t const& desc_a, uint64_t const& desc_b, float* d, bool scale_d) {
|
||||
call_fma_impl(desc_a, desc_b, d, scale_d, cute::make_index_sequence<N_/2>{});
|
||||
}
|
||||
|
||||
static constexpr int M = 64;
|
||||
static constexpr int N = N_;
|
||||
static constexpr int K = 32;
|
||||
static constexpr int kNumAccum = M * N / 128;
|
||||
};
|
||||
|
||||
template <int N>
|
||||
struct FP8MMASelector {
|
||||
|
||||
static constexpr auto select_mma() {
|
||||
using namespace cute::SM90::GMMA;
|
||||
if constexpr (N == 8) return MMA_64x8x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 16) return MMA_64x16x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 24) return MMA_64x24x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 32) return MMA_64x32x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 40) return MMA_64x40x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 48) return MMA_64x48x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 56) return MMA_64x56x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 64) return MMA_64x64x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 72) return MMA_64x72x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 80) return MMA_64x80x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 88) return MMA_64x88x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 96) return MMA_64x96x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 104) return MMA_64x104x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 112) return MMA_64x112x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 120) return MMA_64x120x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 128) return MMA_64x128x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 136) return MMA_64x136x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 144) return MMA_64x144x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 152) return MMA_64x152x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 160) return MMA_64x160x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 168) return MMA_64x168x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 176) return MMA_64x176x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 184) return MMA_64x184x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 192) return MMA_64x192x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 200) return MMA_64x200x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 208) return MMA_64x208x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 216) return MMA_64x216x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 224) return MMA_64x224x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 232) return MMA_64x232x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 240) return MMA_64x240x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 248) return MMA_64x248x32_F32E4M3E4M3_SS_TN();
|
||||
if constexpr (N == 256) return MMA_64x256x32_F32E4M3E4M3_SS_TN();
|
||||
}
|
||||
|
||||
static constexpr auto select_type() {
|
||||
return FP8MMA<N, decltype(select_mma())>();
|
||||
}
|
||||
|
||||
using type = decltype(select_type());
|
||||
};
|
||||
|
||||
template <int N_, typename MMA>
|
||||
struct BF16MMA {
|
||||
|
||||
template <size_t ...Idx>
|
||||
__forceinline__ __device__ static void call_fma_impl(uint64_t const& desc_a, uint64_t const& desc_b, float* d, bool scale_d, cute::index_sequence<Idx...>) {
|
||||
using namespace cute::SM90::GMMA;
|
||||
MMA::fma(desc_a, desc_b, d[Idx]..., (scale_d ? ScaleOut::One : ScaleOut::Zero));
|
||||
}
|
||||
|
||||
__forceinline__ __device__ static void wgmma(uint64_t const& desc_a, uint64_t const& desc_b, float* d, bool scale_d) {
|
||||
call_fma_impl(desc_a, desc_b, d, scale_d, cute::make_index_sequence<N_/2>{});
|
||||
}
|
||||
|
||||
static constexpr int M = 64;
|
||||
static constexpr int N = N_;
|
||||
static constexpr int K = 16;
|
||||
static constexpr int kNumAccum = M * N / 128;
|
||||
};
|
||||
|
||||
template <int N>
|
||||
struct BF16MMASelector {
|
||||
|
||||
static constexpr auto select_mma() {
|
||||
using namespace cute::SM90::GMMA;
|
||||
if constexpr (N == 8) return MMA_64x8x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 16) return MMA_64x16x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 24) return MMA_64x24x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 32) return MMA_64x32x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 40) return MMA_64x40x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 48) return MMA_64x48x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 56) return MMA_64x56x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 64) return MMA_64x64x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 72) return MMA_64x72x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 80) return MMA_64x80x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 88) return MMA_64x88x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 96) return MMA_64x96x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 104) return MMA_64x104x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 112) return MMA_64x112x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 120) return MMA_64x120x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 128) return MMA_64x128x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 136) return MMA_64x136x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 144) return MMA_64x144x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 152) return MMA_64x152x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 160) return MMA_64x160x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 168) return MMA_64x168x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 176) return MMA_64x176x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 184) return MMA_64x184x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 192) return MMA_64x192x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 200) return MMA_64x200x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 208) return MMA_64x208x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 216) return MMA_64x216x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 224) return MMA_64x224x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 232) return MMA_64x232x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 240) return MMA_64x240x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 248) return MMA_64x248x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
if constexpr (N == 256) return MMA_64x256x16_F32BF16BF16_SS<Major::K, Major::K>();
|
||||
}
|
||||
|
||||
static constexpr auto select_type() {
|
||||
return BF16MMA<N, decltype(select_mma())>();
|
||||
}
|
||||
|
||||
using type = decltype(select_type());
|
||||
};
|
||||
|
||||
|
||||
template <typename dtype_t>
|
||||
struct SM90_U32x2_STSM_N {
|
||||
__device__ __forceinline__ static void
|
||||
copy(dtype_t src_0, dtype_t src_1, void* smem_dst) {
|
||||
const uint32_t src[2] = {*reinterpret_cast<uint32_t*>(&src_0), *reinterpret_cast<uint32_t*>(&src_1)};
|
||||
asm volatile("stmatrix.sync.aligned.x2.m8n8.shared.b16 [%0], {%1, %2};\n"
|
||||
:: "l"(smem_dst), "r"(src[0]), "r"(src[1]));
|
||||
}
|
||||
};
|
||||
|
||||
struct SM90_U32x2_LDSM_N {
|
||||
__device__ __forceinline__ static void
|
||||
copy(uint32_t& dst_0, uint32_t& dst_1, void* smem_src) {
|
||||
asm volatile("ldmatrix.sync.aligned.x2.m8n8.shared.b16 {%0, %1}, [%2];\n"
|
||||
: "=r"(dst_0), "=r"(dst_1)
|
||||
: "l"(smem_src));
|
||||
}
|
||||
};
|
||||
|
||||
struct SM90_U32x4_LDSM_N {
|
||||
__device__ __forceinline__ static void
|
||||
copy(uint32_t& dst_0, uint32_t& dst_1, uint32_t& dst_2, uint32_t& dst_3, void* smem_src) {
|
||||
asm volatile("ldmatrix.sync.aligned.x4.m8n8.shared.b16 {%0, %1, %2, %3}, [%4];\n"
|
||||
: "=r"(dst_0), "=r"(dst_1), "=r"(dst_2), "=r"(dst_3)
|
||||
: "l"(smem_src));
|
||||
}
|
||||
};
|
||||
|
||||
__forceinline__ __device__ void warpgroup_arrive() {
|
||||
asm volatile("wgmma.fence.sync.aligned;\n" ::: "memory");
|
||||
}
|
||||
|
||||
__forceinline__ __device__ void warpgroup_commit_batch() {
|
||||
asm volatile("wgmma.commit_group.sync.aligned;\n" ::: "memory");
|
||||
}
|
||||
|
||||
__forceinline__ __device__ void warpgroup_fence_operand(float& reg) {
|
||||
asm volatile("" : "+f"(reg) :: "memory");
|
||||
}
|
||||
|
||||
template <int N>
|
||||
__forceinline__ __device__ void warpgroup_wait() {
|
||||
DG_STATIC_ASSERT(N >= 0 and N <= 7, "WGMMA wait: N must be in range [0, 7]");
|
||||
asm volatile("wgmma.wait_group.sync.aligned %0;\n" :: "n"(N) : "memory");
|
||||
}
|
||||
|
||||
// TODO: replace with CUTLASS solution
|
||||
union GmmaDescriptor {
|
||||
__host__ __device__ constexpr GmmaDescriptor() noexcept: desc_(0) {}
|
||||
|
||||
__host__ __device__ constexpr GmmaDescriptor(uint64_t desc) noexcept: desc_(desc) {}
|
||||
|
||||
__host__ __device__ constexpr GmmaDescriptor(GmmaDescriptor const &t) noexcept: desc_(t.desc_) {}
|
||||
|
||||
__host__ __device__ constexpr GmmaDescriptor(GmmaDescriptor &&t) noexcept: desc_(t.desc_) {}
|
||||
|
||||
__host__ __device__ constexpr GmmaDescriptor &operator=(GmmaDescriptor const &t) noexcept {
|
||||
desc_ = t.desc_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
__host__ __device__ constexpr GmmaDescriptor &operator=(GmmaDescriptor &&t) noexcept {
|
||||
desc_ = t.desc_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
uint64_t desc_;
|
||||
uint32_t reg32_[2];
|
||||
uint16_t reg16_[4];
|
||||
|
||||
struct {
|
||||
uint16_t start_address_: 14, : 2;
|
||||
uint16_t leading_byte_offset_: 14, : 2;
|
||||
uint16_t stride_byte_offset_: 14, : 2;
|
||||
uint8_t : 1, base_offset_: 3, : 4;
|
||||
uint8_t : 6, layout_type_: 2;
|
||||
} bitfield;
|
||||
|
||||
// Decay to an `uint64_t`
|
||||
__host__ __device__ constexpr operator uint64_t() const noexcept { return desc_; }
|
||||
};
|
||||
|
||||
template <class PointerType>
|
||||
__device__ GmmaDescriptor make_smem_desc(PointerType smem_ptr, const int& layout_type,
|
||||
const int& leading_byte_offset = 0,
|
||||
const int& stride_byte_offset = 1024) {
|
||||
GmmaDescriptor desc;
|
||||
const auto& uint_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(smem_ptr));
|
||||
desc.bitfield.start_address_ = uint_ptr >> 4;
|
||||
desc.bitfield.layout_type_ = layout_type;
|
||||
desc.bitfield.leading_byte_offset_ = leading_byte_offset >> 4;
|
||||
desc.bitfield.stride_byte_offset_ = stride_byte_offset >> 4;
|
||||
desc.bitfield.base_offset_ = 0;
|
||||
return desc;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void
|
||||
tma_copy(void const* desc_ptr, uint64_t* barrier_ptr, void* smem_ptr,
|
||||
const uint32_t& crd_0, const uint32_t& crd_1, const uint32_t& num_tma_multicast = 1) {
|
||||
constexpr auto cache_hint = static_cast<uint64_t>(cute::TMA::CacheHintSm90::EVICT_NORMAL);
|
||||
if (num_tma_multicast == 1) {
|
||||
cute::SM90_TMA_LOAD_2D::copy(desc_ptr, barrier_ptr, cache_hint, smem_ptr, crd_0, crd_1);
|
||||
} else if (cute::block_rank_in_cluster() == 0) {
|
||||
cute::SM90_TMA_LOAD_MULTICAST_2D::copy(desc_ptr, barrier_ptr, (1 << num_tma_multicast) - 1, cache_hint, smem_ptr, crd_0, crd_1);
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void
|
||||
tma_3d_copy(void const* desc_ptr, uint64_t* barrier_ptr, void* smem_ptr,
|
||||
const uint32_t& crd_0, const uint32_t& crd_1, const uint32_t& crd_2) {
|
||||
constexpr auto cache_hint = static_cast<uint64_t>(cute::TMA::CacheHintSm90::EVICT_NORMAL);
|
||||
cute::SM90_TMA_LOAD_3D::copy(desc_ptr, barrier_ptr, cache_hint, smem_ptr, crd_0, crd_1, crd_2);
|
||||
}
|
||||
|
||||
// Tensormap related
|
||||
__device__ __forceinline__ void tensor_map_release_cta() {
|
||||
asm volatile ("fence.proxy.tensormap::generic.release.cta;");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tensor_map_acquire_cta(const cute::TmaDescriptor* gmem_desc_ptr) {
|
||||
auto gmem_int_desc = reinterpret_cast<uint64_t>(gmem_desc_ptr);
|
||||
asm volatile ("fence.proxy.tensormap::generic.acquire.cta [%0], 128;" :: "l"(gmem_int_desc) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tensor_map_replace_global_addr_in_smem(cute::TmaDescriptor* smem_desc, const void* new_addr) {
|
||||
auto smem_int_desc = static_cast<uint32_t>(__cvta_generic_to_shared(smem_desc));
|
||||
const auto new_int64_addr = reinterpret_cast<uint64_t>(new_addr);
|
||||
asm volatile ("tensormap.replace.tile.global_address.shared::cta.b1024.b64 [%0], %1;" :: "r"(smem_int_desc), "l"(new_int64_addr));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tensor_map_replace_global_inner_dim_stride_in_smem(cute::TmaDescriptor* smem_desc, const uint32_t& new_dim, const uint64_t& new_stride) {
|
||||
auto smem_int_desc = __cvta_generic_to_shared(smem_desc);
|
||||
asm volatile ("tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 0, %1;" :: "l"(smem_int_desc), "r"(new_dim));
|
||||
#if ((__CUDACC_VER_MAJOR__ > 12) or ((__CUDACC_VER_MAJOR__ == 12) and (__CUDACC_VER_MINOR__ >= 3)))
|
||||
asm volatile("tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], 0, %1;" :: "l"(smem_int_desc), "l"(new_stride));
|
||||
#else
|
||||
DG_STATIC_ASSERT(false, "Invalid CUDA version");
|
||||
#endif
|
||||
}
|
||||
|
||||
} // namespace `deep_gemm::sm90`
|
||||
+18
@@ -0,0 +1,18 @@
|
||||
#pragma once
|
||||
|
||||
namespace deep_gemm {
|
||||
|
||||
enum class GemmType {
|
||||
Normal = 0,
|
||||
MGroupedContiguous = 1,
|
||||
MGroupedMasked = 2,
|
||||
KGroupedContiguous = 3,
|
||||
};
|
||||
|
||||
enum class KernelType {
|
||||
Kernel1D1D = 0,
|
||||
Kernel1D2D = 1,
|
||||
KernelNoSF = 2
|
||||
};
|
||||
|
||||
} // namespace deep_gemm
|
||||
+179
@@ -0,0 +1,179 @@
|
||||
#pragma once
|
||||
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp8.h>
|
||||
#include <cuda/std/cstdint>
|
||||
#include <cuda/std/utility>
|
||||
#include <cute/container/tuple.hpp>
|
||||
|
||||
#include "cute_tie.cuh"
|
||||
|
||||
#ifdef __CLION_IDE__
|
||||
|
||||
__host__ __device__ __forceinline__ void host_device_printf(const char* format, ...) {
|
||||
asm volatile("trap;");
|
||||
}
|
||||
|
||||
#define printf host_device_printf
|
||||
#endif
|
||||
|
||||
#ifndef DG_DEVICE_ASSERT
|
||||
#define DG_DEVICE_ASSERT(cond) \
|
||||
do { \
|
||||
if (not (cond)) { \
|
||||
printf("Assertion failed: %s:%d, condition: %s\n", __FILE__, __LINE__, #cond); \
|
||||
asm("trap;"); \
|
||||
} \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
#ifndef DG_TRAP_ONLY_DEVICE_ASSERT
|
||||
#define DG_TRAP_ONLY_DEVICE_ASSERT(cond) \
|
||||
do { \
|
||||
if (not (cond)) \
|
||||
asm("trap;"); \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
#ifndef DG_STATIC_ASSERT
|
||||
#define DG_STATIC_ASSERT(cond, ...) static_assert(cond, __VA_ARGS__)
|
||||
#endif
|
||||
|
||||
namespace deep_gemm {
|
||||
|
||||
template <typename FuncT>
|
||||
struct PatternVisitor {
|
||||
FuncT func;
|
||||
|
||||
__device__ __host__
|
||||
explicit PatternVisitor(FuncT&& func): func(std::forward<FuncT>(func)) {}
|
||||
|
||||
__device__ __host__
|
||||
auto operator [](const uint32_t& i) {
|
||||
return func(i);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
__device__ __host__ T ceil_div(T a, T b) {
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__device__ __host__ constexpr T constexpr_ceil_div(T a, T b) {
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__device__ __host__ T align(T a, T b) {
|
||||
return ceil_div(a, b) * b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__device__ __host__ constexpr T constexpr_align(T a, T b) {
|
||||
return constexpr_ceil_div(a, b) * b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__device__ __host__ constexpr T constexpr_gcd(T a, T b) {
|
||||
return b == 0 ? a : constexpr_gcd(b, a % b);
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
__forceinline__ __device__ void swap(T& a, T& b) {
|
||||
T temp = a;
|
||||
a = b;
|
||||
b = temp;
|
||||
}
|
||||
|
||||
__forceinline__ __device__ uint32_t get_sm_idx() {
|
||||
uint32_t sm_idx;
|
||||
asm ("mov.u32 %0, %%smid;" : "=r"(sm_idx));
|
||||
return sm_idx;
|
||||
}
|
||||
|
||||
__forceinline__ __device__ uint32_t get_lane_idx() {
|
||||
uint32_t lane_id;
|
||||
asm ("mov.u32 %0, %laneid;" : "=r"(lane_id));
|
||||
return lane_id;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ uint32_t ld_shared(const uint32_t* ptr) {
|
||||
uint32_t ret;
|
||||
asm volatile("ld.shared.u32 %0, [%1];" : "=r"(ret) : "l"(ptr));
|
||||
return ret;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float2 ld_shared(const float2* ptr) {
|
||||
float2 ret;
|
||||
asm volatile("ld.shared.v2.f32 {%0, %1}, [%2];" : "=f"(ret.x), "=f"(ret.y) : "l"(ptr));
|
||||
return ret;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float4 ld_shared(const float4* ptr) {
|
||||
float4 ret;
|
||||
asm volatile("ld.shared.v4.f32 {%0, %1, %2, %3}, [%4];" : "=f"(ret.x), "=f"(ret.y), "=f"(ret.z), "=f"(ret.w) : "l"(ptr));
|
||||
return ret;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ uint4 ld_shared(const uint4* ptr) {
|
||||
uint4 ret;
|
||||
asm volatile("ld.shared.v4.u32 {%0, %1, %2, %3}, [%4];" : "=r"(ret.x), "=r"(ret.y), "=r"(ret.z), "=r"(ret.w) : "l"(ptr));
|
||||
return ret;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float ld_shared(const float* ptr) {
|
||||
float ret;
|
||||
asm volatile("ld.shared.f32 %0, [%1];" : "=f"(ret) : "l"(ptr));
|
||||
return ret;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void st_shared(const float* ptr, float val) {
|
||||
asm volatile("st.shared.f32 [%0], %1;" :: "l"(ptr), "f"(val));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void st_shared(const float2* ptr, float2 val) {
|
||||
asm volatile("st.shared.v2.f32 [%0], {%1, %2};" :: "l"(ptr), "f"(val.x), "f"(val.y));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void st_shared(const uint32_t* ptr, uint32_t val) {
|
||||
asm volatile("st.shared.u32 [%0], %1;" :: "l"(ptr), "r"(val));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void st_shared(const void* ptr, uint32_t x, uint32_t y) {
|
||||
asm volatile("st.shared.v2.u32 [%0], {%1, %2};" :: "l"(ptr), "r"(x), "r"(y));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void st_shared(const void* ptr, uint32_t x, uint32_t y, uint32_t z, uint32_t w) {
|
||||
asm volatile("st.shared.v4.u32 [%0], {%1, %2, %3, %4};" :: "l"(ptr), "r"(x), "r"(y), "r"(z), "r"(w));
|
||||
}
|
||||
|
||||
template <typename old_t>
|
||||
__device__ __forceinline__ int cast_into_bf16_and_pack(old_t& x, old_t& y) {
|
||||
auto bf16x2 = __float22bfloat162_rn({*reinterpret_cast<float*>(&x), *reinterpret_cast<float*>(&y)});
|
||||
return *reinterpret_cast<int*>(&bf16x2);
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void prefetch_l1(void *ptr) {
|
||||
asm volatile("prefetch.global.L1 [%0];" :: "l"(ptr));
|
||||
}
|
||||
|
||||
template <uint32_t kNumBytes>
|
||||
struct Vectorized {
|
||||
static auto zeros() {
|
||||
// TODO: add `ulonglong4` for SM100 once `__ldg` support this
|
||||
if constexpr (kNumBytes > 0 and kNumBytes % 16 == 0) {
|
||||
return make_uint4(0, 0, 0, 0);
|
||||
} else if constexpr (kNumBytes > 0 and kNumBytes % 8 == 0) {
|
||||
return make_uint2(0, 0);
|
||||
} else if constexpr (kNumBytes > 0 and kNumBytes % 4 == 0) {
|
||||
return 0;
|
||||
} else {
|
||||
DG_STATIC_ASSERT(kNumBytes > 0 and kNumBytes % 4 == 0, "Invalid vectorization");
|
||||
}
|
||||
}
|
||||
|
||||
using vec_t = decltype(zeros());
|
||||
};
|
||||
|
||||
} // namespace `deep_gemm`
|
||||
+408
@@ -0,0 +1,408 @@
|
||||
#pragma once
|
||||
|
||||
#pragma clang diagnostic push
|
||||
#pragma clang diagnostic ignored "-Wunknown-attributes"
|
||||
|
||||
#include <cutlass/arch/barrier.h>
|
||||
#include <cutlass/arch/reg_reconfig.h>
|
||||
|
||||
#include <cute/arch/cluster_sm90.hpp>
|
||||
#include <cute/arch/copy_sm90_desc.hpp>
|
||||
#include <cute/arch/copy_sm90_tma.hpp>
|
||||
|
||||
#include <deep_gemm/common/epilogue_utils.cuh>
|
||||
#include <deep_gemm/common/utils.cuh>
|
||||
#include <deep_gemm/common/scheduler.cuh>
|
||||
#include <deep_gemm/common/sm90_utils.cuh>
|
||||
|
||||
namespace deep_gemm {
|
||||
|
||||
using namespace deep_gemm::sm90;
|
||||
|
||||
template <uint32_t kNumFormerIters, uint32_t kGap, uint32_t kEnd, typename func_t>
|
||||
__device__ void dispatch_num_former_iters(uint32_t num_former_iters, const func_t& func) {
|
||||
if (num_former_iters == kNumFormerIters) {
|
||||
func(cute::Int<kNumFormerIters>{});
|
||||
return;
|
||||
}
|
||||
|
||||
if constexpr (kNumFormerIters + kGap <= kEnd)
|
||||
dispatch_num_former_iters<kNumFormerIters + kGap, kGap, kEnd>(num_former_iters, func);
|
||||
}
|
||||
|
||||
template <uint32_t SHAPE_M, uint32_t SHAPE_N, uint32_t SHAPE_K,
|
||||
uint32_t kNumGroups,
|
||||
uint32_t BLOCK_M, uint32_t BLOCK_N, uint32_t BLOCK_K,
|
||||
uint32_t kSwizzleDMode,
|
||||
uint32_t kNumStages, uint32_t kNumLastStages,
|
||||
uint32_t kNumTMAThreads, uint32_t kNumMathThreads,
|
||||
uint32_t kNumTMAMulticast, bool kIsTMAMulticastOnA,
|
||||
uint32_t kNumSMs, GemmType kGemmType,
|
||||
typename epilogue_type_t>
|
||||
__global__ __launch_bounds__(kNumTMAThreads + kNumMathThreads, 1) void
|
||||
sm90_fp8_gemm_1d2d_impl(float* sfb, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const __grid_constant__ cute::TmaDescriptor tensor_map_a,
|
||||
const __grid_constant__ cute::TmaDescriptor tensor_map_b,
|
||||
const __grid_constant__ cute::TmaDescriptor tensor_map_d,
|
||||
const __grid_constant__ cute::TmaDescriptor tensor_map_sfa) {
|
||||
#if (defined(__CUDA_ARCH__) and (__CUDA_ARCH__ >= 900)) or defined(__CLION_IDE__)
|
||||
// Scaling checks
|
||||
DG_STATIC_ASSERT(BLOCK_K == 128, "Only support per-128-channel FP8 scaling");
|
||||
DG_STATIC_ASSERT(constexpr_ceil_div(BLOCK_N, BLOCK_K) == 1 or (constexpr_gcd(BLOCK_N, BLOCK_K) == BLOCK_N - BLOCK_K), "Too much B scales in a single block");
|
||||
|
||||
// Types
|
||||
using WGMMA = typename FP8MMASelector<BLOCK_N>::type;
|
||||
using Barrier = cutlass::arch::ClusterTransactionBarrier;
|
||||
DG_STATIC_ASSERT(BLOCK_M % WGMMA::M == 0, "Invalid block size");
|
||||
|
||||
// Overwrite shape constants if the compiler gives
|
||||
shape_m = SHAPE_M != 0 ? SHAPE_M : shape_m;
|
||||
shape_n = SHAPE_N != 0 ? SHAPE_N : shape_n;
|
||||
shape_k = SHAPE_K != 0 ? SHAPE_K : shape_k;
|
||||
|
||||
// Shared memory
|
||||
static constexpr bool kMustUseUniformedScaleB = (BLOCK_K % BLOCK_N == 0);
|
||||
static constexpr uint32_t SMEM_D_SIZE = BLOCK_M * BLOCK_N * sizeof(__nv_bfloat16);
|
||||
static constexpr uint32_t SMEM_A_SIZE_PER_STAGE = BLOCK_M * BLOCK_K * sizeof(__nv_fp8_e4m3);
|
||||
static constexpr uint32_t SMEM_B_SIZE_PER_STAGE = BLOCK_N * BLOCK_K * sizeof(__nv_fp8_e4m3);
|
||||
static constexpr uint32_t SMEM_SFA_SIZE_PER_STAGE = BLOCK_M * sizeof(float);
|
||||
const uint32_t& shape_k_scales = ceil_div(shape_k, BLOCK_K);
|
||||
const uint32_t& smem_sfb_size = align<uint32_t>(shape_k_scales * (kMustUseUniformedScaleB ? 1 : 2) * sizeof(float), sizeof(Barrier));
|
||||
|
||||
// Configs
|
||||
const uint32_t num_total_k_blocks = ceil_div(shape_k, BLOCK_K);
|
||||
const uint32_t warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
const uint32_t lane_idx = get_lane_idx();
|
||||
|
||||
// Prefetch TMA descriptors at the very beginning
|
||||
if (warp_idx == kNumMathThreads / 32 and cute::elect_one_sync()) {
|
||||
cute::prefetch_tma_descriptor(&tensor_map_a);
|
||||
cute::prefetch_tma_descriptor(&tensor_map_b);
|
||||
cute::prefetch_tma_descriptor(&tensor_map_sfa);
|
||||
cute::prefetch_tma_descriptor(&tensor_map_d);
|
||||
}
|
||||
__syncwarp();
|
||||
|
||||
// Align to 1024 bytes for swizzle-128B
|
||||
extern __shared__ __align__(1024) uint8_t smem_buffer[];
|
||||
DG_STATIC_ASSERT(SMEM_D_SIZE % 1024 == 0, "Shared memory of A/B must be aligned to 1024 bytes");
|
||||
|
||||
// Data on shared memory
|
||||
auto smem_d = reinterpret_cast<__nv_bfloat16*>(smem_buffer);
|
||||
auto smem_a = PatternVisitor([&](const uint32_t& i) {
|
||||
return reinterpret_cast<__nv_fp8_e4m3*>(smem_buffer + SMEM_D_SIZE + i * SMEM_A_SIZE_PER_STAGE);
|
||||
});
|
||||
auto smem_b = PatternVisitor([&](const uint32_t& i) {
|
||||
return reinterpret_cast<__nv_fp8_e4m3*>(smem_buffer + SMEM_D_SIZE + kNumStages * SMEM_A_SIZE_PER_STAGE + i * SMEM_B_SIZE_PER_STAGE);
|
||||
});
|
||||
constexpr uint32_t SMEM_SF_OFFSET = SMEM_D_SIZE + kNumStages * (SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE);
|
||||
auto smem_sfa = PatternVisitor([&](const uint32_t& i) {
|
||||
return reinterpret_cast<float*>(smem_buffer + SMEM_SF_OFFSET + i * SMEM_SFA_SIZE_PER_STAGE);
|
||||
});
|
||||
auto smem_sfb = reinterpret_cast<float*>(smem_buffer + SMEM_SF_OFFSET + kNumStages * SMEM_SFA_SIZE_PER_STAGE);
|
||||
|
||||
// Fill barriers
|
||||
auto barrier_start_ptr = reinterpret_cast<Barrier*>(reinterpret_cast<uint8_t*>(smem_sfb) + smem_sfb_size);
|
||||
auto full_barriers = PatternVisitor([&](const uint32_t& i) { return barrier_start_ptr + i; });
|
||||
auto empty_barriers = PatternVisitor([&](const uint32_t& i) { return barrier_start_ptr + kNumStages + i; });
|
||||
|
||||
// Initialize barriers
|
||||
DG_STATIC_ASSERT(kNumTMAMulticast <= 32, "Too many TMA multicast");
|
||||
if (warp_idx == kNumMathThreads / 32 + 1 and cute::elect_one_sync()) {
|
||||
// NOTES: we always use `lane_idx` to arrive for the `lane_idx`-th CTA in the cluster,
|
||||
// even with TMA multicast disabled, we want to make the behavior aligned
|
||||
#pragma unroll
|
||||
for (uint32_t i = 0; i < kNumStages; ++ i) {
|
||||
full_barriers[i]->init(1);
|
||||
empty_barriers[i]->init(kNumTMAMulticast * kNumMathThreads / 32);
|
||||
}
|
||||
|
||||
// Make initialized barrier visible in async proxy
|
||||
cutlass::arch::fence_barrier_init();
|
||||
}
|
||||
|
||||
// Synchronize all threads to make barrier visible in normal memory model
|
||||
(kNumTMAMulticast > 1) ? cute::cluster_sync() : __syncthreads();
|
||||
|
||||
// Register reconfigurations
|
||||
constexpr uint32_t kNumTMARegisters = 40;
|
||||
constexpr uint32_t kNumMathRegisters = 232;
|
||||
|
||||
// Block scheduler
|
||||
uint32_t m_block_idx, n_block_idx;
|
||||
auto scheduler = Scheduler<kGemmType, BLOCK_M, BLOCK_N, kNumGroups, kNumTMAMulticast, kIsTMAMulticastOnA, kNumSMs>(shape_m, shape_n, shape_k, grouped_layout);
|
||||
|
||||
// Pipeline and TMA phases
|
||||
uint32_t stage_idx = 0, phase = 0;
|
||||
auto advance_pipeline = [&](uint32_t& k_block_idx) {
|
||||
++ k_block_idx;
|
||||
|
||||
// Flip phases only if reach the next first stage
|
||||
stage_idx = stage_idx == kNumStages - 1 ? 0 : stage_idx + 1;
|
||||
phase ^= stage_idx == 0;
|
||||
};
|
||||
|
||||
if (warp_idx >= kNumMathThreads / 32) {
|
||||
// TMA warp-group for loading data
|
||||
cutlass::arch::warpgroup_reg_dealloc<kNumTMARegisters>();
|
||||
|
||||
// NOTES: only one thread (or warp) will be used
|
||||
if (warp_idx == kNumMathThreads / 32 and cute::elect_one_sync()) {
|
||||
// Persistently schedule over blocks
|
||||
while (scheduler.get_next_block(m_block_idx, n_block_idx)) {
|
||||
// Assign TMA multicast number into A and B
|
||||
// NOTES: there may be additional odd rows/columns or cases where multicast is not possible.
|
||||
const bool is_tma_multicast_valid = scheduler.is_tma_multicast_valid(m_block_idx);
|
||||
const uint32_t num_tma_multicast_a = (kIsTMAMulticastOnA and is_tma_multicast_valid) ? kNumTMAMulticast : 1;
|
||||
const uint32_t num_tma_multicast_b = (not kIsTMAMulticastOnA and is_tma_multicast_valid) ? kNumTMAMulticast : 1;
|
||||
DG_STATIC_ASSERT(kNumTMAMulticast <= 2, "Scheduler does not support > 2 TMA multicast");
|
||||
|
||||
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
|
||||
// Wait consumer release
|
||||
empty_barriers[stage_idx]->wait(phase ^ 1);
|
||||
|
||||
// Issue TMA A
|
||||
constexpr bool kWithGroupOffsetA = kGemmType == GemmType::MGroupedMasked;
|
||||
auto& full_barrier = *full_barriers[stage_idx];
|
||||
const uint32_t k_idx = k_block_idx * BLOCK_K;
|
||||
tma_copy(&tensor_map_a, reinterpret_cast<uint64_t*>(&full_barrier),
|
||||
smem_a[stage_idx], k_idx, scheduler.get_global_idx<kWithGroupOffsetA>(shape_m, BLOCK_M, m_block_idx),
|
||||
num_tma_multicast_a);
|
||||
tma_copy(&tensor_map_sfa, reinterpret_cast<uint64_t*>(&full_barrier),
|
||||
smem_sfa[stage_idx], m_block_idx * BLOCK_M, scheduler.get_global_idx<kWithGroupOffsetA>(shape_k_scales, 1, k_block_idx),
|
||||
num_tma_multicast_a);
|
||||
|
||||
// Issue TMA B
|
||||
tma_copy(&tensor_map_b, reinterpret_cast<uint64_t*>(&full_barrier),
|
||||
smem_b[stage_idx], k_idx, scheduler.get_global_idx<true>(shape_n, BLOCK_N, n_block_idx, m_block_idx),
|
||||
num_tma_multicast_b);
|
||||
full_barrier.arrive_and_expect_tx(SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE + SMEM_SFA_SIZE_PER_STAGE);
|
||||
}
|
||||
}
|
||||
|
||||
// To safely deconstruct distributed shared barriers, we need another round of empty waits
|
||||
if constexpr (kNumTMAMulticast > 1) {
|
||||
for (uint32_t i = 0; i < kNumStages; advance_pipeline(i))
|
||||
empty_barriers[stage_idx]->wait(phase ^ 1);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Math warp-groups for WGMMA
|
||||
cutlass::arch::warpgroup_reg_alloc<kNumMathRegisters>();
|
||||
|
||||
// NOTES: use `__shfl_sync` to encourage NVCC to use unified registers
|
||||
const auto math_wg_idx = __shfl_sync(0xffffffff, threadIdx.x / 128, 0);
|
||||
const auto r_0 = warp_idx * 16 + lane_idx / 4, r_1 = r_0 + 8;
|
||||
|
||||
auto a_desc = make_smem_desc(smem_a[0] + math_wg_idx * WGMMA::M * BLOCK_K, 1);
|
||||
auto b_desc = make_smem_desc(smem_b[0], 1);
|
||||
const uint32_t a_desc_lo = __shfl_sync(0xffffffff, a_desc.reg32_[0], 0);
|
||||
const uint32_t b_desc_lo = __shfl_sync(0xffffffff, b_desc.reg32_[0], 0);
|
||||
|
||||
// Persistently schedule over blocks
|
||||
while (scheduler.get_next_block(m_block_idx, n_block_idx)) {
|
||||
// Decide the number of scales B to load
|
||||
DG_TRAP_ONLY_DEVICE_ASSERT(shape_n % 8 == 0);
|
||||
uint32_t num_former_iters = BLOCK_N / 8, num_full_iters = num_former_iters;
|
||||
if constexpr (not kMustUseUniformedScaleB) {
|
||||
num_former_iters = min(BLOCK_N, BLOCK_K - n_block_idx * BLOCK_N % BLOCK_K) / 8;
|
||||
num_full_iters = min(shape_n - n_block_idx * BLOCK_N, BLOCK_N) / 8;
|
||||
}
|
||||
uint32_t num_sfb = shape_k_scales * (num_former_iters >= num_full_iters ? 1 : 2);
|
||||
|
||||
// Load B scales with math warp-groups
|
||||
// NOTES: except the first warp, we want to overlap loading B scales with TMA stores between tasks
|
||||
if (threadIdx.x >= 32) {
|
||||
auto num_previous_lines = scheduler.get_global_idx<true>(ceil_div(shape_n, BLOCK_K), 0, 0, m_block_idx);
|
||||
auto local_sfb = sfb + (num_previous_lines + ((n_block_idx * BLOCK_N) / BLOCK_K)) * shape_k_scales;
|
||||
#pragma unroll
|
||||
for (uint32_t i = threadIdx.x - 32; i < num_sfb; i += kNumMathThreads - 32)
|
||||
st_shared(smem_sfb + i, __ldg(local_sfb + i));
|
||||
}
|
||||
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
|
||||
|
||||
// Accumulation for WGMMA or CUDA promotion
|
||||
constexpr uint32_t WAVE_BLOCK_M = WGMMA::M * (BLOCK_M <= 64 ? 1 : 2);
|
||||
DG_STATIC_ASSERT(BLOCK_M % WAVE_BLOCK_M == 0, "Invalid block sizes");
|
||||
float accum[WGMMA::kNumAccum], final_accum[WGMMA::kNumAccum * (BLOCK_M / WAVE_BLOCK_M)] = {0};
|
||||
|
||||
// Empty barrier arrival
|
||||
auto empty_barrier_arrive = [&]() {
|
||||
if constexpr (kNumTMAMulticast == 1) {
|
||||
lane_idx == 0 ? empty_barriers[stage_idx]->arrive() : void();
|
||||
} else {
|
||||
auto target_cta = scheduler.is_peer_cta_alive ? lane_idx : cute::block_rank_in_cluster();
|
||||
lane_idx < kNumTMAMulticast ? empty_barriers[stage_idx]->arrive(target_cta) : void();
|
||||
}
|
||||
};
|
||||
|
||||
// Skip useless computations
|
||||
if (scheduler.is_computation_valid(m_block_idx, math_wg_idx * WGMMA::M)) {
|
||||
// The compiler must know the dynamic variable `num_former_iters`'s real value
|
||||
constexpr bool kShouldOptimize = BLOCK_K / constexpr_gcd(BLOCK_K, BLOCK_N) <= 4 and not kMustUseUniformedScaleB;
|
||||
constexpr uint32_t kGap = constexpr_gcd(BLOCK_K, BLOCK_N) / 8;
|
||||
constexpr uint32_t kEnd = kShouldOptimize ? BLOCK_K / 8 : 0;
|
||||
|
||||
// Dispatch `num_former_iters` and launch MMAs
|
||||
dispatch_num_former_iters<0, kGap, kEnd>(kShouldOptimize ? num_former_iters : 0, [&](auto _) {
|
||||
#pragma unroll 8
|
||||
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
|
||||
const auto& a_desc_base_lo = a_desc_lo + stage_idx * (SMEM_A_SIZE_PER_STAGE / 16);
|
||||
const auto& b_desc_base_lo = b_desc_lo + stage_idx * (SMEM_B_SIZE_PER_STAGE / 16);
|
||||
|
||||
// Read B scales
|
||||
float scale_b_0 = ld_shared(smem_sfb + k_block_idx), scale_b_1;
|
||||
// NOTES: even some blocks do not need to read the second row, but we still load one to align with other blocks
|
||||
if constexpr (not kMustUseUniformedScaleB)
|
||||
scale_b_1 = ld_shared(smem_sfb + k_block_idx + shape_k_scales);
|
||||
|
||||
// Wait TMA arrivals
|
||||
full_barriers[stage_idx]->wait(phase);
|
||||
|
||||
// TODO: remove some useless computation for unaligned Ms
|
||||
#pragma unroll
|
||||
for (uint32_t local_idx = 0; local_idx < BLOCK_M / WAVE_BLOCK_M; ++ local_idx) {
|
||||
auto m_offset = local_idx * WAVE_BLOCK_M;
|
||||
|
||||
// Read A scales
|
||||
// NOTES: all shared memory read must be prior to `warpgroup_arrive` to avoid next scheduled block polluting the results
|
||||
auto scale_a_0 = ld_shared(smem_sfa[stage_idx] + r_0 + m_offset);
|
||||
auto scale_a_1 = ld_shared(smem_sfa[stage_idx] + r_1 + m_offset);
|
||||
|
||||
// Commit WGMMA instructions
|
||||
#pragma unroll
|
||||
for (uint32_t i = 0; i < WGMMA::kNumAccum; ++ i)
|
||||
warpgroup_fence_operand(accum[i]);
|
||||
warpgroup_arrive();
|
||||
#pragma unroll
|
||||
for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) {
|
||||
a_desc.reg32_[0] = a_desc_base_lo + (m_offset * BLOCK_K + k * WGMMA::K) / 16;
|
||||
b_desc.reg32_[0] = b_desc_base_lo + k * WGMMA::K / 16;
|
||||
WGMMA::wgmma(a_desc, b_desc, accum, k);
|
||||
}
|
||||
warpgroup_commit_batch();
|
||||
#pragma unroll
|
||||
for (uint32_t i = 0; i < WGMMA::kNumAccum; ++ i)
|
||||
warpgroup_fence_operand(accum[i]);
|
||||
warpgroup_wait<0>();
|
||||
|
||||
// Notify barrier arrival at the last warpgroup wave
|
||||
if (local_idx == BLOCK_M / WAVE_BLOCK_M - 1)
|
||||
empty_barrier_arrive();
|
||||
|
||||
// Promote with scales
|
||||
// NOTES: making it as predicates is very important for performance, comparing to two loops
|
||||
float scale_0_0 = scale_a_0 * scale_b_0, scale_1_0 = scale_a_1 * scale_b_0;
|
||||
float scale_0_1, scale_1_1;
|
||||
if constexpr (not kMustUseUniformedScaleB)
|
||||
scale_0_1 = scale_a_0 * scale_b_1, scale_1_1 = scale_a_1 * scale_b_1;
|
||||
|
||||
auto shifted_accum = final_accum + WGMMA::kNumAccum * local_idx;
|
||||
#pragma unroll
|
||||
for (uint32_t i = 0; i < WGMMA::kNumAccum / 4; ++ i) {
|
||||
// NOTES: for unrolled `num_former_iters` cases, we expect the compiler to automatically make it a constant
|
||||
bool predicate = kMustUseUniformedScaleB or i < num_former_iters;
|
||||
shifted_accum[i * 4 + 0] += (predicate ? scale_0_0 : scale_0_1) * accum[i * 4 + 0];
|
||||
shifted_accum[i * 4 + 1] += (predicate ? scale_0_0 : scale_0_1) * accum[i * 4 + 1];
|
||||
shifted_accum[i * 4 + 2] += (predicate ? scale_1_0 : scale_1_1) * accum[i * 4 + 2];
|
||||
shifted_accum[i * 4 + 3] += (predicate ? scale_1_0 : scale_1_1) * accum[i * 4 + 3];
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
} else {
|
||||
#pragma unroll
|
||||
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
|
||||
full_barriers[stage_idx]->wait(phase);
|
||||
empty_barrier_arrive();
|
||||
}
|
||||
}
|
||||
|
||||
// TMA checks
|
||||
constexpr uint32_t kNumElemBytes = sizeof(nv_bfloat16);
|
||||
constexpr uint32_t TMA_D_BLOCK_N = kSwizzleDMode == 0 ? BLOCK_N : (kSwizzleDMode / kNumElemBytes);
|
||||
constexpr uint32_t WGMMA_M_PER_WARP = WGMMA::M / 4;
|
||||
DG_STATIC_ASSERT(BLOCK_M % 8 == 0, "Invalid swizzling atom");
|
||||
DG_STATIC_ASSERT(BLOCK_N % TMA_D_BLOCK_N == 0 and BLOCK_N / TMA_D_BLOCK_N <= 32,
|
||||
"Unaligned TMA store or too many TMA store instructions");
|
||||
DG_STATIC_ASSERT(TMA_D_BLOCK_N % 8 == 0, "Invalid TMA block N");
|
||||
|
||||
// Wait last TMA store to be finished
|
||||
if (threadIdx.x < BLOCK_N / TMA_D_BLOCK_N)
|
||||
cute::tma_store_wait<0>();
|
||||
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
|
||||
|
||||
// Write back to shared memory using STSM and issue TMA stores
|
||||
DG_STATIC_ASSERT(WGMMA::kNumAccum % 4 == 0, "Invalid STSM x2 vectorization");
|
||||
#pragma unroll
|
||||
for (uint32_t local_idx = 0; local_idx < BLOCK_M / WAVE_BLOCK_M; ++ local_idx) {
|
||||
auto m_offset = local_idx * WAVE_BLOCK_M;
|
||||
auto shifted_accum = final_accum + WGMMA::kNumAccum * local_idx;
|
||||
#pragma unroll
|
||||
for (auto i = 0; i < WGMMA::kNumAccum / 4; ++ i) {
|
||||
// Swizzle or padding into the correct address
|
||||
uint8_t* smem_ptr = nullptr;
|
||||
if constexpr (kSwizzleDMode > 0) {
|
||||
// Calculate the swizzling atom offset and in-atom offset
|
||||
constexpr uint32_t kNumBankGroupBytes = 16;
|
||||
auto atom_offset = i / (TMA_D_BLOCK_N / 8), in_atom_offset = i % (TMA_D_BLOCK_N / 8);
|
||||
|
||||
// Calculate the index of the bank group to be written in the atom
|
||||
auto bank_group_index = in_atom_offset + lane_idx * (kSwizzleDMode / kNumBankGroupBytes);
|
||||
|
||||
// Reshape the atom in another view and swizzle
|
||||
// - original: `(BLOCK_M, kSwizzleDMode / kNumBankGroupBytes)`
|
||||
// - new: `(BLOCK_M * kSwizzleDMode / kNumBankGroupBytes / 8, 8)`
|
||||
constexpr bool kHasShortcut = (kSwizzleDMode / kNumBankGroupBytes) == 8;
|
||||
auto row = kHasShortcut ? (in_atom_offset / 8 + lane_idx) : (bank_group_index / 8);
|
||||
auto col = kHasShortcut ? (in_atom_offset) : (bank_group_index % 8);
|
||||
col ^= row % (kSwizzleDMode / 16);
|
||||
|
||||
// Add back into the base pointer
|
||||
// NOTES: think twice before modifying this, as changes may affect the number of instructions
|
||||
smem_ptr = reinterpret_cast<uint8_t*>(smem_d) + // Base pointer
|
||||
warp_idx * (WGMMA_M_PER_WARP * kSwizzleDMode) + // Warp offset
|
||||
m_offset * kSwizzleDMode + // Wave offset
|
||||
atom_offset * BLOCK_M * kSwizzleDMode + // Swizzle atom offset (constants)
|
||||
row * (kNumBankGroupBytes * 8) + col * kNumBankGroupBytes; // In-atom offset
|
||||
} else {
|
||||
// No swizzling, just padding
|
||||
smem_ptr = reinterpret_cast<uint8_t*>(smem_d + (m_offset + warp_idx * WGMMA_M_PER_WARP + lane_idx) * BLOCK_N + i * 8);
|
||||
}
|
||||
|
||||
// NOTES: only 16 lanes' addresses are used
|
||||
SM90_U32x2_STSM_N<nv_bfloat162>::copy(
|
||||
__float22bfloat162_rn({shifted_accum[i * 4 + 0], shifted_accum[i * 4 + 1]}),
|
||||
__float22bfloat162_rn({shifted_accum[i * 4 + 2], shifted_accum[i * 4 + 3]}),
|
||||
smem_ptr
|
||||
);
|
||||
}
|
||||
}
|
||||
cute::tma_store_fence();
|
||||
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
|
||||
|
||||
// Use TMA store to write back to global memory
|
||||
// TODO: compatible with FP32 output
|
||||
constexpr bool kWithGroupOffsetD = kGemmType == GemmType::MGroupedMasked;
|
||||
DG_STATIC_ASSERT(kNumMathThreads >= BLOCK_N / TMA_D_BLOCK_N, "Too many TMA blocks");
|
||||
if (threadIdx.x < BLOCK_N / TMA_D_BLOCK_N) {
|
||||
auto in_block_n_offset = threadIdx.x * TMA_D_BLOCK_N;
|
||||
auto smem_ptr = smem_d + in_block_n_offset * BLOCK_M;
|
||||
cute::SM90_TMA_STORE_2D::copy(&tensor_map_d, smem_ptr,
|
||||
epilogue_type_t::apply_index_n<TMA_D_BLOCK_N>(n_block_idx * BLOCK_N + in_block_n_offset),
|
||||
scheduler.get_global_idx<kWithGroupOffsetD>(shape_m, BLOCK_M, m_block_idx));
|
||||
cute::tma_store_arrive();
|
||||
}
|
||||
__syncwarp();
|
||||
}
|
||||
}
|
||||
#else
|
||||
if (blockIdx.x == 0 and threadIdx.x == 0)
|
||||
DG_DEVICE_ASSERT(false and "This kernel only support sm_90a");
|
||||
#endif
|
||||
}
|
||||
|
||||
}; // namespace deep_gemm
|
||||
|
||||
#pragma clang diagnostic pop
|
||||
+590
@@ -0,0 +1,590 @@
|
||||
#include <cutlass/arch/barrier.h>
|
||||
#include <cutlass/arch/reg_reconfig.h>
|
||||
|
||||
#include <cute/arch/cluster_sm90.hpp>
|
||||
#include <cute/arch/copy_sm90_desc.hpp>
|
||||
#include <cute/arch/copy_sm90_tma.hpp>
|
||||
|
||||
#include <deep_gemm/common/epilogue_utils.cuh>
|
||||
#include <deep_gemm/common/utils.cuh>
|
||||
#include <deep_gemm/common/scheduler.cuh>
|
||||
#include <deep_gemm/common/sm90_utils.cuh>
|
||||
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
|
||||
// LT-PATCH: upstream hard-#defines `__CUDA_ARCH__ 900` here, which forces the wgmma
|
||||
// kernel body on every compile pass and makes this source impossible to place in a
|
||||
// multi-arch fat binary (it emits sm_90-only instructions during e.g. the sm_89 pass
|
||||
// -> ptxas error). Removed so the existing `#if __CUDA_ARCH__ >= 900 ... #else assert
|
||||
// #endif` guard takes effect per-arch: the real body is built only into the sm_90a
|
||||
// cubin, other arches get a host-visible assert stub. The sm_90a pass is unchanged
|
||||
// (nvcc defines __CUDA_ARCH__=900 there regardless).
|
||||
|
||||
namespace deep_gemm {
|
||||
|
||||
using namespace deep_gemm::sm90;
|
||||
|
||||
template <uint32_t kNumFormerIters, uint32_t kGap, uint32_t kEnd, typename func_t>
|
||||
__device__ void dispatch_num_former_iters(uint32_t num_former_iters, const func_t& func) {
|
||||
if (num_former_iters == kNumFormerIters) {
|
||||
func(cute::Int<kNumFormerIters>{});
|
||||
return;
|
||||
}
|
||||
|
||||
if constexpr (kNumFormerIters + kGap <= kEnd)
|
||||
dispatch_num_former_iters<kNumFormerIters + kGap, kGap, kEnd>(num_former_iters, func);
|
||||
}
|
||||
|
||||
template <uint32_t SHAPE_M, uint32_t SHAPE_N, uint32_t SHAPE_K,
|
||||
uint32_t kNumGroups,
|
||||
uint32_t BLOCK_M, uint32_t BLOCK_N, uint32_t BLOCK_K,
|
||||
uint32_t kSwizzleDMode,
|
||||
uint32_t kNumStages, uint32_t kNumLastStages,
|
||||
uint32_t kNumTMAThreads, uint32_t kNumMathThreads,
|
||||
uint32_t kNumTMAMulticast, bool kIsTMAMulticastOnA,
|
||||
uint32_t kNumSMs, GemmType kGemmType,
|
||||
typename epilogue_type_t>
|
||||
__global__ __launch_bounds__(kNumTMAThreads + kNumMathThreads, 1) void
|
||||
sm90_fp8_gemm_1d2d_bias_impl(float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const __grid_constant__ cute::TmaDescriptor tensor_map_a,
|
||||
const __grid_constant__ cute::TmaDescriptor tensor_map_b,
|
||||
const __grid_constant__ cute::TmaDescriptor tensor_map_d,
|
||||
const __grid_constant__ cute::TmaDescriptor tensor_map_sfa) {
|
||||
// LT-PATCH: was `__CUDA_ARCH__ >= 900`. Tightened to Hopper-only (< 1000) so that in a
|
||||
// multi-arch fat binary that also targets Blackwell (sm_100/sm_120), this wgmma body is
|
||||
// NOT emitted for those passes (wgmma is sm_90a-only) -- they get the `#else` assert stub
|
||||
// instead. Blackwell dispatches to the SM89 kernel at runtime, so the stub is never run.
|
||||
#if (defined(__CUDA_ARCH__) and (__CUDA_ARCH__ >= 900) and (__CUDA_ARCH__ < 1000)) or defined(__CLION_IDE__)
|
||||
// Scaling checks
|
||||
DG_STATIC_ASSERT(BLOCK_K == 128, "Only support per-128-channel FP8 scaling");
|
||||
DG_STATIC_ASSERT(constexpr_ceil_div(BLOCK_N, BLOCK_K) == 1 or (constexpr_gcd(BLOCK_N, BLOCK_K) == BLOCK_N - BLOCK_K), "Too much B scales in a single block");
|
||||
|
||||
// Types
|
||||
using WGMMA = typename FP8MMASelector<BLOCK_N>::type;
|
||||
using Barrier = cutlass::arch::ClusterTransactionBarrier;
|
||||
DG_STATIC_ASSERT(BLOCK_M % WGMMA::M == 0, "Invalid block size");
|
||||
|
||||
// Overwrite shape constants if the compiler gives
|
||||
shape_m = SHAPE_M != 0 ? SHAPE_M : shape_m;
|
||||
shape_n = SHAPE_N != 0 ? SHAPE_N : shape_n;
|
||||
shape_k = SHAPE_K != 0 ? SHAPE_K : shape_k;
|
||||
|
||||
// Shared memory
|
||||
static constexpr bool kMustUseUniformedScaleB = (BLOCK_K % BLOCK_N == 0);
|
||||
static constexpr uint32_t SMEM_D_SIZE = BLOCK_M * BLOCK_N * sizeof(__nv_bfloat16);
|
||||
static constexpr uint32_t SMEM_A_SIZE_PER_STAGE = BLOCK_M * BLOCK_K * sizeof(__nv_fp8_e4m3);
|
||||
static constexpr uint32_t SMEM_B_SIZE_PER_STAGE = BLOCK_N * BLOCK_K * sizeof(__nv_fp8_e4m3);
|
||||
static constexpr uint32_t SMEM_SFA_SIZE_PER_STAGE = BLOCK_M * sizeof(float);
|
||||
const uint32_t& shape_k_scales = ceil_div(shape_k, BLOCK_K);
|
||||
const uint32_t& smem_sfb_size = align<uint32_t>(shape_k_scales * (kMustUseUniformedScaleB ? 1 : 2) * sizeof(float), sizeof(Barrier));
|
||||
|
||||
// Configs
|
||||
const uint32_t num_total_k_blocks = ceil_div(shape_k, BLOCK_K);
|
||||
const uint32_t warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
const uint32_t lane_idx = get_lane_idx();
|
||||
|
||||
// Prefetch TMA descriptors at the very beginning
|
||||
if (warp_idx == kNumMathThreads / 32 and cute::elect_one_sync()) {
|
||||
cute::prefetch_tma_descriptor(&tensor_map_a);
|
||||
cute::prefetch_tma_descriptor(&tensor_map_b);
|
||||
cute::prefetch_tma_descriptor(&tensor_map_sfa);
|
||||
cute::prefetch_tma_descriptor(&tensor_map_d);
|
||||
}
|
||||
__syncwarp();
|
||||
|
||||
// Align to 1024 bytes for swizzle-128B
|
||||
extern __shared__ __align__(1024) uint8_t smem_buffer[];
|
||||
DG_STATIC_ASSERT(SMEM_D_SIZE % 1024 == 0, "Shared memory of A/B must be aligned to 1024 bytes");
|
||||
|
||||
// Data on shared memory
|
||||
auto smem_d = reinterpret_cast<__nv_bfloat16*>(smem_buffer);
|
||||
auto smem_a = PatternVisitor([&](const uint32_t& i) {
|
||||
return reinterpret_cast<__nv_fp8_e4m3*>(smem_buffer + SMEM_D_SIZE + i * SMEM_A_SIZE_PER_STAGE);
|
||||
});
|
||||
auto smem_b = PatternVisitor([&](const uint32_t& i) {
|
||||
return reinterpret_cast<__nv_fp8_e4m3*>(smem_buffer + SMEM_D_SIZE + kNumStages * SMEM_A_SIZE_PER_STAGE + i * SMEM_B_SIZE_PER_STAGE);
|
||||
});
|
||||
constexpr uint32_t SMEM_SF_OFFSET = SMEM_D_SIZE + kNumStages * (SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE);
|
||||
auto smem_sfa = PatternVisitor([&](const uint32_t& i) {
|
||||
return reinterpret_cast<float*>(smem_buffer + SMEM_SF_OFFSET + i * SMEM_SFA_SIZE_PER_STAGE);
|
||||
});
|
||||
auto smem_sfb = reinterpret_cast<float*>(smem_buffer + SMEM_SF_OFFSET + kNumStages * SMEM_SFA_SIZE_PER_STAGE);
|
||||
|
||||
// Fill barriers
|
||||
auto barrier_start_ptr = reinterpret_cast<Barrier*>(reinterpret_cast<uint8_t*>(smem_sfb) + smem_sfb_size);
|
||||
auto full_barriers = PatternVisitor([&](const uint32_t& i) { return barrier_start_ptr + i; });
|
||||
auto empty_barriers = PatternVisitor([&](const uint32_t& i) { return barrier_start_ptr + kNumStages + i; });
|
||||
|
||||
// Initialize barriers
|
||||
DG_STATIC_ASSERT(kNumTMAMulticast <= 32, "Too many TMA multicast");
|
||||
if (warp_idx == kNumMathThreads / 32 + 1 and cute::elect_one_sync()) {
|
||||
// NOTES: we always use `lane_idx` to arrive for the `lane_idx`-th CTA in the cluster,
|
||||
// even with TMA multicast disabled, we want to make the behavior aligned
|
||||
#pragma unroll
|
||||
for (uint32_t i = 0; i < kNumStages; ++ i) {
|
||||
full_barriers[i]->init(1);
|
||||
empty_barriers[i]->init(kNumTMAMulticast * kNumMathThreads / 32);
|
||||
}
|
||||
|
||||
// Make initialized barrier visible in async proxy
|
||||
cutlass::arch::fence_barrier_init();
|
||||
}
|
||||
|
||||
// Synchronize all threads to make barrier visible in normal memory model
|
||||
(kNumTMAMulticast > 1) ? cute::cluster_sync() : __syncthreads();
|
||||
|
||||
// Register reconfigurations
|
||||
constexpr uint32_t kNumTMARegisters = 40;
|
||||
constexpr uint32_t kNumMathRegisters = 232;
|
||||
|
||||
// Block scheduler
|
||||
uint32_t m_block_idx, n_block_idx;
|
||||
auto scheduler = Scheduler<kGemmType, BLOCK_M, BLOCK_N, kNumGroups, kNumTMAMulticast, kIsTMAMulticastOnA, kNumSMs>(shape_m, shape_n, shape_k, grouped_layout);
|
||||
|
||||
// Pipeline and TMA phases
|
||||
uint32_t stage_idx = 0, phase = 0;
|
||||
auto advance_pipeline = [&](uint32_t& k_block_idx) {
|
||||
++ k_block_idx;
|
||||
|
||||
// Flip phases only if reach the next first stage
|
||||
stage_idx = stage_idx == kNumStages - 1 ? 0 : stage_idx + 1;
|
||||
phase ^= stage_idx == 0;
|
||||
};
|
||||
|
||||
if (warp_idx >= kNumMathThreads / 32) {
|
||||
// TMA warp-group for loading data
|
||||
cutlass::arch::warpgroup_reg_dealloc<kNumTMARegisters>();
|
||||
|
||||
// NOTES: only one thread (or warp) will be used
|
||||
if (warp_idx == kNumMathThreads / 32 and cute::elect_one_sync()) {
|
||||
// Persistently schedule over blocks
|
||||
while (scheduler.get_next_block(m_block_idx, n_block_idx)) {
|
||||
// Assign TMA multicast number into A and B
|
||||
// NOTES: there may be additional odd rows/columns or cases where multicast is not possible.
|
||||
const bool is_tma_multicast_valid = scheduler.is_tma_multicast_valid(m_block_idx);
|
||||
const uint32_t num_tma_multicast_a = (kIsTMAMulticastOnA and is_tma_multicast_valid) ? kNumTMAMulticast : 1;
|
||||
const uint32_t num_tma_multicast_b = (not kIsTMAMulticastOnA and is_tma_multicast_valid) ? kNumTMAMulticast : 1;
|
||||
DG_STATIC_ASSERT(kNumTMAMulticast <= 2, "Scheduler does not support > 2 TMA multicast");
|
||||
|
||||
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
|
||||
// Wait consumer release
|
||||
empty_barriers[stage_idx]->wait(phase ^ 1);
|
||||
|
||||
// Issue TMA A
|
||||
constexpr bool kWithGroupOffsetA = kGemmType == GemmType::MGroupedMasked;
|
||||
auto& full_barrier = *full_barriers[stage_idx];
|
||||
const uint32_t k_idx = k_block_idx * BLOCK_K;
|
||||
tma_copy(&tensor_map_a, reinterpret_cast<uint64_t*>(&full_barrier),
|
||||
smem_a[stage_idx], k_idx, scheduler.get_global_idx<kWithGroupOffsetA>(shape_m, BLOCK_M, m_block_idx),
|
||||
num_tma_multicast_a);
|
||||
tma_copy(&tensor_map_sfa, reinterpret_cast<uint64_t*>(&full_barrier),
|
||||
smem_sfa[stage_idx], m_block_idx * BLOCK_M, scheduler.get_global_idx<kWithGroupOffsetA>(shape_k_scales, 1, k_block_idx),
|
||||
num_tma_multicast_a);
|
||||
|
||||
// Issue TMA B
|
||||
tma_copy(&tensor_map_b, reinterpret_cast<uint64_t*>(&full_barrier),
|
||||
smem_b[stage_idx], k_idx, scheduler.get_global_idx<true>(shape_n, BLOCK_N, n_block_idx, m_block_idx),
|
||||
num_tma_multicast_b);
|
||||
full_barrier.arrive_and_expect_tx(SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE + SMEM_SFA_SIZE_PER_STAGE);
|
||||
}
|
||||
}
|
||||
|
||||
// To safely deconstruct distributed shared barriers, we need another round of empty waits
|
||||
if constexpr (kNumTMAMulticast > 1) {
|
||||
for (uint32_t i = 0; i < kNumStages; advance_pipeline(i))
|
||||
empty_barriers[stage_idx]->wait(phase ^ 1);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Math warp-groups for WGMMA
|
||||
cutlass::arch::warpgroup_reg_alloc<kNumMathRegisters>();
|
||||
|
||||
// NOTES: use `__shfl_sync` to encourage NVCC to use unified registers
|
||||
const auto math_wg_idx = __shfl_sync(0xffffffff, threadIdx.x / 128, 0);
|
||||
const auto r_0 = warp_idx * 16 + lane_idx / 4, r_1 = r_0 + 8;
|
||||
|
||||
auto a_desc = make_smem_desc(smem_a[0] + math_wg_idx * WGMMA::M * BLOCK_K, 1);
|
||||
auto b_desc = make_smem_desc(smem_b[0], 1);
|
||||
const uint32_t a_desc_lo = __shfl_sync(0xffffffff, a_desc.reg32_[0], 0);
|
||||
const uint32_t b_desc_lo = __shfl_sync(0xffffffff, b_desc.reg32_[0], 0);
|
||||
|
||||
// Persistently schedule over blocks
|
||||
while (scheduler.get_next_block(m_block_idx, n_block_idx)) {
|
||||
// Decide the number of scales B to load
|
||||
DG_TRAP_ONLY_DEVICE_ASSERT(shape_n % 8 == 0);
|
||||
uint32_t num_former_iters = BLOCK_N / 8, num_full_iters = num_former_iters;
|
||||
if constexpr (not kMustUseUniformedScaleB) {
|
||||
num_former_iters = min(BLOCK_N, BLOCK_K - n_block_idx * BLOCK_N % BLOCK_K) / 8;
|
||||
num_full_iters = min(shape_n - n_block_idx * BLOCK_N, BLOCK_N) / 8;
|
||||
}
|
||||
uint32_t num_sfb = shape_k_scales * (num_former_iters >= num_full_iters ? 1 : 2);
|
||||
|
||||
// Load B scales with math warp-groups
|
||||
// NOTES: except the first warp, we want to overlap loading B scales with TMA stores between tasks
|
||||
if (threadIdx.x >= 32) {
|
||||
auto num_previous_lines = scheduler.get_global_idx<true>(ceil_div(shape_n, BLOCK_K), 0, 0, m_block_idx);
|
||||
auto local_sfb = sfb + (num_previous_lines + ((n_block_idx * BLOCK_N) / BLOCK_K)) * shape_k_scales;
|
||||
#pragma unroll
|
||||
for (uint32_t i = threadIdx.x - 32; i < num_sfb; i += kNumMathThreads - 32)
|
||||
st_shared(smem_sfb + i, __ldg(local_sfb + i));
|
||||
}
|
||||
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
|
||||
|
||||
// Accumulation for WGMMA or CUDA promotion
|
||||
constexpr uint32_t WAVE_BLOCK_M = WGMMA::M * (BLOCK_M <= 64 ? 1 : 2);
|
||||
DG_STATIC_ASSERT(BLOCK_M % WAVE_BLOCK_M == 0, "Invalid block sizes");
|
||||
float accum[WGMMA::kNumAccum], final_accum[WGMMA::kNumAccum * (BLOCK_M / WAVE_BLOCK_M)] = {0};
|
||||
|
||||
// Empty barrier arrival
|
||||
auto empty_barrier_arrive = [&]() {
|
||||
if constexpr (kNumTMAMulticast == 1) {
|
||||
lane_idx == 0 ? empty_barriers[stage_idx]->arrive() : void();
|
||||
} else {
|
||||
auto target_cta = scheduler.is_peer_cta_alive ? lane_idx : cute::block_rank_in_cluster();
|
||||
lane_idx < kNumTMAMulticast ? empty_barriers[stage_idx]->arrive(target_cta) : void();
|
||||
}
|
||||
};
|
||||
|
||||
// Skip useless computations
|
||||
if (scheduler.is_computation_valid(m_block_idx, math_wg_idx * WGMMA::M)) {
|
||||
// The compiler must know the dynamic variable `num_former_iters`'s real value
|
||||
constexpr bool kShouldOptimize = BLOCK_K / constexpr_gcd(BLOCK_K, BLOCK_N) <= 4 and not kMustUseUniformedScaleB;
|
||||
constexpr uint32_t kGap = constexpr_gcd(BLOCK_K, BLOCK_N) / 8;
|
||||
constexpr uint32_t kEnd = kShouldOptimize ? BLOCK_K / 8 : 0;
|
||||
|
||||
// Dispatch `num_former_iters` and launch MMAs
|
||||
dispatch_num_former_iters<0, kGap, kEnd>(kShouldOptimize ? num_former_iters : 0, [&](auto _) {
|
||||
#pragma unroll 8
|
||||
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
|
||||
const auto& a_desc_base_lo = a_desc_lo + stage_idx * (SMEM_A_SIZE_PER_STAGE / 16);
|
||||
const auto& b_desc_base_lo = b_desc_lo + stage_idx * (SMEM_B_SIZE_PER_STAGE / 16);
|
||||
|
||||
// Read B scales
|
||||
float scale_b_0 = ld_shared(smem_sfb + k_block_idx), scale_b_1;
|
||||
// NOTES: even some blocks do not need to read the second row, but we still load one to align with other blocks
|
||||
if constexpr (not kMustUseUniformedScaleB)
|
||||
scale_b_1 = ld_shared(smem_sfb + k_block_idx + shape_k_scales);
|
||||
|
||||
// Wait TMA arrivals
|
||||
full_barriers[stage_idx]->wait(phase);
|
||||
|
||||
// TODO: remove some useless computation for unaligned Ms
|
||||
#pragma unroll
|
||||
for (uint32_t local_idx = 0; local_idx < BLOCK_M / WAVE_BLOCK_M; ++ local_idx) {
|
||||
auto m_offset = local_idx * WAVE_BLOCK_M;
|
||||
|
||||
// Read A scales
|
||||
// NOTES: all shared memory read must be prior to `warpgroup_arrive` to avoid next scheduled block polluting the results
|
||||
auto scale_a_0 = ld_shared(smem_sfa[stage_idx] + r_0 + m_offset);
|
||||
auto scale_a_1 = ld_shared(smem_sfa[stage_idx] + r_1 + m_offset);
|
||||
|
||||
// Commit WGMMA instructions
|
||||
#pragma unroll
|
||||
for (uint32_t i = 0; i < WGMMA::kNumAccum; ++ i)
|
||||
warpgroup_fence_operand(accum[i]);
|
||||
warpgroup_arrive();
|
||||
#pragma unroll
|
||||
for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) {
|
||||
a_desc.reg32_[0] = a_desc_base_lo + (m_offset * BLOCK_K + k * WGMMA::K) / 16;
|
||||
b_desc.reg32_[0] = b_desc_base_lo + k * WGMMA::K / 16;
|
||||
WGMMA::wgmma(a_desc, b_desc, accum, k);
|
||||
}
|
||||
warpgroup_commit_batch();
|
||||
#pragma unroll
|
||||
for (uint32_t i = 0; i < WGMMA::kNumAccum; ++ i)
|
||||
warpgroup_fence_operand(accum[i]);
|
||||
warpgroup_wait<0>();
|
||||
|
||||
// Notify barrier arrival at the last warpgroup wave
|
||||
if (local_idx == BLOCK_M / WAVE_BLOCK_M - 1)
|
||||
empty_barrier_arrive();
|
||||
|
||||
// Promote with scales
|
||||
// NOTES: making it as predicates is very important for performance, comparing to two loops
|
||||
float scale_0_0 = scale_a_0 * scale_b_0, scale_1_0 = scale_a_1 * scale_b_0;
|
||||
float scale_0_1, scale_1_1;
|
||||
if constexpr (not kMustUseUniformedScaleB)
|
||||
scale_0_1 = scale_a_0 * scale_b_1, scale_1_1 = scale_a_1 * scale_b_1;
|
||||
|
||||
auto shifted_accum = final_accum + WGMMA::kNumAccum * local_idx;
|
||||
#pragma unroll
|
||||
for (uint32_t i = 0; i < WGMMA::kNumAccum / 4; ++ i) {
|
||||
// NOTES: for unrolled `num_former_iters` cases, we expect the compiler to automatically make it a constant
|
||||
bool predicate = kMustUseUniformedScaleB or i < num_former_iters;
|
||||
shifted_accum[i * 4 + 0] += (predicate ? scale_0_0 : scale_0_1) * accum[i * 4 + 0];
|
||||
shifted_accum[i * 4 + 1] += (predicate ? scale_0_0 : scale_0_1) * accum[i * 4 + 1];
|
||||
shifted_accum[i * 4 + 2] += (predicate ? scale_1_0 : scale_1_1) * accum[i * 4 + 2];
|
||||
shifted_accum[i * 4 + 3] += (predicate ? scale_1_0 : scale_1_1) * accum[i * 4 + 3];
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
} else {
|
||||
#pragma unroll
|
||||
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
|
||||
full_barriers[stage_idx]->wait(phase);
|
||||
empty_barrier_arrive();
|
||||
}
|
||||
}
|
||||
|
||||
// TMA checks
|
||||
constexpr uint32_t kNumElemBytes = sizeof(nv_bfloat16);
|
||||
constexpr uint32_t TMA_D_BLOCK_N = kSwizzleDMode == 0 ? BLOCK_N : (kSwizzleDMode / kNumElemBytes);
|
||||
constexpr uint32_t WGMMA_M_PER_WARP = WGMMA::M / 4;
|
||||
DG_STATIC_ASSERT(BLOCK_M % 8 == 0, "Invalid swizzling atom");
|
||||
DG_STATIC_ASSERT(BLOCK_N % TMA_D_BLOCK_N == 0 and BLOCK_N / TMA_D_BLOCK_N <= 32,
|
||||
"Unaligned TMA store or too many TMA store instructions");
|
||||
DG_STATIC_ASSERT(TMA_D_BLOCK_N % 8 == 0, "Invalid TMA block N");
|
||||
// Wait last TMA store to be finished
|
||||
float* bias_ptr = bias + n_block_idx*BLOCK_N + (lane_idx % 4) * 2;
|
||||
#pragma unroll
|
||||
for(uint32_t local_idx=0; local_idx < BLOCK_M / WAVE_BLOCK_M; ++ local_idx){
|
||||
auto shifted_accum = final_accum + WGMMA::kNumAccum * local_idx;
|
||||
#pragma unroll
|
||||
for (auto i = 0; i < WGMMA::kNumAccum / 4; ++ i) {
|
||||
shifted_accum[4*i + 0] += bias_ptr[8*i + 0];
|
||||
shifted_accum[4*i + 1] += bias_ptr[8*i + 1];
|
||||
shifted_accum[4*i + 2] += bias_ptr[8*i + 0];
|
||||
shifted_accum[4*i + 3] += bias_ptr[8*i + 1];
|
||||
}
|
||||
}
|
||||
|
||||
if (threadIdx.x < BLOCK_N / TMA_D_BLOCK_N)
|
||||
cute::tma_store_wait<0>();
|
||||
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
|
||||
|
||||
// Write back to shared memory using STSM and issue TMA stores
|
||||
DG_STATIC_ASSERT(WGMMA::kNumAccum % 4 == 0, "Invalid STSM x2 vectorization");
|
||||
#pragma unroll
|
||||
for (uint32_t local_idx = 0; local_idx < BLOCK_M / WAVE_BLOCK_M; ++ local_idx) {
|
||||
auto m_offset = local_idx * WAVE_BLOCK_M;
|
||||
auto shifted_accum = final_accum + WGMMA::kNumAccum * local_idx;
|
||||
#pragma unroll
|
||||
for (auto i = 0; i < WGMMA::kNumAccum / 4; ++ i) {
|
||||
// Swizzle or padding into the correct address
|
||||
uint8_t* smem_ptr = nullptr;
|
||||
if constexpr (kSwizzleDMode > 0) {
|
||||
// Calculate the swizzling atom offset and in-atom offset
|
||||
constexpr uint32_t kNumBankGroupBytes = 16;
|
||||
auto atom_offset = i / (TMA_D_BLOCK_N / 8), in_atom_offset = i % (TMA_D_BLOCK_N / 8);
|
||||
|
||||
// Calculate the index of the bank group to be written in the atom
|
||||
auto bank_group_index = in_atom_offset + lane_idx * (kSwizzleDMode / kNumBankGroupBytes);
|
||||
|
||||
// Reshape the atom in another view and swizzle
|
||||
// - original: `(BLOCK_M, kSwizzleDMode / kNumBankGroupBytes)`
|
||||
// - new: `(BLOCK_M * kSwizzleDMode / kNumBankGroupBytes / 8, 8)`
|
||||
constexpr bool kHasShortcut = (kSwizzleDMode / kNumBankGroupBytes) == 8;
|
||||
auto row = kHasShortcut ? (in_atom_offset / 8 + lane_idx) : (bank_group_index / 8);
|
||||
auto col = kHasShortcut ? (in_atom_offset) : (bank_group_index % 8);
|
||||
col ^= row % (kSwizzleDMode / 16);
|
||||
|
||||
// Add back into the base pointer
|
||||
// NOTES: think twice before modifying this, as changes may affect the number of instructions
|
||||
smem_ptr = reinterpret_cast<uint8_t*>(smem_d) + // Base pointer
|
||||
warp_idx * (WGMMA_M_PER_WARP * kSwizzleDMode) + // Warp offset
|
||||
m_offset * kSwizzleDMode + // Wave offset
|
||||
atom_offset * BLOCK_M * kSwizzleDMode + // Swizzle atom offset (constants)
|
||||
row * (kNumBankGroupBytes * 8) + col * kNumBankGroupBytes; // In-atom offset
|
||||
} else {
|
||||
// No swizzling, just padding
|
||||
smem_ptr = reinterpret_cast<uint8_t*>(smem_d + (m_offset + warp_idx * WGMMA_M_PER_WARP + lane_idx) * BLOCK_N + i * 8);
|
||||
}
|
||||
|
||||
// NOTES: only 16 lanes' addresses are used
|
||||
SM90_U32x2_STSM_N<nv_bfloat162>::copy(
|
||||
__float22bfloat162_rn({shifted_accum[i * 4 + 0], shifted_accum[i * 4 + 1]}),
|
||||
__float22bfloat162_rn({shifted_accum[i * 4 + 2], shifted_accum[i * 4 + 3]}),
|
||||
smem_ptr
|
||||
);
|
||||
}
|
||||
}
|
||||
cute::tma_store_fence();
|
||||
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
|
||||
|
||||
// Use TMA store to write back to global memory
|
||||
// TODO: compatible with FP32 output
|
||||
constexpr bool kWithGroupOffsetD = kGemmType == GemmType::MGroupedMasked;
|
||||
DG_STATIC_ASSERT(kNumMathThreads >= BLOCK_N / TMA_D_BLOCK_N, "Too many TMA blocks");
|
||||
if (threadIdx.x < BLOCK_N / TMA_D_BLOCK_N) {
|
||||
auto in_block_n_offset = threadIdx.x * TMA_D_BLOCK_N;
|
||||
auto smem_ptr = smem_d + in_block_n_offset * BLOCK_M;
|
||||
cute::SM90_TMA_STORE_2D::copy(&tensor_map_d, smem_ptr,
|
||||
epilogue_type_t::apply_index_n<TMA_D_BLOCK_N>(n_block_idx * BLOCK_N + in_block_n_offset),
|
||||
scheduler.get_global_idx<kWithGroupOffsetD>(shape_m, BLOCK_M, m_block_idx));
|
||||
cute::tma_store_arrive();
|
||||
}
|
||||
__syncwarp();
|
||||
}
|
||||
}
|
||||
#else
|
||||
if (blockIdx.x == 0 and threadIdx.x == 0)
|
||||
DG_DEVICE_ASSERT(false and "This kernel only support sm_90a");
|
||||
#endif
|
||||
}
|
||||
|
||||
static cudaLaunchConfig_t construct_launch_config(const cudaStream_t& stream, const int& smem_size,
|
||||
const dim3& grid_dim, const dim3& block_dim, const int& cluster_dim) {
|
||||
|
||||
cudaLaunchConfig_t config;
|
||||
config.gridDim = grid_dim;
|
||||
config.blockDim = block_dim;
|
||||
config.dynamicSmemBytes = smem_size;
|
||||
config.stream = stream;
|
||||
config.numAttrs = 0;
|
||||
config.attrs = nullptr;
|
||||
|
||||
// NOTES: must use `static` or the `attr` will be deconstructed
|
||||
static cudaLaunchAttribute attr;
|
||||
if (cluster_dim > 1) {
|
||||
attr.id = cudaLaunchAttributeClusterDimension;
|
||||
attr.val.clusterDim = {static_cast<unsigned>(cluster_dim), 1, 1};
|
||||
config.attrs = &attr;
|
||||
config.numAttrs = 1;
|
||||
}
|
||||
return config;
|
||||
}
|
||||
|
||||
|
||||
// static auto launch_kernel(auto kernel, const cudaLaunchConfig_t& config, float* sfb, float* bias, int* grouped_layout,
|
||||
// uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
// const CUtensorMap tensor_map_a,
|
||||
// const CUtensorMap tensor_map_b,
|
||||
// const CUtensorMap tensor_map_d,
|
||||
// const CUtensorMap tensor_map_sfa) {
|
||||
// // void* ptr_args[] = {&sfb, &bias, &grouped_layout, &shape_m, &shape_n, &shape_k, &tensor_map_a, &tensor_map_b, &tensor_map_d, &tensor_map_sfa};
|
||||
// return
|
||||
// }
|
||||
|
||||
|
||||
template<int N, int K>
|
||||
void sm90_fp8_gemm_1d2d_bias_launch(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa){
|
||||
dim3 grid{num_sms, 1, 1};
|
||||
dim3 block{num_threads, 1, 1};
|
||||
const auto config = construct_launch_config(stream, smem_size, grid, block, cluster_dim);
|
||||
if(num_sms == 132){
|
||||
auto kernel = &sm90_fp8_gemm_1d2d_bias_impl<0, N, K, 1, 256, 128, 128, 128, 3, (K / 128) % 3, 128, 256, 2, true, 132, GemmType::Normal, EpilogueIdentity>;
|
||||
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
|
||||
cudaLaunchKernelEx(&config, kernel, sfb, bias, grouped_layout, shape_m, shape_n, shape_k, tensor_map_a, tensor_map_b, tensor_map_d, tensor_map_sfa);
|
||||
} else if(num_sms == 116) {
|
||||
auto kernel = &sm90_fp8_gemm_1d2d_bias_impl<0, N, K, 1, 256, 128, 128, 128, 3, (K / 128) % 3, 128, 256, 2, true, 116, GemmType::Normal, EpilogueIdentity>;
|
||||
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
|
||||
cudaLaunchKernelEx(&config, kernel, sfb, bias, grouped_layout, shape_m, shape_n, shape_k, tensor_map_a, tensor_map_b, tensor_map_d, tensor_map_sfa);
|
||||
} else if (num_sms == 100) {
|
||||
auto kernel = &sm90_fp8_gemm_1d2d_bias_impl<0, N, K, 1, 256, 128, 128, 128, 3, (K / 128) % 3, 128, 256, 2, true, 100, GemmType::Normal, EpilogueIdentity>;
|
||||
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
|
||||
cudaLaunchKernelEx(&config, kernel, sfb, bias, grouped_layout, shape_m, shape_n, shape_k, tensor_map_a, tensor_map_b, tensor_map_d, tensor_map_sfa);
|
||||
} else {
|
||||
// The supported SM counts are exactly the branches above (the only kernels
|
||||
// instantiated). Fail loudly instead of falling through with no launch,
|
||||
// which would leave the output buffer uninitialized.
|
||||
throw std::runtime_error("Unsupported num_sms=" + std::to_string(num_sms)
|
||||
+ " (blockwise SM90 GEMM is built for 132, 116, and 100 SMs)");
|
||||
}
|
||||
|
||||
// launch_kernel(kernel, config, sfb, bias, grouped_layout, shape_m, shape_n, shape_k, tensor_map_a, tensor_map_b, tensor_map_d, tensor_map_sfa);
|
||||
}
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<2048, 2048>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<4096, 2048>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<8192, 2048>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<16384, 2048>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<2048, 4096>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<4096, 4096>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<8192, 4096>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<16384, 4096>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<2048, 8192>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<4096, 8192>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<8192, 8192>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<16384, 8192>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<2048, 16384>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<4096, 16384>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<8192, 16384>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
template void sm90_fp8_gemm_1d2d_bias_launch<16384, 16384>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
}; // namespace deep_gemm
|
||||
@@ -0,0 +1,287 @@
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/layout/layout.h"
|
||||
#include <cute/tensor.hpp>
|
||||
#include <c10/cuda/CUDAException.h>
|
||||
#include <torch/extension.h>
|
||||
#include <torch/python.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <iostream>
|
||||
|
||||
#include "kernel_traits.cuh"
|
||||
#include "static_switch.h"
|
||||
|
||||
namespace sm89{
|
||||
using namespace cute;
|
||||
|
||||
__device__ static void copy_1d(float* gmem_src, float* smem_dst)
|
||||
{
|
||||
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint((void*)smem_dst);
|
||||
asm volatile("cp.async.ca.shared.global.L2::128B [%0], [%1], %2;\n"
|
||||
:: "r"(smem_int_ptr),
|
||||
"l"(gmem_src),
|
||||
"n"(sizeof(float)));
|
||||
}
|
||||
|
||||
template <typename KernelTraits=gemm_traits<128, 256, 2, 4096, 2, 4, true, half_t, bfloat16_t>>
|
||||
__global__ void gemm_fp8_kernel(float_e4m3_t* Aptr, float* sfa, float_e4m3_t* Bptr, float* sfb, float* bias_ptr, void* out, int M, int N, int K, int TMA_ALIGNED_M){
|
||||
using output_t = typename KernelTraits::out_t;
|
||||
using SmemLayoutA = typename KernelTraits::SmemLayoutA;
|
||||
using SmemLayoutB = typename KernelTraits::SmemLayoutB;
|
||||
using SmemLayoutC = typename KernelTraits::SmemLayoutC;
|
||||
|
||||
constexpr int BM = KernelTraits::BM;
|
||||
constexpr int BN = KernelTraits::BN;
|
||||
constexpr int BK = KernelTraits::BK;
|
||||
constexpr int Ksfa = KernelTraits::KSF;
|
||||
constexpr bool has_bias = KernelTraits::HasBias;
|
||||
extern __shared__ float smem_[];
|
||||
float *bias_shm = smem_;
|
||||
float *sfa_shm = reinterpret_cast<float*>(bias_shm + cosize(typename KernelTraits::SmemLayoutBias{}));
|
||||
|
||||
output_t* C_shm = reinterpret_cast<output_t*>(sfa_shm + cosize(typename KernelTraits::SmemLayoutSFA{}));
|
||||
float_e4m3_t* A_shm = reinterpret_cast<float_e4m3_t*>(sfa_shm + cosize(typename KernelTraits::SmemLayoutSFA{}));
|
||||
float_e4m3_t* B_shm = reinterpret_cast<float_e4m3_t*>(A_shm + cosize(SmemLayoutA{}));
|
||||
|
||||
int idx = threadIdx.x;
|
||||
int ix = blockIdx.x;
|
||||
int iy = blockIdx.y;
|
||||
// sfa += BM * iy;
|
||||
sfb += KernelTraits::NUM_SFB_PER_STEP * ix * Ksfa;
|
||||
|
||||
|
||||
output_t* Cptr = reinterpret_cast<output_t*>(out);
|
||||
|
||||
Tensor A = make_tensor(make_gmem_ptr(Aptr), make_shape(M, K), make_stride(K, Int<1>{}));
|
||||
Tensor B = make_tensor(make_gmem_ptr(Bptr), make_shape(N, K), make_stride(K, Int<1>{}));
|
||||
Tensor D = make_tensor(make_gmem_ptr(Cptr), make_shape(M, N), make_stride(N, Int<1>{}));
|
||||
Tensor SFA = make_tensor(make_gmem_ptr(sfa), make_shape(M, Ksfa), make_stride(Int<1>{}, TMA_ALIGNED_M));
|
||||
|
||||
Tensor gA = local_tile(A, make_tile(Int<BM>{}, Int<BK>{}), make_coord(iy, _));
|
||||
Tensor gB = local_tile(B, make_tile(Int<BN>{}, Int<BK>{}), make_coord(ix, _));
|
||||
Tensor gD = local_tile(D, make_tile(Int<BM>{}, Int<BN>{}), make_coord(iy, ix));
|
||||
Tensor gSFA = local_tile(SFA, make_tile(Int<BM>{}, Int<1>{}), make_coord(iy, _));
|
||||
|
||||
auto sBias = make_tensor(make_smem_ptr(bias_shm), typename KernelTraits::SmemLayoutBias{});
|
||||
if constexpr (has_bias){
|
||||
Tensor Bias = make_tensor(make_gmem_ptr(bias_ptr), make_shape(_1{}, N), make_stride(N, Int<1>{}));
|
||||
Tensor gBias = local_tile(Bias, make_tile(Int<1>{}, Int<BN>{}), make_coord(_, ix));
|
||||
typename KernelTraits::G2SBiasCopy g2s_bias_copy;
|
||||
auto g2s_bias_thr_copy = g2s_bias_copy.get_slice(idx);
|
||||
auto tCBiasgBias = g2s_bias_thr_copy.partition_S(gBias);
|
||||
auto tCBiassBias = g2s_bias_thr_copy.partition_D(sBias);
|
||||
if(idx < BN){
|
||||
copy_1d((float*)&gBias(0) + idx, (float*)&sBias(0) + idx);
|
||||
}
|
||||
}
|
||||
|
||||
auto sSFA = make_tensor(make_smem_ptr(sfa_shm), typename KernelTraits::SmemLayoutSFA{});
|
||||
auto sA = make_tensor(make_smem_ptr(A_shm), SmemLayoutA{});
|
||||
auto sB = make_tensor(make_smem_ptr(B_shm), SmemLayoutB{});
|
||||
|
||||
typename KernelTraits::MMATile tiled_mma;
|
||||
auto thr_mma = tiled_mma.get_slice(threadIdx.x);
|
||||
|
||||
auto tCrA = thr_mma.partition_fragment_A(gA(_, _, 0));
|
||||
auto tCrB = thr_mma.partition_fragment_B(gB(_, _, 0));
|
||||
auto tCrD = thr_mma.partition_fragment_C(gD);
|
||||
clear(tCrD);
|
||||
auto tCrD_fp32 = make_tensor_like<float>(tCrD);
|
||||
clear(tCrD_fp32);
|
||||
|
||||
typename KernelTraits::G2STiledCopy g2s_tiled_copy;
|
||||
auto g2s_thr_copy = g2s_tiled_copy.get_slice(idx);
|
||||
auto tAgA_copy = g2s_thr_copy.partition_S(gA);
|
||||
auto tAsA_copy = g2s_thr_copy.partition_D(sA);
|
||||
auto tBgB_copy = g2s_thr_copy.partition_S(gB);
|
||||
auto tBsB_copy = g2s_thr_copy.partition_D(sB);
|
||||
|
||||
auto s2r_tiled_copy_a = make_tiled_copy_A(typename KernelTraits::S2RCopyAtomA{}, tiled_mma);
|
||||
auto s2r_thr_copy_a = s2r_tiled_copy_a.get_slice(idx);
|
||||
auto tAsA = s2r_thr_copy_a.partition_S(sA);
|
||||
auto tCrA_view = s2r_thr_copy_a.retile_D(tCrA);
|
||||
|
||||
|
||||
auto s2r_tiled_copy_b = make_tiled_copy_B(typename KernelTraits::S2RCopyAtomB{}, tiled_mma);
|
||||
auto s2r_thr_copy_b = s2r_tiled_copy_b.get_slice(idx);
|
||||
auto tBsB = s2r_thr_copy_b.partition_S(sB);
|
||||
auto tCrB_view = s2r_thr_copy_b.retile_D(tCrB);
|
||||
|
||||
auto cA = make_identity_tensor(make_shape(size<0>(sA), size<1>(sA)));
|
||||
auto tAcA = g2s_thr_copy.partition_S(cA);
|
||||
int residual = M - iy*BM;
|
||||
|
||||
int itile_to_read = 0;
|
||||
int ismem_read = 0;
|
||||
int ismem_write = 0;
|
||||
int ismem_read_sfa = 0;
|
||||
constexpr int kStages = KernelTraits::KStages;
|
||||
|
||||
#pragma unroll
|
||||
for(int istage=0; istage<kStages - 1; ++istage){
|
||||
for (size_t m = 0; m < size<1>(tAsA_copy); m++)
|
||||
{
|
||||
for (size_t k = 0; k < size<2>(tAsA_copy); k++)
|
||||
{
|
||||
if(get<0>(tAcA(0, m, k)) < residual){
|
||||
cute::copy(g2s_tiled_copy, tAgA_copy(_, m, k, istage), tAsA_copy(_, m, k, istage));
|
||||
}
|
||||
}
|
||||
}
|
||||
if(idx < KernelTraits::THREADS_SFA_COPY && (BM * iy + idx * KernelTraits::SFA_ELEMS_PER_COPY < M)) {
|
||||
copy_1d((float*)&gSFA(0, 0, istage) + idx*KernelTraits::SFA_ELEMS_PER_COPY, (float*)&sSFA(0, istage) + idx*KernelTraits::SFA_ELEMS_PER_COPY);
|
||||
}
|
||||
cute::copy(g2s_tiled_copy, tBgB_copy(_, _, _, istage), tBsB_copy(_, _, _, istage));
|
||||
cp_async_fence();
|
||||
++itile_to_read;
|
||||
++ismem_write;
|
||||
}
|
||||
|
||||
cp_async_wait<kStages - 2>();
|
||||
__syncthreads();
|
||||
|
||||
cute::copy(s2r_tiled_copy_a, tAsA(_, _, 0, ismem_read), tCrA_view(_, _, 0));
|
||||
cute::copy(s2r_tiled_copy_b, tBsB(_, _, 0, ismem_read), tCrB_view(_, _, 0));
|
||||
|
||||
static constexpr int nk = size<2>(tCrA);
|
||||
auto sfa_tv = typename KernelTraits::SFAThreadLayout{};
|
||||
static constexpr int NTILES = KernelTraits::NTiles;
|
||||
#pragma unroll
|
||||
for(int itile = 0; itile < NTILES; itile++){
|
||||
clear(tCrD);
|
||||
#pragma unroll
|
||||
for(int ik = 0; ik < nk; ik++){
|
||||
int ik_next = (ik + 1) % nk;
|
||||
if(ik == nk - 1) {
|
||||
cp_async_wait<kStages - 2>();
|
||||
__syncthreads();
|
||||
ismem_read = (ismem_read + 1) % kStages;
|
||||
}
|
||||
cute::copy(s2r_tiled_copy_a, tAsA(_, _, ik_next, ismem_read), tCrA_view(_, _, ik_next));
|
||||
cute::copy(s2r_tiled_copy_b, tBsB(_, _, ik_next, ismem_read), tCrB_view(_, _, ik_next));
|
||||
if(ik == 0){
|
||||
if(itile_to_read < NTILES){
|
||||
for (size_t m = 0; m < size<1>(tAsA_copy); m++)
|
||||
{
|
||||
for (size_t k = 0; k < size<2>(tAsA_copy); k++)
|
||||
{
|
||||
if(get<0>(tAcA(0, m, k)) < residual){
|
||||
cute::copy(g2s_tiled_copy, tAgA_copy(_, m, k, itile_to_read), tAsA_copy(_, m, k, ismem_write));
|
||||
}
|
||||
}
|
||||
}
|
||||
cute::copy(g2s_tiled_copy, tBgB_copy(_, _, _, itile_to_read), tBsB_copy(_, _, _, ismem_write));
|
||||
if(idx < KernelTraits::THREADS_SFA_COPY && (BM * iy + idx * KernelTraits::SFA_ELEMS_PER_COPY < M)) {
|
||||
copy_1d((float*)&gSFA(0, 0, itile_to_read) + idx * KernelTraits::SFA_ELEMS_PER_COPY, (float*)&sSFA(0, ismem_write) + idx*KernelTraits::SFA_ELEMS_PER_COPY);
|
||||
}
|
||||
++itile_to_read;
|
||||
ismem_write = (ismem_write + 1) % kStages;
|
||||
}
|
||||
cp_async_fence();
|
||||
}
|
||||
cute::gemm(tiled_mma, tCrD, tCrA(_, _, ik), tCrB(_, _, ik), tCrD);
|
||||
}
|
||||
|
||||
int sf_ind = itile / KernelTraits::TILES_PER_BLOCK;
|
||||
float sfb_val = sfb[sf_ind];
|
||||
#pragma unroll
|
||||
for(int i = 0; i < size<1>(tCrD); i++){ // (MMA, MMA_M, MMA_N) = (4, 4, 4)
|
||||
float sfa_val_1 = sSFA(sfa_tv(idx) + i * KernelTraits::MMA_WARP_M, ismem_read_sfa);
|
||||
float sfa_val_2 = sSFA(sfa_tv(idx) + 8 + i * KernelTraits::MMA_WARP_M, ismem_read_sfa);
|
||||
#pragma unroll
|
||||
for(int j = 0; j < size<2>(tCrD); j++){
|
||||
tCrD_fp32(0, i, j) += sfa_val_1 * sfb_val * float(tCrD(0, i, j));
|
||||
tCrD_fp32(1, i, j) += sfa_val_1 * sfb_val * float(tCrD(1, i, j));
|
||||
tCrD_fp32(2, i, j) += sfa_val_2 * sfb_val * float(tCrD(2, i, j));
|
||||
tCrD_fp32(3, i, j) += sfa_val_2 * sfb_val * float(tCrD(3, i, j));
|
||||
}
|
||||
}
|
||||
ismem_read_sfa = (ismem_read_sfa + 1) % kStages;
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
auto tCrBias = make_tensor<float>(Layout<Shape<_2, Int<size<2>(tCrD_fp32)>>>{});
|
||||
auto bias_threads = typename KernelTraits::BiasThreadLayout{};
|
||||
if constexpr (has_bias){
|
||||
#pragma unroll
|
||||
for(int i = 0; i<size<2>(tCrD_fp32); i++){
|
||||
tCrBias(0, i) = sBias(bias_threads(idx) + i * KernelTraits::MMA_WARP_N);
|
||||
tCrBias(1, i) = sBias(1 + bias_threads(idx) + i * KernelTraits::MMA_WARP_N);
|
||||
}
|
||||
#pragma unroll
|
||||
for(int i = 0; i<size<1>(tCrD_fp32); i++){
|
||||
#pragma unroll
|
||||
for (int j = 0; j < size<2>(tCrD_fp32) ; j++)
|
||||
{
|
||||
tCrD_fp32(0, i, j) += tCrBias(0, j);
|
||||
tCrD_fp32(1, i, j) += tCrBias(1, j);
|
||||
tCrD_fp32(2, i, j) += tCrBias(0, j);
|
||||
tCrD_fp32(3, i, j) += tCrBias(1, j);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
auto sC = make_tensor(make_smem_ptr(C_shm), SmemLayoutC{});
|
||||
auto r2s_tiled_copy_c = make_tiled_copy_C(typename KernelTraits::R2SCopyAtomC{}, tiled_mma);
|
||||
auto r2s_thr_copy_c = r2s_tiled_copy_c.get_slice(idx);
|
||||
auto tCrC_r2s = r2s_thr_copy_c.retile_S(tCrD_fp32);
|
||||
auto tCsC_r2s = r2s_thr_copy_c.partition_D(sC);
|
||||
|
||||
typename KernelTraits::S2GCopyC s2g_tiled_copy_c;
|
||||
auto s2g_thr_copy_c = s2g_tiled_copy_c.get_thread_slice(idx);
|
||||
auto tCsC_s2g = s2g_thr_copy_c.partition_S(sC);
|
||||
auto tCgC_s2g = s2g_thr_copy_c.partition_D(gD);
|
||||
|
||||
int pipe = size<2>(tCsC_r2s);
|
||||
|
||||
auto cC = make_identity_tensor(make_shape(size<0>(gD), size<1>(gD)));
|
||||
auto tCcC = s2g_thr_copy_c.partition_D(cC);
|
||||
|
||||
for(int i = 0; i< size<1>(tCrC_r2s); i++){
|
||||
for(int j = 0; j < size<2>(tCrC_r2s); j+=pipe){
|
||||
for(int step = 0; step < pipe; ++step){
|
||||
auto fragment = make_tensor_like<output_t>(tCrC_r2s(_, i, j + step));
|
||||
cute::copy(tCrC_r2s(_, i, j + step), fragment);
|
||||
cute::copy(r2s_tiled_copy_c, fragment, tCsC_r2s(_, 0, step));
|
||||
}
|
||||
__syncthreads();
|
||||
if (get<0>(tCcC(0, i, j / pipe)) < residual){
|
||||
cute::copy(s2g_tiled_copy_c, tCsC_s2g(_, 0, 0), tCgC_s2g(_, i, j / pipe));
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <bool has_bias, typename accum_type>
|
||||
void fp8_kernel_launch(void* Aptr, void* sfa, void* Bptr, void* sfb, void* bias_ptr, void* out, int M, int N, int K, cudaStream_t stream) {
|
||||
int TMA_ALIGNED_M = ((M + sizeof(float) - 1) / sizeof(float)) * sizeof(float); // SIZEOF(float) = 4
|
||||
BLOCK_K_SWITCH(K_, M_SWITCH(
|
||||
using KernelTraits = gemm_traits<BM, BN, 3, K_, WARP_ROW, WARP_COL, has_bias, accum_type, bfloat16_t>;
|
||||
auto kernel = &gemm_fp8_kernel<KernelTraits>;
|
||||
int BX = (N + KernelTraits::BN - 1) / KernelTraits::BN;
|
||||
int BY = (M + KernelTraits::BM - 1) / KernelTraits::BM;
|
||||
dim3 block(KernelTraits::NUM_THREADS);
|
||||
dim3 gridDim(BX, BY);
|
||||
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, KernelTraits::SmemSize);
|
||||
kernel<<<gridDim, KernelTraits::NUM_THREADS, KernelTraits::SmemSize, stream>>>((float_e4m3_t*)Aptr, (float*)sfa, (float_e4m3_t*)Bptr, (float*)sfb, (float*)bias_ptr, out, M, N, K, TMA_ALIGNED_M);
|
||||
C10_CUDA_KERNEL_LAUNCH_CHECK();))
|
||||
}
|
||||
|
||||
template<bool use_fast_accum>
|
||||
void fp8_bias_gemm_cuda(void* Aptr, void* SFA, void* Bptr, void* SFB, void* bias_ptr, void* out, int M, int N, int K, cudaStream_t stream){
|
||||
using accum_type = std::conditional_t<use_fast_accum, half_t, float>;
|
||||
fp8_kernel_launch<true, accum_type>(Aptr, SFA, Bptr, SFB, bias_ptr, out, M, N, K, stream);
|
||||
}
|
||||
// template<bool use_fast_accum>
|
||||
// void fp8_gemm_cuda(void* Aptr, void* SFA, void* Bptr, void* SFB, void* out, int M, int N, int K, cudaStream_t stream){
|
||||
// using accum_type = std::conditional_t<use_fast_accum, half_t, float>;
|
||||
// BLOCK_K_SWITCH(num_acc_upcast_steps, fp8_kernel_launch<false, num_acc_upcast_steps, accum_type>(Aptr, SFA, Bptr, SFB, nullptr, out, M, N, K, stream);)
|
||||
// }
|
||||
|
||||
// template void fp8_gemm_cuda<true>(void* Aptr, void* SFA, void* Bptr, void* SFB, void* out, int M, int N, int K, cudaStream_t stream);
|
||||
// template void fp8_gemm_cuda<false>(void* Aptr, void* SFA, void* Bptr, void* SFB, void* out, int M, int N, int K, cudaStream_t stream);
|
||||
|
||||
template void fp8_bias_gemm_cuda<true>(void* Aptr, void* SFA, void* Bptr, void* SFB, void* bias_ptr, void* out, int M, int N, int K, cudaStream_t stream);
|
||||
template void fp8_bias_gemm_cuda<false>(void* Aptr, void* SFA, void* Bptr, void* SFB, void* bias_ptr, void* out, int M, int N, int K, cudaStream_t stream);
|
||||
}; // namespace sm89
|
||||
@@ -0,0 +1,121 @@
|
||||
#pragma once
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/layout/layout.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
|
||||
#include "mma_sm89_fp16.hpp"
|
||||
#include "mma_traits_sm89_fp16.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template<int BYTES> struct BytesToType {};
|
||||
template<> struct BytesToType<4> {
|
||||
using Type = uint32_t;
|
||||
static_assert(sizeof(Type) == 4);
|
||||
};
|
||||
template<> struct BytesToType<2> {
|
||||
using Type = uint16_t;
|
||||
static_assert(sizeof(Type) == 2);
|
||||
};
|
||||
|
||||
template<int BM_, int BN_, int KStages_, int K_, int WARP_ROW_=2, int WARP_COL_=2, bool HasBias_=false, typename accum_t_=cutlass::half_t, typename out_t_=cutlass::bfloat16_t>
|
||||
struct gemm_traits {
|
||||
static constexpr int BLOCK_SIZE = 128;
|
||||
static constexpr int K = K_;
|
||||
static constexpr int BM = BM_;
|
||||
static constexpr int BN = BN_;
|
||||
static constexpr int BK = 128;
|
||||
static constexpr int TILES_PER_BLOCK = BLOCK_SIZE / BK;
|
||||
static constexpr int NUM_SFB_PER_STEP = BN / BLOCK_SIZE;
|
||||
static constexpr int NTiles = K / BK;
|
||||
static constexpr int KSF = K / BLOCK_SIZE;
|
||||
static constexpr int KStages = KStages_;
|
||||
static constexpr int WARP_ROW = WARP_ROW_;
|
||||
static constexpr int WARP_COL = WARP_COL_;
|
||||
static constexpr int NUM_WARPS = WARP_ROW * WARP_COL;
|
||||
static constexpr int NUM_THREADS = NUM_WARPS * 32;
|
||||
static constexpr int MMA_WARP_M = WARP_ROW * 16;
|
||||
static constexpr int MMA_WARP_N = WARP_COL * 8;
|
||||
static constexpr int MMA_WARP_K = 32;
|
||||
using accum_t = accum_t_;
|
||||
using out_t = out_t_;
|
||||
using SwizzleLayoutO = std::conditional_t<
|
||||
std::is_same_v<out_t_, cutlass::bfloat16_t>,
|
||||
Swizzle<3, 3, 3>,
|
||||
Swizzle<2, 4, 3>
|
||||
>;
|
||||
using SwizzleLayoutAB = Swizzle<2, 4, 3>;
|
||||
using MMA_Atom_SM89 = std::conditional_t<
|
||||
std::is_same_v<accum_t, cutlass::half_t>,
|
||||
MMA_Atom<SM89_16x8x32_F16E4M3E4M3F16_TN>,
|
||||
MMA_Atom<SM89_16x8x32_F32E4M3E4M3F32_TN>
|
||||
>;
|
||||
static constexpr int INPUT_ELEMS_PER_COPY = sizeof(uint128_t) / sizeof(float_e4m3_t);
|
||||
static constexpr int OUTPUT_ELEMS_PER_COPY = sizeof(uint128_t) / sizeof(out_t_);
|
||||
static constexpr int THREADS_PER_ROW = BK / INPUT_ELEMS_PER_COPY;
|
||||
using GMEMLayout = Layout< Shape <Int<NUM_THREADS / THREADS_PER_ROW>, Int<THREADS_PER_ROW>>, Stride<Int<THREADS_PER_ROW>, _1>>;
|
||||
using G2SCopyAtom = Copy_Atom<SM80_CP_ASYNC_CACHEGLOBAL<cute::uint128_t>, float_e4m3_t>;
|
||||
using G2STiledCopy = decltype(
|
||||
make_tiled_copy(
|
||||
G2SCopyAtom{},
|
||||
GMEMLayout{},
|
||||
Layout<Shape<_1, Int<INPUT_ELEMS_PER_COPY>>>{}
|
||||
)
|
||||
);
|
||||
using S2RCopyAtomA = Copy_Atom<SM75_U32x4_LDSM_N, float_e4m3_t>;
|
||||
using S2RCopyAtomB = Copy_Atom<SM75_U32x2_LDSM_N, float_e4m3_t>;
|
||||
using SmemLayoutAtom = decltype(composition(
|
||||
Swizzle<2, 4, 3>{},
|
||||
make_layout(make_shape(Int<8>{}, Int<BK>{}),
|
||||
make_stride(Int<BK>{}, Int<1>{}))));
|
||||
using SmemLayoutA = decltype(
|
||||
tile_to_shape(SmemLayoutAtom{}, make_shape(Int<BM>{}, Int<BK>{}, Int<KStages>{}))
|
||||
);
|
||||
using SmemLayoutB = decltype(
|
||||
tile_to_shape(SmemLayoutAtom{}, make_shape(Int<BN>{}, Int<BK>{}, Int<KStages>{}))
|
||||
);
|
||||
using MMATile = decltype(
|
||||
make_tiled_mma(
|
||||
MMA_Atom_SM89{},
|
||||
Layout<Shape<Int<WARP_ROW>, Int<WARP_COL>, _1>>{},
|
||||
Tile<Int<MMA_WARP_M>, Int<MMA_WARP_N>, Int<MMA_WARP_K>>{}
|
||||
)
|
||||
);
|
||||
|
||||
static constexpr int ELEMS_PER_TILE = MMA_WARP_M * MMA_WARP_N;
|
||||
static constexpr int NUM_ELEMS_PER_WRITE = NUM_THREADS * sizeof(cute::uint128_t) / sizeof(out_t_);
|
||||
static constexpr int OUT_PIPE = NUM_ELEMS_PER_WRITE / ELEMS_PER_TILE;
|
||||
// using SmemLayoutC = Layout<Shape<Int<BM>, Int<BN>>, Stride<Int<BN>, Int<1>>>;
|
||||
|
||||
using SmemLayoutC = decltype(
|
||||
make_layout(
|
||||
make_shape(Int<MMA_WARP_M>{}, Int<MMA_WARP_N*OUT_PIPE>{}),
|
||||
make_stride(Int<MMA_WARP_N*OUT_PIPE>{}, Int<1>{})
|
||||
)
|
||||
);
|
||||
static constexpr int THREADS_PER_ROW_WRITE = MMA_WARP_N * OUT_PIPE / OUTPUT_ELEMS_PER_COPY;
|
||||
using R2SCopyAtomC = Copy_Atom<UniversalCopy<typename BytesToType<2*sizeof(out_t)>::Type>, out_t>;
|
||||
using S2GCopyAtomC = Copy_Atom<UniversalCopy<cute::uint128_t>, out_t>;
|
||||
using S2GCopyC = decltype(make_tiled_copy(S2GCopyAtomC{},
|
||||
make_layout(make_shape(Int<NUM_THREADS / THREADS_PER_ROW_WRITE>{}, Int<THREADS_PER_ROW_WRITE>{}),
|
||||
make_stride(Int<THREADS_PER_ROW_WRITE>{}, Int<1>{})),
|
||||
make_layout(make_shape(Int<1>{}, Int<OUTPUT_ELEMS_PER_COPY>{}))));
|
||||
|
||||
using G2SBiasCopyAtom = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<float>, float>;
|
||||
using G2SBiasCopy = decltype(make_tiled_copy(G2SBiasCopyAtom{}, make_layout(
|
||||
make_shape(Int<1>{},Int<BN>{}), make_stride(Int<BN>{}, Int<1>{})),
|
||||
make_layout(make_shape(Int<1>{},Int<1>{}), make_stride(Int<1>{}, Int<1>{}))));
|
||||
using sfa_copy_vtype = float;
|
||||
static constexpr int SFA_ELEMS_PER_COPY = sizeof(sfa_copy_vtype)/sizeof(float);
|
||||
static constexpr int THREADS_SFA_COPY = BM * sizeof(float) / sizeof(sfa_copy_vtype);
|
||||
// using G2SSFACopyAtom = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<cute::uint128_t>, float>;
|
||||
static constexpr bool HasBias = HasBias_;
|
||||
using SmemLayoutBias = Layout<Shape<Int<1>, Int<BN>>, Stride<Int<BN>, Int<1>>>;
|
||||
using SmemLayoutSFA = Layout<Shape<Int<BM>, Int<KStages>>, Stride<Int<1>, Int<BM>>>;
|
||||
using BiasThreadLayout = Layout<Shape<Shape<_4, _8>, Shape<Int<WARP_ROW>, Int<WARP_COL>>>, Stride<Stride<_2, _0>, Stride<_0, _8>>>;
|
||||
using SFAThreadLayout = Layout<Shape<Shape<_4, _8>, Shape<Int<WARP_ROW>, Int<WARP_COL>>>, Stride<Stride<_0, _1>, Stride<_16, _0>>>;
|
||||
static constexpr int SmemSize = cute::max(cute::cosize(SmemLayoutA{})+cute::cosize(SmemLayoutB{}), cute::cosize(SmemLayoutC{})*sizeof(out_t)) + cute::cosize(SmemLayoutBias{}) * sizeof(float) + cute::cosize(SmemLayoutSFA{})*sizeof(float);
|
||||
};
|
||||
@@ -0,0 +1,84 @@
|
||||
#pragma once
|
||||
|
||||
#include <cute/config.hpp>
|
||||
#include <cute/arch/mma.hpp>
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#if (__CUDACC_VER_MAJOR__ > 12) || (__CUDACC_VER_MAJOR__ == 12 && __CUDACC_VER_MINOR__ >= 4)
|
||||
# define CUTE_ARCH_MMA_F32_SM89_SUPPORTED
|
||||
#endif
|
||||
|
||||
#if (__CUDACC_VER_MAJOR__ > 12) || (__CUDACC_VER_MAJOR__ == 12 && __CUDACC_VER_MINOR__ >= 8)
|
||||
# define CUTE_ARCH_MMA_F16_SM89_SUPPORTED
|
||||
#endif
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 890)
|
||||
# if defined(CUTE_ARCH_MMA_F32_SM89_SUPPORTED)
|
||||
# define CUTE_ARCH_MMA_F32_SM89_ENABLED
|
||||
# endif
|
||||
|
||||
# if defined(CUTE_ARCH_MMA_F16_SM89_SUPPORTED)
|
||||
# define CUTE_ARCH_MMA_F16_SM89_ENABLED
|
||||
# endif
|
||||
#endif
|
||||
|
||||
namespace cute {
|
||||
struct SM89_16x8x32_F32E4M3E4M3F32_TN
|
||||
{
|
||||
using DRegisters = float[4];
|
||||
using ARegisters = uint32_t[4];
|
||||
using BRegisters = uint32_t[2];
|
||||
using CRegisters = float[4];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(float & d0, float & d1, float & d2, float & d3,
|
||||
uint32_t const& a0, uint32_t const& a1, uint32_t const& a2, uint32_t const& a3,
|
||||
uint32_t const& b0, uint32_t const& b1,
|
||||
float const& c0, float const& c1, float const& c2, float const& c3)
|
||||
{
|
||||
#if defined(CUTE_ARCH_MMA_F32_SM89_ENABLED)
|
||||
asm(
|
||||
"mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 "
|
||||
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};\n"
|
||||
: "=f"(d0), "=f"(d1), "=f"(d2), "=f"(d3)
|
||||
:
|
||||
"r"(a0), "r"(a1), "r"(a2), "r"(a3),
|
||||
"r"(b0), "r"(b1),
|
||||
"f"(c0), "f"(c1), "f"(c2), "f"(c3)
|
||||
);
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM89_16x8x32_F32E4M3E4M3F32_TN without CUTE_ARCH_MMA_F32_SM89_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
// MMA 16x8x32 TN
|
||||
struct SM89_16x8x32_F16E4M3E4M3F16_TN
|
||||
{
|
||||
using DRegisters = uint32_t[2];
|
||||
using ARegisters = uint32_t[4];
|
||||
using BRegisters = uint32_t[2];
|
||||
using CRegisters = uint32_t[2];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint32_t & d0, uint32_t & d1,
|
||||
uint32_t const& a0, uint32_t const& a1, uint32_t const& a2, uint32_t const& a3,
|
||||
uint32_t const& b0, uint32_t const& b1,
|
||||
uint32_t const& c0, uint32_t const& c1)
|
||||
{
|
||||
#if defined(CUTE_ARCH_MMA_F16_SM89_ENABLED)
|
||||
asm(
|
||||
"mma.sync.aligned.m16n8k32.row.col.f16.e4m3.e4m3.f16 "
|
||||
"{%0,%1}, {%2,%3,%4,%5}, {%6,%7}, {%8,%9};\n"
|
||||
: "=r"(d0), "=r"(d1)
|
||||
:
|
||||
"r"(a0), "r"(a1), "r"(a2), "r"(a3),
|
||||
"r"(b0), "r"(b1),
|
||||
"r"(c0), "r"(c1)
|
||||
);
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM89_16x8x32_F32E4M3E4M3F32_TN without CUTE_ARCH_MMA_F16_SM89_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
#pragma once
|
||||
|
||||
#include <cute/atom/mma_traits.hpp>
|
||||
#include <cute/layout.hpp>
|
||||
#include <cute/numeric/numeric_types.hpp>
|
||||
#include "mma_sm89_fp16.hpp"
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
namespace {
|
||||
|
||||
// (T32,V4) -> (M16,N8)
|
||||
using SM80_16x8_Row = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
|
||||
|
||||
}
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM89_16x8x32_F32E4M3E4M3F32_TN> {
|
||||
using ValTypeD = float;
|
||||
using ValTypeA = float_e4m3_t;
|
||||
using ValTypeB = float_e4m3_t;
|
||||
using ValTypeC = float;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_32>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _4,_2, _2>>,
|
||||
Stride<Stride<_64,_1>,Stride<_16,_8,_256>>>;
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_4, _2>>,
|
||||
Stride<Stride<_32,_1>,Stride<_8,_128>>>;
|
||||
using CLayout = SM80_16x8_Row;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM89_16x8x32_F16E4M3E4M3F16_TN> {
|
||||
using ValTypeD = half_t;
|
||||
using ValTypeA = float_e4m3_t;
|
||||
using ValTypeB = float_e4m3_t;
|
||||
using ValTypeC = half_t;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_32>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _4,_2, _2>>,
|
||||
Stride<Stride<_64,_1>,Stride<_16,_8,_256>>>;
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_4, _2>>,
|
||||
Stride<Stride<_32,_1>,Stride<_8,_128>>>;
|
||||
using CLayout = SM80_16x8_Row;
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
#pragma once
|
||||
|
||||
#define BOOL_SWITCH(COND, CONST_NAME, ...) \
|
||||
if (COND) { \
|
||||
constexpr static bool CONST_NAME = true; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
constexpr static bool CONST_NAME = false; \
|
||||
__VA_ARGS__ \
|
||||
}
|
||||
//K/128
|
||||
#define BLOCK_K_SWITCH(COSNT_NAME, ...) \
|
||||
if (K == 2048) { \
|
||||
constexpr static int COSNT_NAME = 2048; \
|
||||
__VA_ARGS__ \
|
||||
} \
|
||||
else if (K == 4096) { \
|
||||
constexpr static int COSNT_NAME = 4096; \
|
||||
__VA_ARGS__ \
|
||||
} else if (K == 8192) { \
|
||||
constexpr static int COSNT_NAME = 8192; \
|
||||
__VA_ARGS__ \
|
||||
} else if (K == 16384) { \
|
||||
constexpr static int COSNT_NAME = 16384; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
TORCH_CHECK(false, "Unsupported K value: ", K); \
|
||||
}
|
||||
|
||||
#define M_SWITCH(...) \
|
||||
constexpr static int BM = 64; \
|
||||
constexpr static int BN = 128; \
|
||||
constexpr static int WARP_ROW = 2; \
|
||||
constexpr static int WARP_COL = 4; \
|
||||
__VA_ARGS__
|
||||
@@ -0,0 +1,206 @@
|
||||
#pragma once
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#include <torch/python.h>
|
||||
#include "exceptions.hpp"
|
||||
|
||||
namespace blockwise {
|
||||
template <typename T>
|
||||
static T ceil_div(const T& a, const T& b) {
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
template <typename T>
|
||||
static constexpr T align(const T& a, const T& b) {
|
||||
return ceil_div(a, b) * b;
|
||||
}
|
||||
|
||||
static int get_tma_aligned_size(const int& x, const int& element_size) {
|
||||
constexpr int kNumTMAAlignmentBytes = 16;
|
||||
DG_HOST_ASSERT(kNumTMAAlignmentBytes % element_size == 0);
|
||||
return align(x, kNumTMAAlignmentBytes / element_size);
|
||||
}
|
||||
static std::pair<int, int> get_inner_outer_dims(const cute::UMMA::Major& major, const int& k, const int& mn) {
|
||||
return major == cute::UMMA::Major::K ? std::make_pair(k, mn) : std::make_pair(mn, k);
|
||||
}
|
||||
|
||||
static int get_non_contiguous_dim(const cute::UMMA::Major& major) {
|
||||
return major == cute::UMMA::Major::K ? -2 : -1;
|
||||
}
|
||||
|
||||
static int get_compiled_dim(const int& dim, const char& name, const std::string& compiled_dims) {
|
||||
for (const char& c: compiled_dims) {
|
||||
if (name == c)
|
||||
return dim;
|
||||
}
|
||||
return 0;
|
||||
}
|
||||
static CUtensorMapDataType aten_dtype_to_tensor_map_dtype(const at::ScalarType& dtype,
|
||||
const bool& allow_tf32) {
|
||||
if (allow_tf32 and dtype == torch::kFloat)
|
||||
return CU_TENSOR_MAP_DATA_TYPE_TFLOAT32;
|
||||
|
||||
switch (dtype) {
|
||||
case torch::kInt: return CU_TENSOR_MAP_DATA_TYPE_INT32;
|
||||
case torch::kFloat: return CU_TENSOR_MAP_DATA_TYPE_FLOAT32;
|
||||
case torch::kBFloat16: return CU_TENSOR_MAP_DATA_TYPE_BFLOAT16;
|
||||
case torch::kFloat8_e4m3fn: return CU_TENSOR_MAP_DATA_TYPE_UINT8;
|
||||
default: DG_HOST_UNREACHABLE("Unsupported dtype");
|
||||
}
|
||||
}
|
||||
|
||||
static CUtensorMapSwizzle mode_into_tensor_map_swizzle(const int& mode, const int& base) {
|
||||
#if CUDA_VERSION >= 12080
|
||||
if (base != 0) {
|
||||
DG_HOST_ASSERT(base == 32 and mode == 128);
|
||||
return CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B;
|
||||
}
|
||||
#endif
|
||||
|
||||
DG_HOST_ASSERT(base == 0);
|
||||
switch (mode) {
|
||||
case 0:
|
||||
case 16: return CU_TENSOR_MAP_SWIZZLE_NONE;
|
||||
case 32: return CU_TENSOR_MAP_SWIZZLE_32B;
|
||||
case 64: return CU_TENSOR_MAP_SWIZZLE_64B;
|
||||
case 128: return CU_TENSOR_MAP_SWIZZLE_128B;
|
||||
default: DG_HOST_UNREACHABLE("Unsupported swizzling mode");
|
||||
}
|
||||
}
|
||||
|
||||
static CUtensorMap make_tma_2d_desc(const torch::Tensor& t,
|
||||
int gmem_inner_dim, int gmem_outer_dim,
|
||||
int smem_inner_dim, int smem_outer_dim,
|
||||
const int& gmem_outer_stride,
|
||||
const int& swizzle_mode, const int& swizzle_base = 0,
|
||||
const bool& allow_tf32 = false) {
|
||||
const auto& elem_size = static_cast<int>(t.element_size());
|
||||
if (swizzle_mode != 0)
|
||||
smem_inner_dim = swizzle_mode / elem_size;
|
||||
|
||||
CUtensorMap tensor_map;
|
||||
const cuuint64_t gmem_dims[2] = {static_cast<cuuint64_t>(gmem_inner_dim), static_cast<cuuint64_t>(gmem_outer_dim)};
|
||||
const cuuint32_t smem_dims[2] = {static_cast<cuuint32_t>(smem_inner_dim), static_cast<cuuint32_t>(smem_outer_dim)};
|
||||
const cuuint64_t gmem_strides[1] = {static_cast<cuuint64_t>(gmem_outer_stride * elem_size), };
|
||||
const cuuint32_t elem_strides[2] = {1, 1};
|
||||
// if (get_env<int>("DG_JIT_DEBUG")) {
|
||||
// printf("Making TMA desc: global memory: %d %d, shared memory: %d %d, outer stride: %d, swizzle: %d (base: %d), elem size: %d\n",
|
||||
// gmem_inner_dim, gmem_outer_dim, smem_inner_dim, smem_outer_dim,
|
||||
// gmem_outer_stride, swizzle_mode, swizzle_base, elem_size);
|
||||
// }
|
||||
cuTensorMapEncodeTiled(
|
||||
&tensor_map, aten_dtype_to_tensor_map_dtype(t.scalar_type(), allow_tf32),
|
||||
2, t.data_ptr(), gmem_dims, gmem_strides, smem_dims, elem_strides,
|
||||
CU_TENSOR_MAP_INTERLEAVE_NONE, mode_into_tensor_map_swizzle(swizzle_mode, swizzle_base),
|
||||
CU_TENSOR_MAP_L2_PROMOTION_L2_256B, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
|
||||
return tensor_map;
|
||||
}
|
||||
|
||||
static CUtensorMap make_tma_3d_desc(const torch::Tensor& t,
|
||||
const int& gmem_dim_0, const int& gmem_dim_1, const int& gmem_dim_2,
|
||||
const int& smem_dim_0, const int& smem_dim_1, const int& smem_dim_2,
|
||||
const int& gmem_stride_0, const int& gmem_stride_1,
|
||||
const int& swizzle_mode, const int& swizzle_base = 0,
|
||||
const bool& allow_tf32 = false) {
|
||||
const auto& elem_size = static_cast<int>(t.element_size());
|
||||
if (swizzle_mode != 0)
|
||||
DG_HOST_ASSERT(smem_dim_0 == swizzle_mode / elem_size);
|
||||
|
||||
CUtensorMap tensor_map;
|
||||
const cuuint64_t gmem_dims[3] = {static_cast<cuuint64_t>(gmem_dim_0), static_cast<cuuint64_t>(gmem_dim_1), static_cast<cuuint64_t>(gmem_dim_2),};
|
||||
const cuuint32_t smem_dims[3] = {static_cast<cuuint32_t>(smem_dim_0), static_cast<cuuint32_t>(smem_dim_1), static_cast<cuuint32_t>(smem_dim_2)};
|
||||
const cuuint64_t gmem_strides[2] = {static_cast<cuuint64_t>(gmem_stride_0 * elem_size), static_cast<cuuint64_t>(gmem_stride_1 * elem_size)};
|
||||
const cuuint32_t elem_strides[3] = {1, 1, 1};
|
||||
// if (get_env<int>("DG_JIT_DEBUG")) {
|
||||
// printf("Making 3D TMA desc: global memory: %d %d %d, shared memory: %d %d %d, outer stride: %d %d, swizzle: %d, elem size: %d\n",
|
||||
// gmem_dim_0, gmem_dim_1, gmem_dim_2, smem_dim_0, smem_dim_1, smem_dim_2,
|
||||
// gmem_stride_0, gmem_stride_1, swizzle_mode, elem_size);
|
||||
// }
|
||||
cuTensorMapEncodeTiled(
|
||||
&tensor_map, aten_dtype_to_tensor_map_dtype(t.scalar_type(), allow_tf32),
|
||||
3, t.data_ptr(), gmem_dims, gmem_strides, smem_dims, elem_strides,
|
||||
CU_TENSOR_MAP_INTERLEAVE_NONE, mode_into_tensor_map_swizzle(swizzle_mode, swizzle_base),
|
||||
CU_TENSOR_MAP_L2_PROMOTION_L2_256B, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
|
||||
return tensor_map;
|
||||
}
|
||||
|
||||
static CUtensorMap make_tma_a_desc(const cute::UMMA::Major& major,
|
||||
const torch::Tensor& t,
|
||||
const int& shape_m, const int& shape_k,
|
||||
const int& block_m, const int& block_k,
|
||||
const int& outer_stride,
|
||||
const int& num_groups,
|
||||
const int& swizzle_mode, const int& swizzle_base = 0,
|
||||
const bool& allow_tf32 = false) {
|
||||
if (num_groups > 1)
|
||||
DG_HOST_ASSERT(major == cute::UMMA::Major::K);
|
||||
const auto& [gmem_inner_dim, gmem_outer_dim] = get_inner_outer_dims(major, shape_k, shape_m * num_groups);
|
||||
const auto& [smem_inner_dim, smem_outer_dim] = get_inner_outer_dims(major, block_k, block_m);
|
||||
return make_tma_2d_desc(t,
|
||||
gmem_inner_dim, gmem_outer_dim,
|
||||
smem_inner_dim, smem_outer_dim,
|
||||
outer_stride,
|
||||
swizzle_mode, swizzle_base,
|
||||
allow_tf32);
|
||||
}
|
||||
|
||||
static CUtensorMap make_tma_b_desc(const cute::UMMA::Major& major,
|
||||
const torch::Tensor& t,
|
||||
const int& shape_n, const int& shape_k,
|
||||
const int& block_n, const int& block_k,
|
||||
const int& outer_stride,
|
||||
const int& num_groups,
|
||||
const int& swizzle_mode, const int& swizzle_base = 0,
|
||||
const bool& allow_tf32 = false) {
|
||||
const auto& [gmem_inner_dim, gmem_outer_dim] = get_inner_outer_dims(major, shape_k, shape_n);
|
||||
const auto& [smem_inner_dim, smem_outer_dim] = get_inner_outer_dims(major, block_k, block_n);
|
||||
|
||||
// `num_groups` is always applied into the outer dimensions
|
||||
return make_tma_2d_desc(t,
|
||||
gmem_inner_dim, gmem_outer_dim * num_groups,
|
||||
smem_inner_dim, smem_outer_dim,
|
||||
outer_stride,
|
||||
swizzle_mode, swizzle_base,
|
||||
allow_tf32);
|
||||
}
|
||||
|
||||
static CUtensorMap make_tma_cd_desc(const torch::Tensor& t,
|
||||
const int& shape_m, const int& shape_n,
|
||||
const int& block_m, const int& block_n,
|
||||
const int& outer_stride,
|
||||
const int& num_groups,
|
||||
const int& swizzle_mode, const int& swizzle_base = 0,
|
||||
const bool& allow_tf32 = false) {
|
||||
// Swizzling requires the inner box dim to be less or equal than `kSwizzleCDMode`
|
||||
// bytes, so `BLOCK_N * sizeof(T) / kSwizzleCDMode` TMA stores are required
|
||||
return make_tma_2d_desc(t,
|
||||
shape_n, shape_m * num_groups,
|
||||
block_n, block_m,
|
||||
outer_stride,
|
||||
swizzle_mode, swizzle_base,
|
||||
allow_tf32);
|
||||
}
|
||||
|
||||
static CUtensorMap make_tma_sf_desc(const cute::UMMA::Major& major,
|
||||
const torch::Tensor& t,
|
||||
int shape_mn, int shape_k,
|
||||
const int& block_mn, const int& block_k,
|
||||
const int& num_groups,
|
||||
const int& swizzle_mode, const int& swizzle_base = 0,
|
||||
const bool& allow_tf32 = false) {
|
||||
DG_HOST_ASSERT(major == cute::UMMA::Major::MN);
|
||||
|
||||
// TODO: maybe swizzle SF as well
|
||||
DG_HOST_ASSERT(swizzle_mode == 0);
|
||||
|
||||
shape_mn = get_tma_aligned_size(shape_mn, static_cast<int>(t.element_size()));
|
||||
return make_tma_2d_desc(t,
|
||||
shape_mn, ceil_div(shape_k, block_k * (t.scalar_type() == torch::kFloat ? 1 : 4)) * num_groups,
|
||||
block_mn, 1,
|
||||
shape_mn,
|
||||
swizzle_mode, swizzle_base,
|
||||
allow_tf32);
|
||||
}
|
||||
|
||||
} // namespace deep_gemm
|
||||
@@ -0,0 +1,39 @@
|
||||
#pragma once
|
||||
#include <cuda.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <nvrtc.h>
|
||||
|
||||
#include <torch/python.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
|
||||
#include "kernels/geforce/static_switch.h"
|
||||
|
||||
namespace sm89 {
|
||||
template<bool use_fast_accum>
|
||||
void fp8_bias_gemm_cuda(void* Aptr, void* SFA, void* Bptr, void* SFB, void* bias_ptr, void* out, int M, int N, int K, cudaStream_t stream);
|
||||
}
|
||||
|
||||
namespace blockwise {
|
||||
static void sm89_fp8_gemm_1d2d_bias(const torch::Tensor& a, const torch::Tensor& sfa,
|
||||
const torch::Tensor& b, const torch::Tensor& sfb,
|
||||
const torch::Tensor& bias,
|
||||
const torch::Tensor& d,
|
||||
const int& m, const int& n, const int& k,
|
||||
const bool use_fast_accum) {
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
if (use_fast_accum) {
|
||||
sm89::fp8_bias_gemm_cuda<true>(
|
||||
a.data_ptr(), sfa.data_ptr(),
|
||||
b.data_ptr(), sfb.data_ptr(),
|
||||
bias.data_ptr(), d.data_ptr(),
|
||||
m, n, k, stream);
|
||||
} else {
|
||||
sm89::fp8_bias_gemm_cuda<false>(
|
||||
a.data_ptr(), sfa.data_ptr(),
|
||||
b.data_ptr(), sfb.data_ptr(),
|
||||
bias.data_ptr(), d.data_ptr(),
|
||||
m, n, k, stream);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
|
||||
|
||||
#pragma once
|
||||
#include <cuda.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <nvrtc.h>
|
||||
|
||||
#include <torch/python.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
|
||||
#include <cute/arch/mma_sm100_desc.hpp>
|
||||
#include "runtime_utils.hpp"
|
||||
|
||||
#include "config.hpp"
|
||||
#include "static_switch.hpp"
|
||||
|
||||
namespace deep_gemm{
|
||||
template<int N, int K>
|
||||
void sm90_fp8_gemm_1d2d_bias_launch(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
|
||||
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
|
||||
const CUtensorMap tensor_map_a,
|
||||
const CUtensorMap tensor_map_b,
|
||||
const CUtensorMap tensor_map_d,
|
||||
const CUtensorMap tensor_map_sfa);
|
||||
};
|
||||
|
||||
namespace blockwise{
|
||||
|
||||
static void sm90_fp8_gemm_1d2d_bias(const torch::Tensor& a, const torch::Tensor& sfa,
|
||||
const torch::Tensor& b, const torch::Tensor& sfb,
|
||||
const torch::Tensor& bias,
|
||||
const std::optional<torch::Tensor>& c,
|
||||
const torch::Tensor& d,
|
||||
const int& m, const int& n, const int& k, const int num_sms) {
|
||||
// DG_HOST_ASSERT(not c.has_value() and d.scalar_type() == torch::kBFloat16);
|
||||
const auto& config = GemmConfig<90>();
|
||||
|
||||
// Requires no TMA splits
|
||||
// DG_HOST_ASSERT(config.smem_config.swizzle_a_mode == config.block_k);
|
||||
// DG_HOST_ASSERT(config.smem_config.swizzle_b_mode == config.block_k);
|
||||
int smem_size = k == 16384 || k == 8192 ? 216624 : config.smem_config.smem_size;
|
||||
const auto& tensor_map_a = make_tma_a_desc(cute::UMMA::Major::K, a, m, k,
|
||||
config.block_m,
|
||||
config.block_k,
|
||||
static_cast<int>(a.stride(-2)), 1,
|
||||
config.smem_config.swizzle_a_mode);
|
||||
const auto& tensor_map_b = make_tma_b_desc(cute::UMMA::Major::K, b, n, k,
|
||||
config.block_n,
|
||||
config.block_k,
|
||||
static_cast<int>(b.stride(-2)), 1,
|
||||
config.smem_config.swizzle_b_mode);
|
||||
const auto& tensor_map_d = make_tma_cd_desc(d, m, static_cast<int>(d.size(-1)),
|
||||
config.block_m,
|
||||
config.block_n,
|
||||
static_cast<int>(d.stride(-2)), 1,
|
||||
config.smem_config.swizzle_cd_mode);
|
||||
const auto& tensor_map_sfa = make_tma_sf_desc(cute::UMMA::Major::MN, sfa, m, k,
|
||||
config.block_m, config.block_k, 1, 0);
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
// Launch
|
||||
DIM_SWITCH(k, K,
|
||||
DIM_SWITCH(n, N,
|
||||
deep_gemm::sm90_fp8_gemm_1d2d_bias_launch<N, K>(num_sms, config.thread_config.num_threads, config.multicast_config.num_multicast, smem_size, stream, (float*)sfb.data_ptr(), (float*)bias.data_ptr(), nullptr, m, n, k, tensor_map_a, tensor_map_b, tensor_map_d, tensor_map_sfa);)
|
||||
)
|
||||
}
|
||||
};
|
||||
@@ -0,0 +1,60 @@
|
||||
#pragma once
|
||||
|
||||
#define DIM_SWITCH(VAR_NAME, CONST_NAME, ...) \
|
||||
if (VAR_NAME == 4096) { \
|
||||
constexpr static int CONST_NAME = 4096; \
|
||||
__VA_ARGS__ \
|
||||
} else if (VAR_NAME == 2048){ \
|
||||
constexpr static int CONST_NAME = 2048; \
|
||||
__VA_ARGS__ \
|
||||
} else if (VAR_NAME == 8192){ \
|
||||
constexpr static int CONST_NAME = 8192; \
|
||||
__VA_ARGS__ \
|
||||
} else if(VAR_NAME == 16384) { \
|
||||
constexpr static int CONST_NAME = 16384; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
TORCH_CHECK(false, "Unsupported DIM_SWITCH value: ", VAR_NAME); \
|
||||
}
|
||||
|
||||
#define BOOL_SWITCH(COND, CONST_NAME, ...) \
|
||||
if (COND) { \
|
||||
constexpr static bool CONST_NAME = true; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
constexpr static bool CONST_NAME = false; \
|
||||
__VA_ARGS__ \
|
||||
} \
|
||||
//K/128
|
||||
#define BLOCK_K_SWITCH(COSNT_NAME, ...) \
|
||||
if (K == 2048) { \
|
||||
constexpr static int COSNT_NAME = 16; \
|
||||
__VA_ARGS__ \
|
||||
} \
|
||||
else if (K == 4096) { \
|
||||
constexpr static int COSNT_NAME = 32; \
|
||||
__VA_ARGS__ \
|
||||
} else if (K == 8192) { \
|
||||
constexpr static int COSNT_NAME = 64; \
|
||||
__VA_ARGS__ \
|
||||
} else if (K == 16384) { \
|
||||
constexpr static int COSNT_NAME = 128; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
TORCH_CHECK(false, "Unsupported K value: ", K); \
|
||||
}
|
||||
|
||||
#define M_SWITCH(...) \
|
||||
if (M <= 1024) { \
|
||||
constexpr static int BM = 128; \
|
||||
constexpr static int BN = 128; \
|
||||
constexpr static int WARP_ROW = 2; \
|
||||
constexpr static int WARP_COL = 2; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
constexpr static int BM = 128; \
|
||||
constexpr static int BN = 256; \
|
||||
constexpr static int WARP_ROW = 2; \
|
||||
constexpr static int WARP_COL = 4; \
|
||||
__VA_ARGS__ \
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
#pragma once
|
||||
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp8.h>
|
||||
#include <cuda/std/cstdint>
|
||||
#include <cuda/std/utility>
|
||||
#include <cute/container/tuple.hpp>
|
||||
|
||||
#ifdef __CLION_IDE__
|
||||
|
||||
__host__ __device__ __forceinline__ void host_device_printf(const char* format, ...) {
|
||||
asm volatile("trap;");
|
||||
}
|
||||
|
||||
#define printf host_device_printf
|
||||
#endif
|
||||
|
||||
#ifndef DG_DEVICE_ASSERT
|
||||
#define DG_DEVICE_ASSERT(cond) \
|
||||
do { \
|
||||
if (not (cond)) { \
|
||||
printf("Assertion failed: %s:%d, condition: %s\n", __FILE__, __LINE__, #cond); \
|
||||
asm("trap;"); \
|
||||
} \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
#ifndef DG_TRAP_ONLY_DEVICE_ASSERT
|
||||
#define DG_TRAP_ONLY_DEVICE_ASSERT(cond) \
|
||||
do { \
|
||||
if (not (cond)) \
|
||||
asm("trap;"); \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
#ifndef DG_STATIC_ASSERT
|
||||
#define DG_STATIC_ASSERT(cond, ...) static_assert(cond, __VA_ARGS__)
|
||||
#endif
|
||||
@@ -0,0 +1,88 @@
|
||||
/**
|
||||
* @file configs.cuh
|
||||
* @brief Configuration constants and compile-time settings for ltx-kernels.
|
||||
*
|
||||
* This header defines the tunable parameters and constants used throughout
|
||||
* the ltx-kernels communication library. These values are chosen to balance
|
||||
* performance across different GPU architectures.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cstdint>
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
namespace ltx_kernels {
|
||||
// =============================================================================
|
||||
// Synchronization Configuration
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* @brief Default barrier timeout in seconds.
|
||||
*
|
||||
* If a barrier wait exceeds this timeout, the kernel traps to indicate a deadlock or
|
||||
* communication failure. All2All converts it to clock cycles using the device's peak SM clock
|
||||
* (cudaDeviceGetAttribute(cudaDevAttrClockRate)), so the wall-clock guard holds regardless of GPU.
|
||||
*/
|
||||
constexpr double DEFAULT_BARRIER_TIMEOUT_SECONDS = 10.0;
|
||||
|
||||
// =============================================================================
|
||||
// Hardware Limits
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* @brief Maximum number of peer GPUs supported for IPC communication.
|
||||
*
|
||||
* This limits the size of static arrays for buffer pointers and barrier signals.
|
||||
* Set to 8 to support up to 8-way tensor parallelism (common for DGX systems).
|
||||
*/
|
||||
constexpr int MAX_NUM_PEERS = 8;
|
||||
|
||||
// =============================================================================
|
||||
// Kernel Configuration
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* @brief Default number of threads per block for All2All kernels.
|
||||
*
|
||||
* Used by send_recv_all2all and gather_heads kernels. The value 512 provides
|
||||
* good occupancy while leaving registers for complex pointer arithmetic.
|
||||
*/
|
||||
constexpr int DEFAULT_KERNEL_THREADS = 512;
|
||||
|
||||
/**
|
||||
* @brief Number of threads per block for the AllGather kernel.
|
||||
*
|
||||
* AllGather uses more threads (1024) because its memory access pattern
|
||||
* is simpler (no head selection), allowing higher thread-level parallelism.
|
||||
*/
|
||||
constexpr int ALLGATHER_KERNEL_THREADS = 1024;
|
||||
|
||||
} // namespace ltx_kernels
|
||||
|
||||
// =============================================================================
|
||||
// Torch/CUDA Compatibility Fixes
|
||||
// =============================================================================
|
||||
|
||||
/*
|
||||
* PyTorch sometimes disables CUDA half/bfloat16 operators and conversions
|
||||
* to avoid ambiguity in template resolution. We re-enable them here since
|
||||
* our kernels explicitly handle these types.
|
||||
*/
|
||||
|
||||
#ifdef __CUDA_NO_HALF_CONVERSIONS__
|
||||
#undef __CUDA_NO_HALF_CONVERSIONS__
|
||||
#endif
|
||||
#ifdef __CUDA_NO_HALF_OPERATORS__
|
||||
#undef __CUDA_NO_HALF_OPERATORS__
|
||||
#endif
|
||||
#ifdef __CUDA_NO_HALF2_OPERATORS__
|
||||
#undef __CUDA_NO_HALF2_OPERATORS__
|
||||
#endif
|
||||
#ifdef __CUDA_NO_BFLOAT16_CONVERSIONS__
|
||||
#undef __CUDA_NO_BFLOAT16_CONVERSIONS__
|
||||
#endif
|
||||
#ifdef __CUDA_NO_BFLOAT162_OPERATORS__
|
||||
#undef __CUDA_NO_BFLOAT162_OPERATORS__
|
||||
#endif
|
||||
@@ -0,0 +1,170 @@
|
||||
/**
|
||||
* @file exceptions.cuh
|
||||
* @brief Exception handling and assertion macros for CUDA/C++ code.
|
||||
*
|
||||
* This header provides a unified exception type and assertion macros for
|
||||
* both host and device code. The macros capture file and line information
|
||||
* for easier debugging of errors.
|
||||
*
|
||||
* ## Usage Examples
|
||||
*
|
||||
* ```cpp
|
||||
* // Check CUDA API call
|
||||
* CUDA_CHECK(cudaMalloc(&ptr, size));
|
||||
*
|
||||
* // Host-side assertion
|
||||
* EP_HOST_ASSERT(tensor.is_contiguous());
|
||||
*
|
||||
* // Device-side assertion (inside kernel)
|
||||
* EP_DEVICE_ASSERT(threadIdx.x < MAX_THREADS);
|
||||
*
|
||||
* // Compile-time assertion
|
||||
* EP_STATIC_ASSERT(sizeof(int4) == 16, "int4 must be 16 bytes");
|
||||
* ```
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <exception>
|
||||
#include <string>
|
||||
|
||||
#include "configs.cuh"
|
||||
|
||||
// =============================================================================
|
||||
// Static Assertions
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* @brief Compile-time assertion macro.
|
||||
*
|
||||
* @param cond Condition that must be true at compile time
|
||||
* @param reason Human-readable error message if condition fails
|
||||
*/
|
||||
#ifndef EP_STATIC_ASSERT
|
||||
#define EP_STATIC_ASSERT(cond, reason) static_assert(cond, reason)
|
||||
#endif
|
||||
|
||||
// =============================================================================
|
||||
// Exception Type
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* @class EPException
|
||||
* @brief Custom exception type with file/line information.
|
||||
*
|
||||
* EPException captures the location (file, line) and context (name, error)
|
||||
* of the error for debugging. It inherits from std::exception for
|
||||
* compatibility with standard C++ exception handling.
|
||||
*
|
||||
* ## Message Format
|
||||
*
|
||||
* The what() message has the format:
|
||||
* "Failed: <name> error <file>:<line> '<error message>'"
|
||||
*/
|
||||
class EPException : public std::exception {
|
||||
private:
|
||||
std::string message = {}; ///< Formatted error message
|
||||
|
||||
public:
|
||||
/**
|
||||
* @brief Constructs an EPException with location and error information.
|
||||
*
|
||||
* @param name Category of error (e.g., "CUDA", "Assertion")
|
||||
* @param file Source file where error occurred (__FILE__)
|
||||
* @param line Line number where error occurred (__LINE__)
|
||||
* @param error Description of the error
|
||||
*/
|
||||
explicit EPException(const char *name, const char *file, const int line, const std::string &error) {
|
||||
message = std::string("Failed: ") + name + " error " + file + ":" + std::to_string(line) + " '" + error + "'";
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Returns the formatted error message.
|
||||
* @return C-string containing the error message
|
||||
*/
|
||||
const char *what() const noexcept override { return message.c_str(); }
|
||||
};
|
||||
|
||||
// =============================================================================
|
||||
// Runtime Assertion Macros
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* @brief Checks CUDA API return value and throws on error.
|
||||
*
|
||||
* Use this macro to wrap all CUDA runtime API calls. If the call fails,
|
||||
* an EPException is thrown with the CUDA error string.
|
||||
*
|
||||
* @param cmd CUDA API call expression
|
||||
* @throws EPException if the CUDA call returns an error
|
||||
*
|
||||
* Example:
|
||||
* ```cpp
|
||||
* CUDA_CHECK(cudaMalloc(&ptr, size));
|
||||
* CUDA_CHECK(cudaMemcpy(dst, src, size, cudaMemcpyDeviceToDevice));
|
||||
* ```
|
||||
*/
|
||||
#ifndef CUDA_CHECK
|
||||
#define CUDA_CHECK(cmd) \
|
||||
do { \
|
||||
cudaError_t e = (cmd); \
|
||||
if (e != cudaSuccess) { \
|
||||
throw EPException("CUDA", __FILE__, __LINE__, cudaGetErrorString(e)); \
|
||||
} \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
/**
|
||||
* @brief Host-side assertion that throws on failure.
|
||||
*
|
||||
* Use this for runtime checks in host code. If the condition is false,
|
||||
* an EPException is thrown with the condition as the error message.
|
||||
*
|
||||
* @param cond Condition to check (must be true)
|
||||
* @throws EPException if condition is false
|
||||
*
|
||||
* Example:
|
||||
* ```cpp
|
||||
* EP_HOST_ASSERT(tensor.dim() == 4);
|
||||
* EP_HOST_ASSERT(rank >= 0 && rank < world_size);
|
||||
* ```
|
||||
*/
|
||||
#ifndef EP_HOST_ASSERT
|
||||
#define EP_HOST_ASSERT(cond) \
|
||||
do { \
|
||||
if (not(cond)) { \
|
||||
throw EPException("Assertion", __FILE__, __LINE__, #cond); \
|
||||
} \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
/**
|
||||
* @brief Device-side assertion that traps on failure.
|
||||
*
|
||||
* Use this for runtime checks inside CUDA kernels. If the condition is
|
||||
* false, prints an error message and executes a trap instruction to
|
||||
* halt the GPU.
|
||||
*
|
||||
* @warning This causes the entire kernel to abort. Use sparingly and
|
||||
* consider removing from release builds for performance.
|
||||
*
|
||||
* @param cond Condition to check (must be true)
|
||||
*
|
||||
* Example:
|
||||
* ```cpp
|
||||
* __global__ void my_kernel(int* data, int size) {
|
||||
* int idx = threadIdx.x + blockIdx.x * blockDim.x;
|
||||
* EP_DEVICE_ASSERT(idx < size);
|
||||
* data[idx] = 42;
|
||||
* }
|
||||
* ```
|
||||
*/
|
||||
#ifndef EP_DEVICE_ASSERT
|
||||
#define EP_DEVICE_ASSERT(cond) \
|
||||
do { \
|
||||
if (not(cond)) { \
|
||||
printf("Assertion failed: %s:%d, condition: %s\n", __FILE__, __LINE__, #cond); \
|
||||
asm("trap;"); \
|
||||
} \
|
||||
} while (0)
|
||||
#endif
|
||||
@@ -0,0 +1,360 @@
|
||||
/**
|
||||
* @file utils.cuh
|
||||
* @brief Low-level CUDA utility functions for memory operations and synchronization.
|
||||
*
|
||||
* This header provides optimized PTX assembly wrappers for memory operations
|
||||
* that bypass cache hierarchy or use specific memory ordering semantics.
|
||||
* These are critical for achieving peak bandwidth in multi-GPU communication.
|
||||
*
|
||||
* ## Memory Operation Types
|
||||
*
|
||||
* - **Non-allocating stores (st_na)**: Bypass L1 cache to avoid polluting it
|
||||
* with data that won't be reused locally
|
||||
* - **Non-caching loads (ld_nc)**: Bypass L1 cache for streaming reads
|
||||
* - **Acquire/Release**: Memory ordering for synchronization
|
||||
* - **System scope (sys)**: Visibility across all GPUs, not just this one
|
||||
*
|
||||
* ## Cache Hints
|
||||
*
|
||||
* - L1::no_allocate: Don't allocate in L1 on miss (streaming pattern)
|
||||
* - L2::256B: Use 256-byte L2 cache lines
|
||||
* - volatile: Bypass all caches, always go to memory
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
#include <stdint.h>
|
||||
|
||||
// =============================================================================
|
||||
// PTX Instruction Selection
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* Store instruction macro. When DISABLE_AGGRESSIVE_PTX_INSTRS is not defined,
|
||||
* uses non-allocating stores to avoid polluting L1 cache with write-only data.
|
||||
*/
|
||||
#ifndef DISABLE_AGGRESSIVE_PTX_INSTRS
|
||||
#define ST_NA_FUNC "st.global.L1::no_allocate"
|
||||
#else
|
||||
#define ST_NA_FUNC "st.global"
|
||||
#endif
|
||||
|
||||
/**
|
||||
* Load instruction macro. When DISABLE_AGGRESSIVE_PTX_INSTRS is not defined,
|
||||
* uses non-caching loads optimized for streaming access patterns.
|
||||
*/
|
||||
#ifndef DISABLE_AGGRESSIVE_PTX_INSTRS
|
||||
#define LD_NC_FUNC "ld.global.nc.L1::no_allocate.L2::256B"
|
||||
#else
|
||||
#define LD_NC_FUNC "ld.volatile.global.L2::256B"
|
||||
#endif
|
||||
|
||||
namespace ltx_kernels {
|
||||
|
||||
// =============================================================================
|
||||
// Round-Robin SM Distribution Helpers
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* @brief Compute target rank for a given SM using round-robin distribution.
|
||||
*
|
||||
* Round-robin assignment ensures all SMs are utilized even when num_sms
|
||||
* is not evenly divisible by world_size.
|
||||
*
|
||||
* @param sm_id The SM/block ID (blockIdx.x)
|
||||
* @param world_size Total number of ranks
|
||||
* @return Target rank for this SM
|
||||
*/
|
||||
__device__ __forceinline__ int get_target_rank(int sm_id, int world_size) { return sm_id % world_size; }
|
||||
|
||||
/**
|
||||
* @brief Compute local SM index within a rank's SM group.
|
||||
*
|
||||
* With round-robin, SM i is the (i / world_size)-th SM assigned to its rank.
|
||||
*
|
||||
* @param sm_id The SM/block ID (blockIdx.x)
|
||||
* @param world_size Total number of ranks
|
||||
* @return Local index of this SM within its assigned rank's group
|
||||
*/
|
||||
__device__ __forceinline__ int get_rank_local_sm_id(int sm_id, int world_size) { return sm_id / world_size; }
|
||||
|
||||
/**
|
||||
* @brief Compute number of SMs assigned to a specific rank.
|
||||
*
|
||||
* With round-robin distribution:
|
||||
* - Ranks [0, extra) get (base + 1) SMs each
|
||||
* - Ranks [extra, world_size) get base SMs each
|
||||
* where base = num_sms / world_size, extra = num_sms % world_size
|
||||
*
|
||||
* @param target_rank The rank to query
|
||||
* @param num_sms Total number of SMs launched
|
||||
* @param world_size Total number of ranks
|
||||
* @return Number of SMs assigned to target_rank
|
||||
*/
|
||||
__device__ __forceinline__ int get_num_sms_for_rank(int target_rank, int num_sms, int world_size) {
|
||||
int base_sms = num_sms / world_size;
|
||||
int extra_sms = num_sms % world_size;
|
||||
return base_sms + (target_rank < extra_sms ? 1 : 0);
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Control Flow
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* @brief Triggers a GPU trap (fatal error).
|
||||
*
|
||||
* Used for unrecoverable errors like synchronization timeout.
|
||||
* Causes the kernel to abort and report an error to the host.
|
||||
*/
|
||||
__device__ __forceinline__ void trap() { asm("trap;"); }
|
||||
|
||||
// =============================================================================
|
||||
// Memory Ordering Operations (for synchronization)
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* @brief System-scope store with release ordering.
|
||||
*
|
||||
* Ensures all prior memory operations are visible before this store.
|
||||
* System scope means visibility across all GPUs (for IPC communication).
|
||||
*
|
||||
* @param ptr Pointer to global memory
|
||||
* @param val Value to store
|
||||
*/
|
||||
__device__ __forceinline__ void st_release_sys_global(const int *ptr, int val) {
|
||||
asm volatile("st.release.sys.global.s32 [%0], %1;" ::"l"(ptr), "r"(val) : "memory");
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief System-scope store with relaxed ordering.
|
||||
*
|
||||
* No ordering guarantees - fastest store but requires external synchronization.
|
||||
*
|
||||
* @param ptr Pointer to global memory
|
||||
* @param val Value to store
|
||||
*/
|
||||
__device__ __forceinline__ void st_relaxed_sys_global(const int *ptr, int val) {
|
||||
asm volatile("st.relaxed.sys.global.s32 [%0], %1;" ::"l"(ptr), "r"(val) : "memory");
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief CTA-scope store with release ordering.
|
||||
*
|
||||
* Ensures visibility within the thread block (CTA = Cooperative Thread Array).
|
||||
*
|
||||
* @param ptr Pointer to global memory
|
||||
* @param val Value to store
|
||||
*/
|
||||
__device__ __forceinline__ void st_release_cta(const int *ptr, int val) {
|
||||
asm volatile("st.release.cta.s32 [%0], %1;" ::"l"(ptr), "r"(val) : "memory");
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief System-scope load with acquire ordering (32-bit).
|
||||
*
|
||||
* Ensures subsequent memory operations are ordered after this load.
|
||||
* System scope for IPC visibility across GPUs.
|
||||
*
|
||||
* @param ptr Pointer to global memory
|
||||
* @return Loaded value
|
||||
*/
|
||||
__device__ __forceinline__ int ld_acquire_sys_global(const int *ptr) {
|
||||
int ret;
|
||||
asm volatile("ld.acquire.sys.global.s32 %0, [%1];" : "=r"(ret) : "l"(ptr));
|
||||
return ret;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief System-scope load with acquire ordering (64-bit).
|
||||
*
|
||||
* @param ptr Pointer to global memory
|
||||
* @return Loaded value
|
||||
*/
|
||||
__device__ __forceinline__ uint64_t ld_acquire_sys_global(const uint64_t *ptr) {
|
||||
uint64_t ret;
|
||||
asm volatile("ld.acquire.sys.global.u64 %0, [%1];" : "=l"(ret) : "l"(ptr));
|
||||
return ret;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief GPU-scope load with acquire ordering.
|
||||
*
|
||||
* Visibility limited to this GPU (not for IPC).
|
||||
*
|
||||
* @param ptr Pointer to global memory
|
||||
* @return Loaded value
|
||||
*/
|
||||
__device__ __forceinline__ int ld_acquire_global(const int *ptr) {
|
||||
int ret;
|
||||
asm volatile("ld.acquire.gpu.global.s32 %0, [%1];" : "=r"(ret) : "l"(ptr));
|
||||
return ret;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Volatile load bypassing all caches.
|
||||
*
|
||||
* Always reads from memory, never from cache. Used for polling
|
||||
* synchronization variables that may be updated by other GPUs.
|
||||
*
|
||||
* @param ptr Pointer to global memory
|
||||
* @return Loaded value
|
||||
*/
|
||||
__device__ __forceinline__ int ld_volatile_global(const int *ptr) {
|
||||
int ret;
|
||||
asm volatile("ld.volatile.global.s32 %0, [%1];" : "=r"(ret) : "l"(ptr));
|
||||
return ret;
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Optimized Bulk Memory Operations
|
||||
// =============================================================================
|
||||
|
||||
/**
|
||||
* @brief Non-allocating 128-bit store.
|
||||
*
|
||||
* Stores an int4 (128 bits / 16 bytes) without allocating in L1 cache.
|
||||
* Optimal for write-streaming patterns where data won't be read locally.
|
||||
*
|
||||
* @param ptr Destination pointer (must be 16-byte aligned)
|
||||
* @param value Data to store
|
||||
*/
|
||||
__device__ __forceinline__ void st_na_global(const int4 *ptr, const int4 &value) {
|
||||
asm volatile(ST_NA_FUNC ".v4.s32 [%0], {%1, %2, %3, %4};" ::"l"(ptr), "r"(value.x), "r"(value.y), "r"(value.z),
|
||||
"r"(value.w));
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Non-caching 128-bit load.
|
||||
*
|
||||
* Loads an int4 bypassing L1 cache with optimized L2 caching (256B lines).
|
||||
* Optimal for read-streaming patterns.
|
||||
*
|
||||
* @param ptr Source pointer (must be 16-byte aligned)
|
||||
* @return Loaded int4 value
|
||||
*/
|
||||
__device__ __forceinline__ int4 ld_nc_global(const int4 *ptr) {
|
||||
int4 ret;
|
||||
asm volatile(LD_NC_FUNC ".v4.s32 {%0, %1, %2, %3}, [%4];"
|
||||
: "=r"(ret.x), "=r"(ret.y), "=r"(ret.z), "=r"(ret.w)
|
||||
: "l"(ptr));
|
||||
return ret;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Barrier synchronization pattern for multi-GPU communication.
|
||||
*
|
||||
* This function implements a barrier synchronization protocol used in All2All
|
||||
* and AllGather operations. It signals completion to target ranks and waits
|
||||
* for all expected signals to arrive before resetting the barrier.
|
||||
*
|
||||
* Protocol:
|
||||
* 1. Thread 0 of each block signals completion to the target rank
|
||||
* 2. Block 0 waits for all ranks to signal (with timeout protection)
|
||||
* 3. Once all signals received, reset the barrier counters
|
||||
*
|
||||
* @param barrier_signal_ptrs Array of pointers to barrier signal buffers for each rank
|
||||
* @param target_rank The rank this block is sending data to
|
||||
* @param rank This GPU's rank
|
||||
* @param world_size Total number of GPUs/ranks
|
||||
* @param expected_count Number of signals expected (typically num_sms_per_rank)
|
||||
* @param sm_id The SM/block ID (blockIdx.x)
|
||||
* @param thread_id The thread ID within the block (threadIdx.x)
|
||||
* @param timeout_cycles Number of cycles to wait before timeout
|
||||
*/
|
||||
__device__ __forceinline__ void barrier_wait_and_reset(int **barrier_signal_ptrs, int target_rank, int rank,
|
||||
int world_size, int expected_count, int sm_id, int thread_id,
|
||||
uint64_t timeout_cycles) {
|
||||
// Release: fence so peers see our data writes, then sync before signaling.
|
||||
__threadfence_system();
|
||||
__syncthreads();
|
||||
|
||||
// Thread 0 signals completion to target rank
|
||||
if (thread_id == 0) {
|
||||
atomicAdd_system(barrier_signal_ptrs[target_rank] + rank, 1);
|
||||
}
|
||||
|
||||
// Synchronize before checking signals
|
||||
__syncthreads();
|
||||
|
||||
// Only block 0 waits for all signals and resets the barrier
|
||||
if (sm_id == 0 && thread_id < world_size) {
|
||||
auto start_time = clock64();
|
||||
while (true) {
|
||||
// Acquire: seeing the signal guarantees the peer's data is visible.
|
||||
int recv_count = ld_acquire_sys_global(barrier_signal_ptrs[rank] + thread_id);
|
||||
if (recv_count == expected_count) {
|
||||
break;
|
||||
}
|
||||
if (clock64() - start_time >= timeout_cycles) {
|
||||
printf("All2All barrier timeout: rank=%d, waiting_for_source=%d, expected=%d, got=%d\n", rank, thread_id,
|
||||
expected_count, recv_count);
|
||||
trap();
|
||||
}
|
||||
}
|
||||
// Reset barrier for next use
|
||||
atomicSub_system(barrier_signal_ptrs[rank] + thread_id, expected_count);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Barrier synchronization for round-robin SM distribution.
|
||||
*
|
||||
* Similar to barrier_wait_and_reset, but handles the case where SMs are
|
||||
* distributed round-robin across ranks, resulting in different target ranks
|
||||
* receiving different numbers of signals.
|
||||
*
|
||||
* With round-robin: target ranks [0, extra) receive (base + 1) signals from
|
||||
* each source, and target ranks [extra, world_size) receive base signals
|
||||
* from each source. Note that ALL sources send the same count to a given
|
||||
* receiver - the count depends on the receiver's rank position.
|
||||
*
|
||||
* @param barrier_signal_ptrs Array of pointers to barrier signal buffers for each rank
|
||||
* @param target_rank The rank this block is sending data to
|
||||
* @param rank This GPU's rank
|
||||
* @param world_size Total number of GPUs/ranks
|
||||
* @param num_sms Total number of SMs launched (used to compute expected counts)
|
||||
* @param sm_id The SM/block ID (blockIdx.x)
|
||||
* @param thread_id The thread ID within the block (threadIdx.x)
|
||||
* @param timeout_cycles Number of cycles to wait before timeout
|
||||
*/
|
||||
__device__ __forceinline__ void barrier_wait_and_reset_roundrobin(int **barrier_signal_ptrs, int target_rank, int rank,
|
||||
int world_size, int num_sms, int sm_id, int thread_id,
|
||||
uint64_t timeout_cycles) {
|
||||
// Release: fence so peers see our data writes, then sync before signaling.
|
||||
__threadfence_system();
|
||||
__syncthreads();
|
||||
|
||||
// Thread 0 signals completion to target rank
|
||||
if (thread_id == 0) {
|
||||
atomicAdd_system(barrier_signal_ptrs[target_rank] + rank, 1);
|
||||
}
|
||||
|
||||
// Synchronize before checking signals
|
||||
__syncthreads();
|
||||
|
||||
// Only block 0 waits for all signals and resets the barrier
|
||||
// Each thread handles one source rank
|
||||
if (sm_id == 0 && thread_id < world_size) {
|
||||
// All sources send the same number of signals to THIS receiver.
|
||||
// The count depends on how many SMs target this rank (the receiver).
|
||||
int expected_from_each_source = get_num_sms_for_rank(rank, num_sms, world_size);
|
||||
|
||||
auto start_time = clock64();
|
||||
while (true) {
|
||||
// Acquire: seeing the signal guarantees the peer's data is visible.
|
||||
int recv_count = ld_acquire_sys_global(barrier_signal_ptrs[rank] + thread_id);
|
||||
if (recv_count == expected_from_each_source) {
|
||||
break;
|
||||
}
|
||||
if (clock64() - start_time >= timeout_cycles) {
|
||||
printf("All2All barrier timeout (roundrobin): rank=%d, waiting_for_source=%d, expected=%d, got=%d\n", rank,
|
||||
thread_id, expected_from_each_source, recv_count);
|
||||
trap();
|
||||
}
|
||||
}
|
||||
// Reset barrier for next use
|
||||
atomicSub_system(barrier_signal_ptrs[rank] + thread_id, expected_from_each_source);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace ltx_kernels
|
||||
@@ -0,0 +1,114 @@
|
||||
/**
|
||||
* @file event.hpp
|
||||
* @brief CUDA stream and event synchronization utilities.
|
||||
*
|
||||
* This header provides wrapper types and helper functions for managing
|
||||
* CUDA events and stream synchronization in PyTorch/ATen environment.
|
||||
* These utilities are used to coordinate asynchronous operations across
|
||||
* multiple CUDA streams.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <memory>
|
||||
|
||||
#include "cuda/exceptions.cuh"
|
||||
|
||||
namespace ltx_kernels {
|
||||
|
||||
/**
|
||||
* @struct EventHandle
|
||||
* @brief RAII wrapper for a CUDA event with automatic recording.
|
||||
*
|
||||
* EventHandle encapsulates a torch::Event and automatically records it
|
||||
* on the specified (or current) CUDA stream upon construction. This
|
||||
* provides a convenient way to capture the completion point of stream
|
||||
* operations for synchronization purposes.
|
||||
*
|
||||
* ## Usage Example
|
||||
*
|
||||
* ```cpp
|
||||
* // Record event on current stream
|
||||
* EventHandle ev1;
|
||||
*
|
||||
* // Record event on specific stream
|
||||
* EventHandle ev2(my_stream);
|
||||
*
|
||||
* // Make current stream wait for the event
|
||||
* ev1.current_stream_wait();
|
||||
* ```
|
||||
*/
|
||||
struct EventHandle {
|
||||
/// Shared pointer to the underlying torch::Event
|
||||
std::shared_ptr<torch::Event> event;
|
||||
|
||||
/**
|
||||
* @brief Constructs an EventHandle and records on the current CUDA stream.
|
||||
*
|
||||
* The event captures the completion point of all operations submitted
|
||||
* to the current stream before this constructor is called.
|
||||
*/
|
||||
EventHandle() {
|
||||
event = std::make_shared<torch::Event>(torch::kCUDA);
|
||||
event->record(at::cuda::getCurrentCUDAStream());
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Constructs an EventHandle and records on the specified stream.
|
||||
*
|
||||
* @param stream The CUDA stream to record the event on
|
||||
*/
|
||||
explicit EventHandle(const at::cuda::CUDAStream &stream) {
|
||||
event = std::make_shared<torch::Event>(torch::kCUDA);
|
||||
event->record(stream);
|
||||
}
|
||||
|
||||
/// Copy constructor (shares the underlying event)
|
||||
EventHandle(const EventHandle &other) = default;
|
||||
|
||||
/**
|
||||
* @brief Makes the current CUDA stream wait for this event.
|
||||
*
|
||||
* After this call returns, operations submitted to the current stream
|
||||
* will not execute until the event has been reached on its recording stream.
|
||||
*/
|
||||
void current_stream_wait() const { at::cuda::getCurrentCUDAStream().unwrap().wait(*event); }
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief Creates and records a CUDA event on the specified stream.
|
||||
*
|
||||
* @param s The CUDA stream to record on
|
||||
* @return A torch::Event that has been recorded on stream s
|
||||
*/
|
||||
inline torch::Event create_event(const at::cuda::CUDAStream &s) {
|
||||
auto event = torch::Event(torch::kCUDA);
|
||||
event.record(s);
|
||||
return event;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Makes stream s_0 wait for stream s_1's current position.
|
||||
*
|
||||
* After this call, operations on s_0 will not execute until all operations
|
||||
* currently queued on s_1 have completed.
|
||||
*
|
||||
* @param s_0 The stream that will wait
|
||||
* @param s_1 The stream to wait for
|
||||
* @pre s_0 and s_1 must be different streams
|
||||
*/
|
||||
inline void stream_wait(const at::cuda::CUDAStream &s_0, const at::cuda::CUDAStream &s_1) {
|
||||
EP_HOST_ASSERT(s_0.id() != s_1.id());
|
||||
s_0.unwrap().wait(create_event(s_1));
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Makes a stream wait for a previously recorded event.
|
||||
*
|
||||
* @param s The stream that will wait
|
||||
* @param event The event to wait for
|
||||
*/
|
||||
inline void stream_wait(const at::cuda::CUDAStream &s, const EventHandle &event) { s.unwrap().wait(*event.event); }
|
||||
|
||||
} // namespace ltx_kernels
|
||||
@@ -0,0 +1,73 @@
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <torch/extension.h>
|
||||
#include <torch/python.h>
|
||||
|
||||
#include <vector>
|
||||
|
||||
void fp6_pack_cuda(
|
||||
at::Tensor& x,
|
||||
at::Tensor& out,
|
||||
cudaStream_t stream
|
||||
);
|
||||
|
||||
void fp6_unpack_cuda(
|
||||
at::Tensor& x,
|
||||
at::Tensor& out,
|
||||
cudaStream_t stream
|
||||
);
|
||||
|
||||
at::Tensor fp6_pack(at::Tensor &x) {
|
||||
// TORCH_CHECK(x.dtype() == torch::kUInt8, "Input tensor must be uint8");
|
||||
TORCH_CHECK(x.is_cuda(), "Input tensor must be on CUDA");
|
||||
TORCH_CHECK(x.is_contiguous(), "Input tensor must be contiguous");
|
||||
TORCH_CHECK(x.dim() == 2, "Input tensor must be 2D [m, n]");
|
||||
|
||||
int64_t m = x.size(0);
|
||||
int64_t n = x.size(1);
|
||||
|
||||
TORCH_CHECK(n % 8 == 0, "n must be divisible by 8, got ", n);
|
||||
|
||||
// Output shape: [m, n*3/4] since 4 elements of 8-bit = 32 bits, 4 elements of 6-bit = 24 bits = 3 bytes
|
||||
int64_t n_packed = n * 3 / 4;
|
||||
|
||||
auto options = torch::TensorOptions()
|
||||
.dtype(torch::kUInt8)
|
||||
.device(x.device());
|
||||
|
||||
at::Tensor out = torch::empty({m, n_packed}, options);
|
||||
|
||||
at::cuda::CUDAGuard device_guard{x.get_device()};
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
fp6_pack_cuda(x, out, stream);
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
at::Tensor fp6_unpack(at::Tensor &x, int64_t original_n) {
|
||||
TORCH_CHECK(x.dtype() == torch::kUInt8, "Input tensor must be uint8");
|
||||
TORCH_CHECK(x.is_cuda(), "Input tensor must be on CUDA");
|
||||
TORCH_CHECK(x.is_contiguous(), "Input tensor must be contiguous");
|
||||
TORCH_CHECK(x.dim() == 2, "Input tensor must be 2D [m, n_packed]");
|
||||
TORCH_CHECK(original_n % 8 == 0, "original_n must be divisible by 8, got ", original_n);
|
||||
|
||||
int64_t m = x.size(0);
|
||||
int64_t n_packed = x.size(1);
|
||||
|
||||
TORCH_CHECK(n_packed == original_n * 3 / 4,
|
||||
"Packed size mismatch: expected ", original_n * 3 / 4, " got ", n_packed);
|
||||
|
||||
auto options = torch::TensorOptions()
|
||||
.dtype(torch::kUInt8)
|
||||
.device(x.device());
|
||||
|
||||
at::Tensor out = torch::empty({m, original_n}, options);
|
||||
|
||||
at::cuda::CUDAGuard device_guard{x.get_device()};
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
fp6_unpack_cuda(x, out, stream);
|
||||
|
||||
return out;
|
||||
}
|
||||
@@ -0,0 +1,211 @@
|
||||
#include <c10/cuda/CUDAException.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <cuda.h>
|
||||
|
||||
#include <ATen/ATen.h>
|
||||
#include <torch/types.h>
|
||||
|
||||
// Device function to pack 8-bit to 6-bit
|
||||
// 8-bit layout: s e_1 e_2 e_3 m_1 m_2 m_3 m_4 (bits 7-0)
|
||||
// 6-bit layout: s e_3 m_1 m_2 m_3 m_4 (bits 5-0)
|
||||
// Drop e_1 (bit 6) and e_2 (bit 5)
|
||||
__device__ __forceinline__ uint8_t pack_8bit_to_6bit(uint8_t input) {
|
||||
// Extract the sign bit (bit 7)
|
||||
uint8_t sign = (input >> 7) & 0x1;
|
||||
|
||||
// Extract e_3 (bit 4)
|
||||
uint8_t e_3 = (input >> 4) & 0x1;
|
||||
|
||||
// Extract mantissa bits (bits 3-0)
|
||||
uint8_t mantissa = input & 0x0F;
|
||||
|
||||
// Pack into 6-bit format: s e_3 m_1 m_2 m_3 m_4
|
||||
uint8_t result = (sign << 5) | (e_3 << 4) | mantissa;
|
||||
|
||||
return result & 0x3F; // Mask to 6 bits
|
||||
}
|
||||
|
||||
// Device function to pack 4 x 6-bit values into 3 bytes
|
||||
__device__ __forceinline__ void pack_4x6bit_to_3bytes(const uint8_t* input_6bit, uint8_t* output_3bytes) {
|
||||
uint8_t v0 = input_6bit[0] & 0x3F;
|
||||
uint8_t v1 = input_6bit[1] & 0x3F;
|
||||
uint8_t v2 = input_6bit[2] & 0x3F;
|
||||
uint8_t v3 = input_6bit[3] & 0x3F;
|
||||
|
||||
// Pack: [v0: 6 bits][v1: 6 bits][v2: 6 bits][v3: 6 bits] = 24 bits = 3 bytes
|
||||
output_3bytes[0] = (v0 << 2) | (v1 >> 4);
|
||||
output_3bytes[1] = (v1 << 4) | (v2 >> 2);
|
||||
output_3bytes[2] = (v2 << 6) | v3;
|
||||
}
|
||||
|
||||
// CUDA kernel for packing 2D tensor
|
||||
// Input: [m, n] uint8 tensor
|
||||
// Output: [m, n*3/4] uint8 tensor
|
||||
__global__ void fp6_pack_kernel(
|
||||
const uint8_t* __restrict__ input,
|
||||
uint8_t* __restrict__ output,
|
||||
int m,
|
||||
int n,
|
||||
int n_packed
|
||||
) {
|
||||
// Each thread processes one row and 4 elements at a time
|
||||
int row = blockIdx.x;
|
||||
int col_group = blockIdx.y * blockDim.x + threadIdx.x;
|
||||
|
||||
if (row >= m) return;
|
||||
|
||||
// Calculate input and output positions
|
||||
int input_col = col_group * 4;
|
||||
if (input_col >= n) return;
|
||||
|
||||
int output_col = col_group * 3;
|
||||
|
||||
const uint8_t* input_row = input + row * n;
|
||||
uint8_t* output_row = output + row * n_packed;
|
||||
|
||||
uint8_t temp_6bit[4];
|
||||
|
||||
// Pack 4 elements
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; i++) {
|
||||
if (input_col + i < n) {
|
||||
temp_6bit[i] = pack_8bit_to_6bit(input_row[input_col + i]);
|
||||
} else {
|
||||
temp_6bit[i] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
// Write 3 bytes to output
|
||||
uint8_t temp_3bytes[3];
|
||||
pack_4x6bit_to_3bytes(temp_6bit, temp_3bytes);
|
||||
|
||||
if (output_col < n_packed) output_row[output_col] = temp_3bytes[0];
|
||||
if (output_col + 1 < n_packed) output_row[output_col + 1] = temp_3bytes[1];
|
||||
if (output_col + 2 < n_packed) output_row[output_col + 2] = temp_3bytes[2];
|
||||
}
|
||||
|
||||
// Device function to unpack 6-bit to 8-bit
|
||||
__device__ __forceinline__ uint8_t unpack_6bit_to_8bit(uint8_t input) {
|
||||
input = input & 0x3F; // Ensure only 6 bits
|
||||
|
||||
uint8_t sign = (input >> 5) & 0x1;
|
||||
uint8_t e_3 = (input >> 4) & 0x1;
|
||||
uint8_t mantissa = input & 0x0F;
|
||||
|
||||
// Reconstruct 8-bit with e_1 and e_2 set to 0
|
||||
uint8_t result = (sign << 7) | (e_3 << 4) | mantissa;
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
// Device function to unpack 3 bytes into 4 x 6-bit values
|
||||
__device__ __forceinline__ void unpack_3bytes_to_4x6bit(const uint8_t* input_3bytes, uint8_t* output_6bit) {
|
||||
output_6bit[0] = (input_3bytes[0] >> 2) & 0x3F;
|
||||
output_6bit[1] = ((input_3bytes[0] << 4) | (input_3bytes[1] >> 4)) & 0x3F;
|
||||
output_6bit[2] = ((input_3bytes[1] << 2) | (input_3bytes[2] >> 6)) & 0x3F;
|
||||
output_6bit[3] = input_3bytes[2] & 0x3F;
|
||||
}
|
||||
|
||||
// CUDA kernel for unpacking 2D tensor
|
||||
// Input: [m, n_packed] uint8 tensor
|
||||
// Output: [m, n] uint8 tensor
|
||||
__global__ void fp6_unpack_kernel(
|
||||
const uint8_t* __restrict__ input,
|
||||
uint8_t* __restrict__ output,
|
||||
int m,
|
||||
int n_packed,
|
||||
int n
|
||||
) {
|
||||
// Each thread processes one row and 4 elements at a time
|
||||
int row = blockIdx.x;
|
||||
int col_group = blockIdx.y * blockDim.x + threadIdx.x;
|
||||
|
||||
if (row >= m) return;
|
||||
|
||||
// Calculate input and output positions
|
||||
int input_col = col_group * 3;
|
||||
if (input_col >= n_packed) return;
|
||||
|
||||
int output_col = col_group * 4;
|
||||
|
||||
const uint8_t* input_row = input + row * n_packed;
|
||||
uint8_t* output_row = output + row * n;
|
||||
|
||||
// Read 3 bytes
|
||||
uint8_t temp_3bytes[3];
|
||||
temp_3bytes[0] = (input_col < n_packed) ? input_row[input_col] : 0;
|
||||
temp_3bytes[1] = (input_col + 1 < n_packed) ? input_row[input_col + 1] : 0;
|
||||
temp_3bytes[2] = (input_col + 2 < n_packed) ? input_row[input_col + 2] : 0;
|
||||
|
||||
// Unpack to 4 x 6-bit values
|
||||
uint8_t temp_6bit[4];
|
||||
unpack_3bytes_to_4x6bit(temp_3bytes, temp_6bit);
|
||||
|
||||
// Convert to 8-bit and write
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; i++) {
|
||||
if (output_col + i < n) {
|
||||
output_row[output_col + i] = unpack_6bit_to_8bit(temp_6bit[i]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Host function to launch pack kernel
|
||||
void fp6_pack_cuda(
|
||||
at::Tensor& x,
|
||||
at::Tensor& out,
|
||||
cudaStream_t stream
|
||||
) {
|
||||
int m = x.size(0);
|
||||
int n = x.size(1);
|
||||
int n_packed = out.size(1);
|
||||
|
||||
const uint8_t* input_ptr = (uint8_t*)x.data_ptr();
|
||||
uint8_t* output_ptr = (uint8_t*)out.data_ptr();
|
||||
|
||||
// Each thread handles 4 input elements -> 3 output bytes
|
||||
int num_groups = (n + 3) / 4;
|
||||
|
||||
int threads = 256;
|
||||
dim3 blocks(m, (num_groups + threads - 1) / threads);
|
||||
|
||||
fp6_pack_kernel<<<blocks, threads, 0, stream>>>(
|
||||
input_ptr,
|
||||
output_ptr,
|
||||
m,
|
||||
n,
|
||||
n_packed
|
||||
);
|
||||
|
||||
C10_CUDA_KERNEL_LAUNCH_CHECK();
|
||||
}
|
||||
|
||||
// Host function to launch unpack kernel
|
||||
void fp6_unpack_cuda(
|
||||
at::Tensor& x,
|
||||
at::Tensor& out,
|
||||
cudaStream_t stream
|
||||
) {
|
||||
int m = x.size(0);
|
||||
int n_packed = x.size(1);
|
||||
int n = out.size(1);
|
||||
|
||||
const uint8_t* input_ptr = (uint8_t*)x.data_ptr();
|
||||
uint8_t* output_ptr = (uint8_t*)out.data_ptr();
|
||||
|
||||
// Each thread handles 3 input bytes -> 4 output elements
|
||||
int num_groups = (n + 3) / 4;
|
||||
|
||||
int threads = 256;
|
||||
dim3 blocks(m, (num_groups + threads - 1) / threads);
|
||||
|
||||
fp6_unpack_kernel<<<blocks, threads, 0, stream>>>(
|
||||
input_ptr,
|
||||
output_ptr,
|
||||
m,
|
||||
n_packed,
|
||||
n
|
||||
);
|
||||
|
||||
C10_CUDA_KERNEL_LAUNCH_CHECK();
|
||||
}
|
||||
@@ -0,0 +1,163 @@
|
||||
/******************************************************************************
|
||||
* Copyright (c) 2023, Tri Dao.
|
||||
******************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct HadamardParamsBase {
|
||||
using index_t = int64_t;
|
||||
|
||||
int batch, dim, log_N;
|
||||
|
||||
index_t x_batch_stride;
|
||||
index_t out_batch_stride;
|
||||
|
||||
float scale;
|
||||
|
||||
// Common data pointers.
|
||||
void *__restrict__ x_ptr;
|
||||
void *__restrict__ out_ptr;
|
||||
};
|
||||
|
||||
struct UnifiedHadamardParamsBase{
|
||||
using index_t = int64_t;
|
||||
|
||||
int batch, dim, log_N;
|
||||
int batch_fma_change;
|
||||
|
||||
index_t x_batch_stride;
|
||||
index_t out_batch_stride;
|
||||
index_t fma_batch_stride;
|
||||
index_t cos_freq_batch_stride;
|
||||
index_t sin_freq_batch_stride;
|
||||
|
||||
float scale;
|
||||
|
||||
// Common data pointers.
|
||||
void *__restrict__ x_ptr;
|
||||
void *__restrict__ out_ptr;
|
||||
void *__restrict__ out_scales_ptr;
|
||||
|
||||
void *__restrict__ y_scale_ptr;
|
||||
void *__restrict__ z_shift_ptr;
|
||||
|
||||
void *__restrict__ weights_ptr;
|
||||
|
||||
void *__restrict__ cos_freq_ptr;
|
||||
void *__restrict__ sin_freq_ptr;
|
||||
|
||||
};
|
||||
|
||||
|
||||
struct DequantHadamardParamsBase {
|
||||
using index_t = int64_t;
|
||||
|
||||
int batch, dim, log_N;
|
||||
|
||||
index_t x_batch_stride;
|
||||
index_t out_batch_stride;
|
||||
|
||||
float scale;
|
||||
|
||||
// Common data pointers.
|
||||
void *__restrict__ x_ptr;
|
||||
void *__restrict__ scales_ptr;
|
||||
void *__restrict__ out_ptr;
|
||||
};
|
||||
|
||||
|
||||
struct QuantHadamardParamsBase {
|
||||
using index_t = int64_t;
|
||||
|
||||
int batch, dim, log_N;
|
||||
|
||||
index_t x_batch_stride;
|
||||
index_t out_batch_stride;
|
||||
|
||||
float scale;
|
||||
|
||||
// Common data pointers.
|
||||
void *__restrict__ x_ptr;
|
||||
void *__restrict__ out_ptr;
|
||||
void *__restrict__ out_scales_ptr;
|
||||
};
|
||||
|
||||
struct NormFMAHadamardParamsBase {
|
||||
using index_t = int64_t;
|
||||
|
||||
int batch, dim, log_N;
|
||||
int seqlen;
|
||||
|
||||
index_t x_batch_stride;
|
||||
index_t out_batch_stride;
|
||||
index_t fma_batch_stride;
|
||||
|
||||
float scale;
|
||||
|
||||
// Common data pointers.
|
||||
void *__restrict__ x_ptr;
|
||||
void *__restrict__ out_ptr;
|
||||
void *__restrict__ y_scale_ptr;
|
||||
void *__restrict__ z_shift_ptr;
|
||||
void *__restrict__ weights_ptr;
|
||||
};
|
||||
|
||||
|
||||
struct NormRopeHadamardParamsBase {
|
||||
using index_t = int64_t;
|
||||
|
||||
int batch, dim, log_N;
|
||||
|
||||
index_t x_batch_stride;
|
||||
index_t out_batch_stride;
|
||||
index_t cos_freq_batch_stride;
|
||||
index_t sin_freq_batch_stride;
|
||||
|
||||
float scale;
|
||||
|
||||
// Common data pointers.
|
||||
void *__restrict__ x_ptr;
|
||||
void *__restrict__ out_ptr;
|
||||
void *__restrict__ cos_freq_ptr;
|
||||
void *__restrict__ sin_freq_ptr;
|
||||
void *__restrict__ weights_ptr;
|
||||
};
|
||||
|
||||
|
||||
struct NormHadamardParamsBase {
|
||||
using index_t = int64_t;
|
||||
|
||||
int batch, dim, log_N;
|
||||
|
||||
index_t x_batch_stride;
|
||||
index_t out_batch_stride;
|
||||
|
||||
float scale;
|
||||
|
||||
// Common data pointers.
|
||||
void *__restrict__ x_ptr;
|
||||
void *__restrict__ out_ptr;
|
||||
void *__restrict__ weights_ptr;
|
||||
};
|
||||
|
||||
|
||||
struct RopeHadamardParamsBase {
|
||||
using index_t = int64_t;
|
||||
|
||||
int batch, dim, log_N;
|
||||
|
||||
index_t x_batch_stride;
|
||||
index_t out_batch_stride;
|
||||
index_t cos_freq_batch_stride;
|
||||
index_t sin_freq_batch_stride;
|
||||
|
||||
float scale;
|
||||
|
||||
// Common data pointers.
|
||||
void *__restrict__ x_ptr;
|
||||
void *__restrict__ out_ptr;
|
||||
void *__restrict__ cos_freq_ptr;
|
||||
void *__restrict__ sin_freq_ptr;
|
||||
};
|
||||
@@ -0,0 +1,319 @@
|
||||
/******************************************************************************
|
||||
* Copyright (c) 2023, Tri Dao.
|
||||
******************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp16.h>
|
||||
|
||||
#define FULL_MASK 0xffffffff
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template<typename TYPE> struct QuantMax {};
|
||||
template<> struct QuantMax<int8_t> { static constexpr float value = 127.0; };
|
||||
template<> struct QuantMax<at::Float8_e4m3fn> { static constexpr float value = 256.0; };
|
||||
|
||||
struct uint8 {
|
||||
uint4 u;
|
||||
uint4 v;
|
||||
};
|
||||
|
||||
template<int BYTES> struct BytesToType {};
|
||||
|
||||
template<>
|
||||
struct BytesToType<32> {
|
||||
using Type = uint8;
|
||||
static_assert(sizeof(Type) == 32);
|
||||
};
|
||||
|
||||
template<> struct BytesToType<16> {
|
||||
using Type = uint4;
|
||||
static_assert(sizeof(Type) == 16);
|
||||
};
|
||||
|
||||
template<> struct BytesToType<8> {
|
||||
using Type = uint64_t;
|
||||
static_assert(sizeof(Type) == 8);
|
||||
};
|
||||
|
||||
template<> struct BytesToType<4> {
|
||||
using Type = uint32_t;
|
||||
static_assert(sizeof(Type) == 4);
|
||||
};
|
||||
|
||||
template<> struct BytesToType<2> {
|
||||
using Type = uint16_t;
|
||||
static_assert(sizeof(Type) == 2);
|
||||
};
|
||||
|
||||
template<> struct BytesToType<1> {
|
||||
using Type = uint8_t;
|
||||
static_assert(sizeof(Type) == 1);
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<typename T>
|
||||
struct SumOp {
|
||||
__device__ inline T operator()(T const & x, T const & y) { return x + y; }
|
||||
};
|
||||
|
||||
template<typename T>
|
||||
struct MaxOp {
|
||||
__device__ inline T operator()(T const & x, T const & y) { return max(x, y); }
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MaxOp<float> {
|
||||
// This is slightly faster
|
||||
__device__ inline float operator()(float const &x, float const &y) { return max(x, y); }
|
||||
};
|
||||
|
||||
|
||||
template<int THREADS>
|
||||
struct Allreduce {
|
||||
static_assert(THREADS == 32 || THREADS == 16 || THREADS == 8 || THREADS == 4);
|
||||
template<typename T, typename Operator>
|
||||
static __device__ inline T run(T x, Operator &op) {
|
||||
constexpr int OFFSET = THREADS / 2;
|
||||
x = op(x, __shfl_xor_sync(uint32_t(-1), x, OFFSET));
|
||||
return Allreduce<OFFSET>::run(x, op);
|
||||
}
|
||||
};
|
||||
|
||||
template<>
|
||||
struct Allreduce<2> {
|
||||
template<typename T, typename Operator>
|
||||
static __device__ inline T run(T x, Operator &op) {
|
||||
x = op(x, __shfl_xor_sync(uint32_t(-1), x, 1));
|
||||
return x;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// https://stackoverflow.com/questions/35311711/whats-the-right-way-to-compute-integral-base-2-logarithms-at-compile-time
|
||||
constexpr int cilog2(int val) { return val > 0 ? 1 + cilog2(val >> 1) : -1; }
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<int kLogN, int kNChunks>
|
||||
__device__ __forceinline__ void hadamard_mult_thread(float x[kNChunks][1 << kLogN]) {
|
||||
constexpr int N = 1 << kLogN;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kLogN; ++i) {
|
||||
const int stride = 1 << i;
|
||||
#pragma unroll
|
||||
for (int j = 0; j < N / 2; ++j) {
|
||||
const int lo = j & (stride - 1);
|
||||
const int idx = (j - lo) * 2 + lo;
|
||||
#pragma unroll
|
||||
for (int c = 0; c < kNChunks; ++c) {
|
||||
const float a = x[c][idx];
|
||||
const float b = x[c][idx + stride];
|
||||
x[c][idx] = a + b;
|
||||
x[c][idx + stride] = a - b;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<int kLogWarpSize, int kStepStart, int kNChunks, int kNItems>
|
||||
__device__ __forceinline__ void hadamard_mult_warp(float x[kNChunks][kNItems]) {
|
||||
constexpr int N = 1 << kLogWarpSize;
|
||||
int lane_id = threadIdx.x % N;
|
||||
#pragma unroll
|
||||
for (int step = kStepStart; step < kLogWarpSize; ++step) {
|
||||
const int lane_mask = 1 << step;
|
||||
const float sign = (lane_id & lane_mask) ? -1.f : 1.f;
|
||||
#pragma unroll
|
||||
for (int c = 0; c < kNChunks; ++c) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kNItems; ++i) {
|
||||
float x_val_other = __shfl_xor_sync(FULL_MASK, x[c][i], lane_mask);
|
||||
x[c][i] = sign * x[c][i] + x_val_other;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <int kNChunks, int kNElts, typename input_t>
|
||||
inline __device__ void load_input(input_t *x, float x_vals[kNChunks][kNElts], int dim) {
|
||||
using vec_t = typename BytesToType<sizeof(input_t) * kNElts>::Type;
|
||||
input_t x_vals_load[kNChunks][kNElts] = {0};
|
||||
#pragma unroll
|
||||
for (int c = 0; c < kNChunks; ++c) {
|
||||
if ((c * blockDim.x + threadIdx.x) * kNElts < dim) {
|
||||
reinterpret_cast<vec_t*>(x_vals_load)[c] = reinterpret_cast<const vec_t*>(x)[c * blockDim.x + threadIdx.x];
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int c = 0; c < kNChunks; ++c) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kNElts; ++i) { x_vals[c][i] = float(x_vals_load[c][i]); }
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
template <int kNChunks, int kNElts, typename output_t, bool do_round>
|
||||
inline __device__ void store_output(output_t *out, float out_vals[kNChunks][kNElts], int dim, float scale=1.f) {
|
||||
using vec_t = typename BytesToType<sizeof(output_t) * kNElts>::Type;
|
||||
output_t out_vals_store[kNChunks][kNElts];
|
||||
#pragma unroll
|
||||
for (int c = 0; c < kNChunks; ++c) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < kNElts; ++i) {
|
||||
if constexpr (do_round){
|
||||
out_vals_store[c][i] = round(out_vals[c][i] * scale);
|
||||
} else {
|
||||
out_vals_store[c][i] = out_vals[c][i] * scale;
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int c = 0; c < kNChunks; ++c) {
|
||||
if ((c * blockDim.x + threadIdx.x) * kNElts < dim) {
|
||||
reinterpret_cast<vec_t*>(out)[c * blockDim.x + threadIdx.x] = reinterpret_cast<const vec_t*>(out_vals_store)[c];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Pre=true means the exchange before the hadamard_mult_warp, Pre=false means after.
|
||||
template <int kNChunks, int kChunksPerExchange, int kNElts, int kWarpSize, int kNWarps, bool Pre, typename vec_t>
|
||||
inline __device__ void exchange_smem_pre(float x_vals[kNChunks][kNElts], vec_t *smem) {
|
||||
constexpr int kNThreads = kWarpSize * kNWarps;
|
||||
constexpr int kNExchangePerVec = kNElts / (sizeof(vec_t) / sizeof(float));
|
||||
const int warp_id = threadIdx.x / kWarpSize;
|
||||
const int lane_id = threadIdx.x % kWarpSize;
|
||||
const int row_t = threadIdx.x % kNWarps;
|
||||
const int col_t = threadIdx.x / kNWarps;
|
||||
// We use the XOR swizzle trick (new_col = col ^ row) to avoid / reduce smem bank conflicts.
|
||||
#pragma unroll
|
||||
for (int c0 = 0; c0 < kNChunks / kChunksPerExchange; ++c0) {
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int c1 = 0; c1 < kChunksPerExchange; ++c1) {
|
||||
#pragma unroll
|
||||
for (int r = 0; r < kNExchangePerVec; ++r) {
|
||||
smem[(c1 * kNExchangePerVec + r) * kNThreads + (Pre ? warp_id * kWarpSize + lane_id ^ warp_id : row_t * kWarpSize + col_t ^ row_t)] = reinterpret_cast<vec_t*>(x_vals[c0 * kChunksPerExchange + c1])[r];
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int c1 = 0; c1 < kChunksPerExchange; ++c1) {
|
||||
#pragma unroll
|
||||
for (int r = 0; r < kNExchangePerVec; ++r) {
|
||||
reinterpret_cast<vec_t*>(x_vals[c0 * kChunksPerExchange + c1])[r] = smem[(c1 * kNExchangePerVec + r) * kNThreads + (Pre ? row_t * kWarpSize + col_t ^ row_t : warp_id * kWarpSize + lane_id ^ warp_id)];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
inline __device__ float gelu_approximate(float x){
|
||||
constexpr float sqrthalfpi2 = 0.7978845608028653558798921198687637369517172623298693153318516593f;
|
||||
constexpr float factor = 0.044715f;
|
||||
return 0.5f*x*(1.0f + tanhf(sqrthalfpi2*(x + factor*x*x*x)));
|
||||
}
|
||||
|
||||
template <int kNChunks, int kNElts>
|
||||
inline __device__ void fused_gelu(float x_vals[kNChunks][kNElts]){
|
||||
#pragma unroll
|
||||
for (size_t c = 0; c < kNChunks; c++)
|
||||
{
|
||||
#pragma unroll
|
||||
for (size_t i = 0; i < kNElts; i++)
|
||||
{
|
||||
x_vals[c][i] = gelu_approximate(x_vals[c][i]);
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
template <int kNChunks, int kNElts, int kNWarps, bool norm_affine>
|
||||
inline __device__ void fused_rms_norm(float x_vals[kNChunks][kNElts], float weights_vals[kNChunks][kNElts], float* smem_sum, float dim){
|
||||
float thread_squared_sum = 0.0f;
|
||||
const int warp_id = threadIdx.x / 32;
|
||||
|
||||
#pragma unroll
|
||||
for (size_t c = 0; c < kNChunks; c++)
|
||||
{
|
||||
#pragma unroll
|
||||
for (size_t i = 0; i < kNElts; i++)
|
||||
{
|
||||
thread_squared_sum += x_vals[c][i] * x_vals[c][i];
|
||||
}
|
||||
|
||||
}
|
||||
SumOp<float> sum_op;
|
||||
float warp_sum = Allreduce<32>::run(thread_squared_sum, sum_op);
|
||||
|
||||
if(threadIdx.x % 32 == 0){
|
||||
smem_sum[warp_id] = warp_sum;
|
||||
}
|
||||
__syncthreads();
|
||||
float norm = 0.0f;
|
||||
#pragma unroll
|
||||
for (size_t i = 0; i < kNWarps; i++)
|
||||
{
|
||||
norm += smem_sum[i];
|
||||
}
|
||||
|
||||
norm *= 1.0f/dim;
|
||||
norm = rsqrtf(norm + 0.0000001f);
|
||||
|
||||
#pragma unroll
|
||||
for (size_t c = 0; c < kNChunks; c++)
|
||||
{
|
||||
#pragma unroll
|
||||
for (size_t i = 0; i < kNElts; i++)
|
||||
{
|
||||
if constexpr (norm_affine){
|
||||
x_vals[c][i] *= (norm * weights_vals[c][i]);
|
||||
} else {
|
||||
x_vals[c][i] *= norm;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
template <int kNChunks, int kNElts>
|
||||
inline __device__ void fused_rope(float x_vals[kNChunks][kNElts], float sin_freqs_vals[kNChunks][kNElts], float cos_freqs_vals[kNChunks][kNElts]){
|
||||
#pragma unroll
|
||||
for (size_t c = 0; c < kNChunks; c++)
|
||||
{
|
||||
#pragma unroll
|
||||
for (size_t i = 0; i < kNElts; i+=2)
|
||||
{
|
||||
float x_1 = x_vals[c][i];
|
||||
float x_2 = x_vals[c][i+1];
|
||||
x_vals[c][i] = -x_2*sin_freqs_vals[c][i] + x_1*cos_freqs_vals[c][i];
|
||||
x_vals[c][i+1] = x_1*sin_freqs_vals[c][i+1] + x_2*cos_freqs_vals[c][i+1];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
template <int kNChunks, int kNElts, bool add_one_scale>
|
||||
inline __device__ void fused_multiply_add(float x_vals[kNChunks][kNElts], float y_scale_vals[kNChunks][kNElts], float z_shift_vals[kNChunks][kNElts]) {
|
||||
|
||||
#pragma unroll
|
||||
for (size_t c = 0; c < kNChunks; c++)
|
||||
{
|
||||
#pragma unroll
|
||||
for (size_t i = 0; i < kNElts; i++)
|
||||
{
|
||||
if constexpr (add_one_scale){
|
||||
x_vals[c][i] = x_vals[c][i] * (1.0f + y_scale_vals[c][i]) + z_shift_vals[c][i];
|
||||
} else {
|
||||
x_vals[c][i] = x_vals[c][i] * y_scale_vals[c][i] + z_shift_vals[c][i];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
|
||||
/******************************************************************************
|
||||
* Copyright (c) 2023, Tri Dao.
|
||||
******************************************************************************/
|
||||
|
||||
// This file is auto-generated. See "code_gen.py"
|
||||
|
||||
|
||||
#pragma once
|
||||
|
||||
|
||||
__device__ __forceinline__ void hadamard_mult_thread_12(float x[12]) {
|
||||
float out[12];
|
||||
out[0] = + x[0] - x[1] + x[2] + x[3] + x[4] + x[5] + x[6] + x[7] + x[8] + x[9] + x[10] + x[11];
|
||||
out[1] = - x[0] - x[1] + x[2] - x[3] + x[4] - x[5] + x[6] - x[7] + x[8] - x[9] + x[10] - x[11];
|
||||
out[2] = + x[0] + x[1] + x[2] - x[3] + x[4] + x[5] - x[6] - x[7] - x[8] - x[9] + x[10] + x[11];
|
||||
out[3] = + x[0] - x[1] - x[2] - x[3] + x[4] - x[5] - x[6] + x[7] - x[8] + x[9] + x[10] - x[11];
|
||||
out[4] = + x[0] + x[1] + x[2] + x[3] + x[4] - x[5] + x[6] + x[7] - x[8] - x[9] - x[10] - x[11];
|
||||
out[5] = + x[0] - x[1] + x[2] - x[3] - x[4] - x[5] + x[6] - x[7] - x[8] + x[9] - x[10] + x[11];
|
||||
out[6] = + x[0] + x[1] - x[2] - x[3] + x[4] + x[5] + x[6] - x[7] + x[8] + x[9] - x[10] - x[11];
|
||||
out[7] = + x[0] - x[1] - x[2] + x[3] + x[4] - x[5] - x[6] - x[7] + x[8] - x[9] - x[10] + x[11];
|
||||
out[8] = + x[0] + x[1] - x[2] - x[3] - x[4] - x[5] + x[6] + x[7] + x[8] - x[9] + x[10] + x[11];
|
||||
out[9] = + x[0] - x[1] - x[2] + x[3] - x[4] + x[5] + x[6] - x[7] - x[8] - x[9] + x[10] - x[11];
|
||||
out[10] = + x[0] + x[1] + x[2] + x[3] - x[4] - x[5] - x[6] - x[7] + x[8] + x[9] + x[10] - x[11];
|
||||
out[11] = + x[0] - x[1] + x[2] - x[3] - x[4] + x[5] - x[6] + x[7] + x[8] - x[9] - x[10] - x[11];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 12; i++) { x[i] = out[i]; }
|
||||
}
|
||||
|
||||
|
||||
__device__ __forceinline__ void hadamard_mult_thread_20(float x[20]) {
|
||||
float out[20];
|
||||
out[0] = + x[0] - x[1] - x[2] - x[3] - x[4] + x[5] - x[6] - x[7] - x[8] - x[9] + x[10] + x[11] - x[12] - x[13] + x[14] + x[15] - x[16] + x[17] + x[18] - x[19];
|
||||
out[1] = - x[0] + x[1] - x[2] - x[3] - x[4] - x[5] + x[6] - x[7] - x[8] - x[9] + x[10] + x[11] + x[12] - x[13] - x[14] - x[15] + x[16] - x[17] + x[18] + x[19];
|
||||
out[2] = - x[0] - x[1] + x[2] - x[3] - x[4] - x[5] - x[6] + x[7] - x[8] - x[9] - x[10] + x[11] + x[12] + x[13] - x[14] + x[15] - x[16] + x[17] - x[18] + x[19];
|
||||
out[3] = - x[0] - x[1] - x[2] + x[3] - x[4] - x[5] - x[6] - x[7] + x[8] - x[9] - x[10] - x[11] + x[12] + x[13] + x[14] + x[15] + x[16] - x[17] + x[18] - x[19];
|
||||
out[4] = - x[0] - x[1] - x[2] - x[3] + x[4] - x[5] - x[6] - x[7] - x[8] + x[9] + x[10] - x[11] - x[12] + x[13] + x[14] - x[15] + x[16] + x[17] - x[18] + x[19];
|
||||
out[5] = - x[0] + x[1] + x[2] + x[3] + x[4] + x[5] - x[6] - x[7] - x[8] - x[9] - x[10] + x[11] - x[12] - x[13] + x[14] + x[15] + x[16] - x[17] - x[18] + x[19];
|
||||
out[6] = + x[0] - x[1] + x[2] + x[3] + x[4] - x[5] + x[6] - x[7] - x[8] - x[9] + x[10] - x[11] + x[12] - x[13] - x[14] + x[15] + x[16] + x[17] - x[18] - x[19];
|
||||
out[7] = + x[0] + x[1] - x[2] + x[3] + x[4] - x[5] - x[6] + x[7] - x[8] - x[9] - x[10] + x[11] - x[12] + x[13] - x[14] - x[15] + x[16] + x[17] + x[18] - x[19];
|
||||
out[8] = + x[0] + x[1] + x[2] - x[3] + x[4] - x[5] - x[6] - x[7] + x[8] - x[9] - x[10] - x[11] + x[12] - x[13] + x[14] - x[15] - x[16] + x[17] + x[18] + x[19];
|
||||
out[9] = + x[0] + x[1] + x[2] + x[3] - x[4] - x[5] - x[6] - x[7] - x[8] + x[9] + x[10] - x[11] - x[12] + x[13] - x[14] + x[15] - x[16] - x[17] + x[18] + x[19];
|
||||
out[10] = - x[0] - x[1] + x[2] + x[3] - x[4] + x[5] - x[6] + x[7] + x[8] - x[9] + x[10] - x[11] - x[12] - x[13] - x[14] - x[15] + x[16] + x[17] + x[18] + x[19];
|
||||
out[11] = - x[0] - x[1] - x[2] + x[3] + x[4] - x[5] + x[6] - x[7] + x[8] + x[9] - x[10] + x[11] - x[12] - x[13] - x[14] + x[15] - x[16] + x[17] + x[18] + x[19];
|
||||
out[12] = + x[0] - x[1] - x[2] - x[3] + x[4] + x[5] - x[6] + x[7] - x[8] + x[9] - x[10] - x[11] + x[12] - x[13] - x[14] + x[15] + x[16] - x[17] + x[18] + x[19];
|
||||
out[13] = + x[0] + x[1] - x[2] - x[3] - x[4] + x[5] + x[6] - x[7] + x[8] - x[9] - x[10] - x[11] - x[12] + x[13] - x[14] + x[15] + x[16] + x[17] - x[18] + x[19];
|
||||
out[14] = - x[0] + x[1] + x[2] - x[3] - x[4] - x[5] + x[6] + x[7] - x[8] + x[9] - x[10] - x[11] - x[12] - x[13] + x[14] + x[15] + x[16] + x[17] + x[18] - x[19];
|
||||
out[15] = - x[0] + x[1] - x[2] - x[3] + x[4] - x[5] - x[6] + x[7] + x[8] - x[9] + x[10] - x[11] - x[12] - x[13] - x[14] + x[15] - x[16] - x[17] - x[18] - x[19];
|
||||
out[16] = + x[0] - x[1] + x[2] - x[3] - x[4] - x[5] - x[6] - x[7] + x[8] + x[9] - x[10] + x[11] - x[12] - x[13] - x[14] - x[15] + x[16] - x[17] - x[18] - x[19];
|
||||
out[17] = - x[0] + x[1] - x[2] + x[3] - x[4] + x[5] - x[6] - x[7] - x[8] + x[9] - x[10] - x[11] + x[12] - x[13] - x[14] - x[15] - x[16] + x[17] - x[18] - x[19];
|
||||
out[18] = - x[0] - x[1] + x[2] - x[3] + x[4] + x[5] + x[6] - x[7] - x[8] - x[9] - x[10] - x[11] - x[12] + x[13] - x[14] - x[15] - x[16] - x[17] + x[18] - x[19];
|
||||
out[19] = + x[0] - x[1] - x[2] + x[3] - x[4] - x[5] + x[6] + x[7] - x[8] - x[9] - x[10] - x[11] - x[12] - x[13] + x[14] - x[15] - x[16] - x[17] - x[18] + x[19];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 20; i++) { x[i] = out[i]; }
|
||||
}
|
||||
|
||||
|
||||
__device__ __forceinline__ void hadamard_mult_thread_28(float x[28]) {
|
||||
float out[28];
|
||||
out[0] = + x[0] - x[1] - x[2] - x[3] - x[4] - x[5] - x[6] + x[7] + x[8] - x[9] - x[10] - x[11] - x[12] + x[13] + x[14] - x[15] + x[16] - x[17] - x[18] + x[19] - x[20] + x[21] - x[22] - x[23] + x[24] + x[25] - x[26] - x[27];
|
||||
out[1] = - x[0] + x[1] - x[2] - x[3] - x[4] - x[5] - x[6] + x[7] + x[8] + x[9] - x[10] - x[11] - x[12] - x[13] - x[14] + x[15] - x[16] + x[17] - x[18] - x[19] + x[20] - x[21] + x[22] - x[23] - x[24] + x[25] + x[26] - x[27];
|
||||
out[2] = - x[0] - x[1] + x[2] - x[3] - x[4] - x[5] - x[6] - x[7] + x[8] + x[9] + x[10] - x[11] - x[12] - x[13] + x[14] - x[15] + x[16] - x[17] + x[18] - x[19] - x[20] - x[21] - x[22] + x[23] - x[24] - x[25] + x[26] + x[27];
|
||||
out[3] = - x[0] - x[1] - x[2] + x[3] - x[4] - x[5] - x[6] - x[7] - x[8] + x[9] + x[10] + x[11] - x[12] - x[13] - x[14] + x[15] - x[16] + x[17] - x[18] + x[19] - x[20] + x[21] - x[22] - x[23] + x[24] - x[25] - x[26] + x[27];
|
||||
out[4] = - x[0] - x[1] - x[2] - x[3] + x[4] - x[5] - x[6] - x[7] - x[8] - x[9] + x[10] + x[11] + x[12] - x[13] - x[14] - x[15] + x[16] - x[17] + x[18] - x[19] + x[20] + x[21] + x[22] - x[23] - x[24] + x[25] - x[26] - x[27];
|
||||
out[5] = - x[0] - x[1] - x[2] - x[3] - x[4] + x[5] - x[6] - x[7] - x[8] - x[9] - x[10] + x[11] + x[12] + x[13] + x[14] - x[15] - x[16] + x[17] - x[18] + x[19] - x[20] - x[21] + x[22] + x[23] - x[24] - x[25] + x[26] - x[27];
|
||||
out[6] = - x[0] - x[1] - x[2] - x[3] - x[4] - x[5] + x[6] + x[7] - x[8] - x[9] - x[10] - x[11] + x[12] + x[13] - x[14] + x[15] - x[16] - x[17] + x[18] - x[19] + x[20] - x[21] - x[22] + x[23] + x[24] - x[25] - x[26] + x[27];
|
||||
out[7] = - x[0] - x[1] + x[2] + x[3] + x[4] + x[5] - x[6] + x[7] - x[8] - x[9] - x[10] - x[11] - x[12] - x[13] - x[14] + x[15] + x[16] - x[17] - x[18] + x[19] + x[20] + x[21] - x[22] + x[23] - x[24] - x[25] + x[26] - x[27];
|
||||
out[8] = - x[0] - x[1] - x[2] + x[3] + x[4] + x[5] + x[6] - x[7] + x[8] - x[9] - x[10] - x[11] - x[12] - x[13] + x[14] - x[15] + x[16] + x[17] - x[18] - x[19] + x[20] - x[21] + x[22] - x[23] + x[24] - x[25] - x[26] + x[27];
|
||||
out[9] = + x[0] - x[1] - x[2] - x[3] + x[4] + x[5] + x[6] - x[7] - x[8] + x[9] - x[10] - x[11] - x[12] - x[13] + x[14] + x[15] - x[16] + x[17] + x[18] - x[19] - x[20] + x[21] - x[22] + x[23] - x[24] + x[25] - x[26] - x[27];
|
||||
out[10] = + x[0] + x[1] - x[2] - x[3] - x[4] + x[5] + x[6] - x[7] - x[8] - x[9] + x[10] - x[11] - x[12] - x[13] - x[14] + x[15] + x[16] - x[17] + x[18] + x[19] - x[20] - x[21] + x[22] - x[23] + x[24] - x[25] + x[26] - x[27];
|
||||
out[11] = + x[0] + x[1] + x[2] - x[3] - x[4] - x[5] + x[6] - x[7] - x[8] - x[9] - x[10] + x[11] - x[12] - x[13] - x[14] - x[15] + x[16] + x[17] - x[18] + x[19] + x[20] - x[21] - x[22] + x[23] - x[24] + x[25] - x[26] + x[27];
|
||||
out[12] = + x[0] + x[1] + x[2] + x[3] - x[4] - x[5] - x[6] - x[7] - x[8] - x[9] - x[10] - x[11] + x[12] - x[13] + x[14] - x[15] - x[16] + x[17] + x[18] - x[19] + x[20] + x[21] - x[22] - x[23] + x[24] - x[25] + x[26] - x[27];
|
||||
out[13] = - x[0] + x[1] + x[2] + x[3] + x[4] - x[5] - x[6] - x[7] - x[8] - x[9] - x[10] - x[11] - x[12] + x[13] + x[14] + x[15] - x[16] - x[17] + x[18] + x[19] - x[20] - x[21] + x[22] - x[23] - x[24] + x[25] - x[26] + x[27];
|
||||
out[14] = - x[0] + x[1] - x[2] + x[3] + x[4] - x[5] + x[6] + x[7] - x[8] - x[9] + x[10] + x[11] - x[12] - x[13] + x[14] - x[15] - x[16] - x[17] - x[18] - x[19] - x[20] - x[21] - x[22] + x[23] + x[24] + x[25] + x[26] - x[27];
|
||||
out[15] = + x[0] - x[1] + x[2] - x[3] + x[4] + x[5] - x[6] - x[7] + x[8] - x[9] - x[10] + x[11] + x[12] - x[13] - x[14] + x[15] - x[16] - x[17] - x[18] - x[19] - x[20] - x[21] - x[22] - x[23] + x[24] + x[25] + x[26] + x[27];
|
||||
out[16] = - x[0] + x[1] - x[2] + x[3] - x[4] + x[5] + x[6] - x[7] - x[8] + x[9] - x[10] - x[11] + x[12] + x[13] - x[14] - x[15] + x[16] - x[17] - x[18] - x[19] - x[20] + x[21] - x[22] - x[23] - x[24] + x[25] + x[26] + x[27];
|
||||
out[17] = + x[0] - x[1] + x[2] - x[3] + x[4] - x[5] + x[6] + x[7] - x[8] - x[9] + x[10] - x[11] - x[12] + x[13] - x[14] - x[15] - x[16] + x[17] - x[18] - x[19] - x[20] + x[21] + x[22] - x[23] - x[24] - x[25] + x[26] + x[27];
|
||||
out[18] = + x[0] + x[1] - x[2] + x[3] - x[4] + x[5] - x[6] + x[7] + x[8] - x[9] - x[10] + x[11] - x[12] - x[13] - x[14] - x[15] - x[16] - x[17] + x[18] - x[19] - x[20] + x[21] + x[22] + x[23] - x[24] - x[25] - x[26] + x[27];
|
||||
out[19] = - x[0] + x[1] + x[2] - x[3] + x[4] - x[5] + x[6] - x[7] + x[8] + x[9] - x[10] - x[11] + x[12] - x[13] - x[14] - x[15] - x[16] - x[17] - x[18] + x[19] - x[20] + x[21] + x[22] + x[23] + x[24] - x[25] - x[26] - x[27];
|
||||
out[20] = + x[0] - x[1] + x[2] + x[3] - x[4] + x[5] - x[6] - x[7] - x[8] + x[9] + x[10] - x[11] - x[12] + x[13] - x[14] - x[15] - x[16] - x[17] - x[18] - x[19] + x[20] - x[21] + x[22] + x[23] + x[24] + x[25] - x[26] - x[27];
|
||||
out[21] = - x[0] + x[1] + x[2] - x[3] - x[4] + x[5] + x[6] - x[7] + x[8] - x[9] + x[10] + x[11] - x[12] + x[13] + x[14] + x[15] - x[16] - x[17] - x[18] - x[19] + x[20] + x[21] - x[22] - x[23] - x[24] - x[25] - x[26] - x[27];
|
||||
out[22] = + x[0] - x[1] + x[2] + x[3] - x[4] - x[5] + x[6] + x[7] - x[8] + x[9] - x[10] + x[11] + x[12] - x[13] + x[14] + x[15] + x[16] - x[17] - x[18] - x[19] - x[20] - x[21] + x[22] - x[23] - x[24] - x[25] - x[26] - x[27];
|
||||
out[23] = + x[0] + x[1] - x[2] + x[3] + x[4] - x[5] - x[6] - x[7] + x[8] - x[9] + x[10] - x[11] + x[12] + x[13] - x[14] + x[15] + x[16] + x[17] - x[18] - x[19] - x[20] - x[21] - x[22] + x[23] - x[24] - x[25] - x[26] - x[27];
|
||||
out[24] = - x[0] + x[1] + x[2] - x[3] + x[4] + x[5] - x[6] + x[7] - x[8] + x[9] - x[10] + x[11] - x[12] + x[13] - x[14] - x[15] + x[16] + x[17] + x[18] - x[19] - x[20] - x[21] - x[22] - x[23] + x[24] - x[25] - x[26] - x[27];
|
||||
out[25] = - x[0] - x[1] + x[2] + x[3] - x[4] + x[5] + x[6] + x[7] + x[8] - x[9] + x[10] - x[11] + x[12] - x[13] - x[14] - x[15] - x[16] + x[17] + x[18] + x[19] - x[20] - x[21] - x[22] - x[23] - x[24] + x[25] - x[26] - x[27];
|
||||
out[26] = + x[0] - x[1] - x[2] + x[3] + x[4] - x[5] + x[6] - x[7] + x[8] + x[9] - x[10] + x[11] - x[12] + x[13] - x[14] - x[15] - x[16] - x[17] + x[18] + x[19] + x[20] - x[21] - x[22] - x[23] - x[24] - x[25] + x[26] - x[27];
|
||||
out[27] = + x[0] + x[1] - x[2] - x[3] + x[4] + x[5] - x[6] + x[7] - x[8] + x[9] + x[10] - x[11] + x[12] - x[13] + x[14] - x[15] - x[16] - x[17] - x[18] + x[19] + x[20] - x[21] - x[22] - x[23] - x[24] - x[25] - x[26] + x[27];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 28; i++) { x[i] = out[i]; }
|
||||
}
|
||||
|
||||
|
||||
__device__ __forceinline__ void hadamard_mult_thread_40(float x[40]) {
|
||||
float out[40];
|
||||
out[0] = + x[0] - x[1] - x[2] - x[3] - x[4] - x[5] - x[6] - x[7] - x[8] - x[9] - x[10] - x[11] - x[12] - x[13] - x[14] - x[15] - x[16] - x[17] - x[18] - x[19] + x[20] - x[21] - x[22] - x[23] - x[24] - x[25] - x[26] - x[27] - x[28] - x[29] - x[30] - x[31] - x[32] - x[33] - x[34] - x[35] - x[36] - x[37] - x[38] - x[39];
|
||||
out[1] = + x[0] + x[1] - x[2] + x[3] + x[4] - x[5] - x[6] - x[7] - x[8] + x[9] - x[10] + x[11] - x[12] + x[13] + x[14] + x[15] + x[16] - x[17] - x[18] + x[19] + x[20] + x[21] - x[22] + x[23] + x[24] - x[25] - x[26] - x[27] - x[28] + x[29] - x[30] + x[31] - x[32] + x[33] + x[34] + x[35] + x[36] - x[37] - x[38] + x[39];
|
||||
out[2] = + x[0] + x[1] + x[2] - x[3] + x[4] + x[5] - x[6] - x[7] - x[8] - x[9] + x[10] - x[11] + x[12] - x[13] + x[14] + x[15] + x[16] + x[17] - x[18] - x[19] + x[20] + x[21] + x[22] - x[23] + x[24] + x[25] - x[26] - x[27] - x[28] - x[29] + x[30] - x[31] + x[32] - x[33] + x[34] + x[35] + x[36] + x[37] - x[38] - x[39];
|
||||
out[3] = + x[0] - x[1] + x[2] + x[3] - x[4] + x[5] + x[6] - x[7] - x[8] - x[9] - x[10] + x[11] - x[12] + x[13] - x[14] + x[15] + x[16] + x[17] + x[18] - x[19] + x[20] - x[21] + x[22] + x[23] - x[24] + x[25] + x[26] - x[27] - x[28] - x[29] - x[30] + x[31] - x[32] + x[33] - x[34] + x[35] + x[36] + x[37] + x[38] - x[39];
|
||||
out[4] = + x[0] - x[1] - x[2] + x[3] + x[4] - x[5] + x[6] + x[7] - x[8] - x[9] - x[10] - x[11] + x[12] - x[13] + x[14] - x[15] + x[16] + x[17] + x[18] + x[19] + x[20] - x[21] - x[22] + x[23] + x[24] - x[25] + x[26] + x[27] - x[28] - x[29] - x[30] - x[31] + x[32] - x[33] + x[34] - x[35] + x[36] + x[37] + x[38] + x[39];
|
||||
out[5] = + x[0] + x[1] - x[2] - x[3] + x[4] + x[5] - x[6] + x[7] + x[8] - x[9] - x[10] - x[11] - x[12] + x[13] - x[14] + x[15] - x[16] + x[17] + x[18] + x[19] + x[20] + x[21] - x[22] - x[23] + x[24] + x[25] - x[26] + x[27] + x[28] - x[29] - x[30] - x[31] - x[32] + x[33] - x[34] + x[35] - x[36] + x[37] + x[38] + x[39];
|
||||
out[6] = + x[0] + x[1] + x[2] - x[3] - x[4] + x[5] + x[6] - x[7] + x[8] + x[9] - x[10] - x[11] - x[12] - x[13] + x[14] - x[15] + x[16] - x[17] + x[18] + x[19] + x[20] + x[21] + x[22] - x[23] - x[24] + x[25] + x[26] - x[27] + x[28] + x[29] - x[30] - x[31] - x[32] - x[33] + x[34] - x[35] + x[36] - x[37] + x[38] + x[39];
|
||||
out[7] = + x[0] + x[1] + x[2] + x[3] - x[4] - x[5] + x[6] + x[7] - x[8] + x[9] + x[10] - x[11] - x[12] - x[13] - x[14] + x[15] - x[16] + x[17] - x[18] + x[19] + x[20] + x[21] + x[22] + x[23] - x[24] - x[25] + x[26] + x[27] - x[28] + x[29] + x[30] - x[31] - x[32] - x[33] - x[34] + x[35] - x[36] + x[37] - x[38] + x[39];
|
||||
out[8] = + x[0] + x[1] + x[2] + x[3] + x[4] - x[5] - x[6] + x[7] + x[8] - x[9] + x[10] + x[11] - x[12] - x[13] - x[14] - x[15] + x[16] - x[17] + x[18] - x[19] + x[20] + x[21] + x[22] + x[23] + x[24] - x[25] - x[26] + x[27] + x[28] - x[29] + x[30] + x[31] - x[32] - x[33] - x[34] - x[35] + x[36] - x[37] + x[38] - x[39];
|
||||
out[9] = + x[0] - x[1] + x[2] + x[3] + x[4] + x[5] - x[6] - x[7] + x[8] + x[9] - x[10] + x[11] + x[12] - x[13] - x[14] - x[15] - x[16] + x[17] - x[18] + x[19] + x[20] - x[21] + x[22] + x[23] + x[24] + x[25] - x[26] - x[27] + x[28] + x[29] - x[30] + x[31] + x[32] - x[33] - x[34] - x[35] - x[36] + x[37] - x[38] + x[39];
|
||||
out[10] = + x[0] + x[1] - x[2] + x[3] + x[4] + x[5] + x[6] - x[7] - x[8] + x[9] + x[10] - x[11] + x[12] + x[13] - x[14] - x[15] - x[16] - x[17] + x[18] - x[19] + x[20] + x[21] - x[22] + x[23] + x[24] + x[25] + x[26] - x[27] - x[28] + x[29] + x[30] - x[31] + x[32] + x[33] - x[34] - x[35] - x[36] - x[37] + x[38] - x[39];
|
||||
out[11] = + x[0] - x[1] + x[2] - x[3] + x[4] + x[5] + x[6] + x[7] - x[8] - x[9] + x[10] + x[11] - x[12] + x[13] + x[14] - x[15] - x[16] - x[17] - x[18] + x[19] + x[20] - x[21] + x[22] - x[23] + x[24] + x[25] + x[26] + x[27] - x[28] - x[29] + x[30] + x[31] - x[32] + x[33] + x[34] - x[35] - x[36] - x[37] - x[38] + x[39];
|
||||
out[12] = + x[0] + x[1] - x[2] + x[3] - x[4] + x[5] + x[6] + x[7] + x[8] - x[9] - x[10] + x[11] + x[12] - x[13] + x[14] + x[15] - x[16] - x[17] - x[18] - x[19] + x[20] + x[21] - x[22] + x[23] - x[24] + x[25] + x[26] + x[27] + x[28] - x[29] - x[30] + x[31] + x[32] - x[33] + x[34] + x[35] - x[36] - x[37] - x[38] - x[39];
|
||||
out[13] = + x[0] - x[1] + x[2] - x[3] + x[4] - x[5] + x[6] + x[7] + x[8] + x[9] - x[10] - x[11] + x[12] + x[13] - x[14] + x[15] + x[16] - x[17] - x[18] - x[19] + x[20] - x[21] + x[22] - x[23] + x[24] - x[25] + x[26] + x[27] + x[28] + x[29] - x[30] - x[31] + x[32] + x[33] - x[34] + x[35] + x[36] - x[37] - x[38] - x[39];
|
||||
out[14] = + x[0] - x[1] - x[2] + x[3] - x[4] + x[5] - x[6] + x[7] + x[8] + x[9] + x[10] - x[11] - x[12] + x[13] + x[14] - x[15] + x[16] + x[17] - x[18] - x[19] + x[20] - x[21] - x[22] + x[23] - x[24] + x[25] - x[26] + x[27] + x[28] + x[29] + x[30] - x[31] - x[32] + x[33] + x[34] - x[35] + x[36] + x[37] - x[38] - x[39];
|
||||
out[15] = + x[0] - x[1] - x[2] - x[3] + x[4] - x[5] + x[6] - x[7] + x[8] + x[9] + x[10] + x[11] - x[12] - x[13] + x[14] + x[15] - x[16] + x[17] + x[18] - x[19] + x[20] - x[21] - x[22] - x[23] + x[24] - x[25] + x[26] - x[27] + x[28] + x[29] + x[30] + x[31] - x[32] - x[33] + x[34] + x[35] - x[36] + x[37] + x[38] - x[39];
|
||||
out[16] = + x[0] - x[1] - x[2] - x[3] - x[4] + x[5] - x[6] + x[7] - x[8] + x[9] + x[10] + x[11] + x[12] - x[13] - x[14] + x[15] + x[16] - x[17] + x[18] + x[19] + x[20] - x[21] - x[22] - x[23] - x[24] + x[25] - x[26] + x[27] - x[28] + x[29] + x[30] + x[31] + x[32] - x[33] - x[34] + x[35] + x[36] - x[37] + x[38] + x[39];
|
||||
out[17] = + x[0] + x[1] - x[2] - x[3] - x[4] - x[5] + x[6] - x[7] + x[8] - x[9] + x[10] + x[11] + x[12] + x[13] - x[14] - x[15] + x[16] + x[17] - x[18] + x[19] + x[20] + x[21] - x[22] - x[23] - x[24] - x[25] + x[26] - x[27] + x[28] - x[29] + x[30] + x[31] + x[32] + x[33] - x[34] - x[35] + x[36] + x[37] - x[38] + x[39];
|
||||
out[18] = + x[0] + x[1] + x[2] - x[3] - x[4] - x[5] - x[6] + x[7] - x[8] + x[9] - x[10] + x[11] + x[12] + x[13] + x[14] - x[15] - x[16] + x[17] + x[18] - x[19] + x[20] + x[21] + x[22] - x[23] - x[24] - x[25] - x[26] + x[27] - x[28] + x[29] - x[30] + x[31] + x[32] + x[33] + x[34] - x[35] - x[36] + x[37] + x[38] - x[39];
|
||||
out[19] = + x[0] - x[1] + x[2] + x[3] - x[4] - x[5] - x[6] - x[7] + x[8] - x[9] + x[10] - x[11] + x[12] + x[13] + x[14] + x[15] - x[16] - x[17] + x[18] + x[19] + x[20] - x[21] + x[22] + x[23] - x[24] - x[25] - x[26] - x[27] + x[28] - x[29] + x[30] - x[31] + x[32] + x[33] + x[34] + x[35] - x[36] - x[37] + x[38] + x[39];
|
||||
out[20] = + x[0] - x[1] - x[2] - x[3] - x[4] - x[5] - x[6] - x[7] - x[8] - x[9] - x[10] - x[11] - x[12] - x[13] - x[14] - x[15] - x[16] - x[17] - x[18] - x[19] - x[20] + x[21] + x[22] + x[23] + x[24] + x[25] + x[26] + x[27] + x[28] + x[29] + x[30] + x[31] + x[32] + x[33] + x[34] + x[35] + x[36] + x[37] + x[38] + x[39];
|
||||
out[21] = + x[0] + x[1] - x[2] + x[3] + x[4] - x[5] - x[6] - x[7] - x[8] + x[9] - x[10] + x[11] - x[12] + x[13] + x[14] + x[15] + x[16] - x[17] - x[18] + x[19] - x[20] - x[21] + x[22] - x[23] - x[24] + x[25] + x[26] + x[27] + x[28] - x[29] + x[30] - x[31] + x[32] - x[33] - x[34] - x[35] - x[36] + x[37] + x[38] - x[39];
|
||||
out[22] = + x[0] + x[1] + x[2] - x[3] + x[4] + x[5] - x[6] - x[7] - x[8] - x[9] + x[10] - x[11] + x[12] - x[13] + x[14] + x[15] + x[16] + x[17] - x[18] - x[19] - x[20] - x[21] - x[22] + x[23] - x[24] - x[25] + x[26] + x[27] + x[28] + x[29] - x[30] + x[31] - x[32] + x[33] - x[34] - x[35] - x[36] - x[37] + x[38] + x[39];
|
||||
out[23] = + x[0] - x[1] + x[2] + x[3] - x[4] + x[5] + x[6] - x[7] - x[8] - x[9] - x[10] + x[11] - x[12] + x[13] - x[14] + x[15] + x[16] + x[17] + x[18] - x[19] - x[20] + x[21] - x[22] - x[23] + x[24] - x[25] - x[26] + x[27] + x[28] + x[29] + x[30] - x[31] + x[32] - x[33] + x[34] - x[35] - x[36] - x[37] - x[38] + x[39];
|
||||
out[24] = + x[0] - x[1] - x[2] + x[3] + x[4] - x[5] + x[6] + x[7] - x[8] - x[9] - x[10] - x[11] + x[12] - x[13] + x[14] - x[15] + x[16] + x[17] + x[18] + x[19] - x[20] + x[21] + x[22] - x[23] - x[24] + x[25] - x[26] - x[27] + x[28] + x[29] + x[30] + x[31] - x[32] + x[33] - x[34] + x[35] - x[36] - x[37] - x[38] - x[39];
|
||||
out[25] = + x[0] + x[1] - x[2] - x[3] + x[4] + x[5] - x[6] + x[7] + x[8] - x[9] - x[10] - x[11] - x[12] + x[13] - x[14] + x[15] - x[16] + x[17] + x[18] + x[19] - x[20] - x[21] + x[22] + x[23] - x[24] - x[25] + x[26] - x[27] - x[28] + x[29] + x[30] + x[31] + x[32] - x[33] + x[34] - x[35] + x[36] - x[37] - x[38] - x[39];
|
||||
out[26] = + x[0] + x[1] + x[2] - x[3] - x[4] + x[5] + x[6] - x[7] + x[8] + x[9] - x[10] - x[11] - x[12] - x[13] + x[14] - x[15] + x[16] - x[17] + x[18] + x[19] - x[20] - x[21] - x[22] + x[23] + x[24] - x[25] - x[26] + x[27] - x[28] - x[29] + x[30] + x[31] + x[32] + x[33] - x[34] + x[35] - x[36] + x[37] - x[38] - x[39];
|
||||
out[27] = + x[0] + x[1] + x[2] + x[3] - x[4] - x[5] + x[6] + x[7] - x[8] + x[9] + x[10] - x[11] - x[12] - x[13] - x[14] + x[15] - x[16] + x[17] - x[18] + x[19] - x[20] - x[21] - x[22] - x[23] + x[24] + x[25] - x[26] - x[27] + x[28] - x[29] - x[30] + x[31] + x[32] + x[33] + x[34] - x[35] + x[36] - x[37] + x[38] - x[39];
|
||||
out[28] = + x[0] + x[1] + x[2] + x[3] + x[4] - x[5] - x[6] + x[7] + x[8] - x[9] + x[10] + x[11] - x[12] - x[13] - x[14] - x[15] + x[16] - x[17] + x[18] - x[19] - x[20] - x[21] - x[22] - x[23] - x[24] + x[25] + x[26] - x[27] - x[28] + x[29] - x[30] - x[31] + x[32] + x[33] + x[34] + x[35] - x[36] + x[37] - x[38] + x[39];
|
||||
out[29] = + x[0] - x[1] + x[2] + x[3] + x[4] + x[5] - x[6] - x[7] + x[8] + x[9] - x[10] + x[11] + x[12] - x[13] - x[14] - x[15] - x[16] + x[17] - x[18] + x[19] - x[20] + x[21] - x[22] - x[23] - x[24] - x[25] + x[26] + x[27] - x[28] - x[29] + x[30] - x[31] - x[32] + x[33] + x[34] + x[35] + x[36] - x[37] + x[38] - x[39];
|
||||
out[30] = + x[0] + x[1] - x[2] + x[3] + x[4] + x[5] + x[6] - x[7] - x[8] + x[9] + x[10] - x[11] + x[12] + x[13] - x[14] - x[15] - x[16] - x[17] + x[18] - x[19] - x[20] - x[21] + x[22] - x[23] - x[24] - x[25] - x[26] + x[27] + x[28] - x[29] - x[30] + x[31] - x[32] - x[33] + x[34] + x[35] + x[36] + x[37] - x[38] + x[39];
|
||||
out[31] = + x[0] - x[1] + x[2] - x[3] + x[4] + x[5] + x[6] + x[7] - x[8] - x[9] + x[10] + x[11] - x[12] + x[13] + x[14] - x[15] - x[16] - x[17] - x[18] + x[19] - x[20] + x[21] - x[22] + x[23] - x[24] - x[25] - x[26] - x[27] + x[28] + x[29] - x[30] - x[31] + x[32] - x[33] - x[34] + x[35] + x[36] + x[37] + x[38] - x[39];
|
||||
out[32] = + x[0] + x[1] - x[2] + x[3] - x[4] + x[5] + x[6] + x[7] + x[8] - x[9] - x[10] + x[11] + x[12] - x[13] + x[14] + x[15] - x[16] - x[17] - x[18] - x[19] - x[20] - x[21] + x[22] - x[23] + x[24] - x[25] - x[26] - x[27] - x[28] + x[29] + x[30] - x[31] - x[32] + x[33] - x[34] - x[35] + x[36] + x[37] + x[38] + x[39];
|
||||
out[33] = + x[0] - x[1] + x[2] - x[3] + x[4] - x[5] + x[6] + x[7] + x[8] + x[9] - x[10] - x[11] + x[12] + x[13] - x[14] + x[15] + x[16] - x[17] - x[18] - x[19] - x[20] + x[21] - x[22] + x[23] - x[24] + x[25] - x[26] - x[27] - x[28] - x[29] + x[30] + x[31] - x[32] - x[33] + x[34] - x[35] - x[36] + x[37] + x[38] + x[39];
|
||||
out[34] = + x[0] - x[1] - x[2] + x[3] - x[4] + x[5] - x[6] + x[7] + x[8] + x[9] + x[10] - x[11] - x[12] + x[13] + x[14] - x[15] + x[16] + x[17] - x[18] - x[19] - x[20] + x[21] + x[22] - x[23] + x[24] - x[25] + x[26] - x[27] - x[28] - x[29] - x[30] + x[31] + x[32] - x[33] - x[34] + x[35] - x[36] - x[37] + x[38] + x[39];
|
||||
out[35] = + x[0] - x[1] - x[2] - x[3] + x[4] - x[5] + x[6] - x[7] + x[8] + x[9] + x[10] + x[11] - x[12] - x[13] + x[14] + x[15] - x[16] + x[17] + x[18] - x[19] - x[20] + x[21] + x[22] + x[23] - x[24] + x[25] - x[26] + x[27] - x[28] - x[29] - x[30] - x[31] + x[32] + x[33] - x[34] - x[35] + x[36] - x[37] - x[38] + x[39];
|
||||
out[36] = + x[0] - x[1] - x[2] - x[3] - x[4] + x[5] - x[6] + x[7] - x[8] + x[9] + x[10] + x[11] + x[12] - x[13] - x[14] + x[15] + x[16] - x[17] + x[18] + x[19] - x[20] + x[21] + x[22] + x[23] + x[24] - x[25] + x[26] - x[27] + x[28] - x[29] - x[30] - x[31] - x[32] + x[33] + x[34] - x[35] - x[36] + x[37] - x[38] - x[39];
|
||||
out[37] = + x[0] + x[1] - x[2] - x[3] - x[4] - x[5] + x[6] - x[7] + x[8] - x[9] + x[10] + x[11] + x[12] + x[13] - x[14] - x[15] + x[16] + x[17] - x[18] + x[19] - x[20] - x[21] + x[22] + x[23] + x[24] + x[25] - x[26] + x[27] - x[28] + x[29] - x[30] - x[31] - x[32] - x[33] + x[34] + x[35] - x[36] - x[37] + x[38] - x[39];
|
||||
out[38] = + x[0] + x[1] + x[2] - x[3] - x[4] - x[5] - x[6] + x[7] - x[8] + x[9] - x[10] + x[11] + x[12] + x[13] + x[14] - x[15] - x[16] + x[17] + x[18] - x[19] - x[20] - x[21] - x[22] + x[23] + x[24] + x[25] + x[26] - x[27] + x[28] - x[29] + x[30] - x[31] - x[32] - x[33] - x[34] + x[35] + x[36] - x[37] - x[38] + x[39];
|
||||
out[39] = + x[0] - x[1] + x[2] + x[3] - x[4] - x[5] - x[6] - x[7] + x[8] - x[9] + x[10] - x[11] + x[12] + x[13] + x[14] + x[15] - x[16] - x[17] + x[18] + x[19] - x[20] + x[21] - x[22] - x[23] + x[24] + x[25] + x[26] + x[27] - x[28] + x[29] - x[30] + x[31] - x[32] - x[33] - x[34] - x[35] + x[36] + x[37] - x[38] - x[39];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 40; i++) { x[i] = out[i]; }
|
||||
}
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
// Inspired by https://github.com/NVIDIA/DALI/blob/main/include/dali/core/static_switch.h
|
||||
// and https://github.com/pytorch/pytorch/blob/master/aten/src/ATen/Dispatch.h
|
||||
|
||||
#pragma once
|
||||
|
||||
/// @param COND - a boolean expression to switch by
|
||||
/// @param CONST_NAME - a name given for the constexpr bool variable.
|
||||
/// @param ... - code to execute for true and false
|
||||
///
|
||||
/// Usage:
|
||||
/// ```
|
||||
/// BOOL_SWITCH(flag, BoolConst, [&] {
|
||||
/// some_function<BoolConst>(...);
|
||||
/// });
|
||||
/// ```
|
||||
#define BOOL_SWITCH(COND, CONST_NAME, ...) \
|
||||
if (COND) { \
|
||||
constexpr static bool CONST_NAME = true; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
constexpr static bool CONST_NAME = false; \
|
||||
__VA_ARGS__ \
|
||||
} \
|
||||
@@ -0,0 +1,31 @@
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include <torch/extension.h>
|
||||
#include <torch/python.h>
|
||||
|
||||
|
||||
at::Tensor rms_norm_rope(at::Tensor &x, c10::optional<at::Tensor>& weights_, at::Tensor &cos_freqs, at::Tensor &sin_freqs, bool out_16bit);
|
||||
|
||||
at::Tensor fp6_pack(at::Tensor &x);
|
||||
at::Tensor fp6_unpack(at::Tensor &x, int64_t original_n);
|
||||
|
||||
at::Tensor rms_norm_split_rope(
|
||||
at::Tensor &x,
|
||||
at::Tensor &sin_freqs,
|
||||
at::Tensor &cos_freqs,
|
||||
at::Tensor &weights,
|
||||
bool out_fp8
|
||||
);
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("rms_norm_rope", &rms_norm_rope,
|
||||
"fused norm + rope + cvt");
|
||||
m.def("fp6_pack", &fp6_pack,
|
||||
"Pack 8-bit to 6-bit by dropping e_1 and e_2 bits");
|
||||
m.def("fp6_unpack", &fp6_unpack,
|
||||
"Unpack 6-bit to 8-bit (with e_1 and e_2 set to 0)");
|
||||
m.def("rms_norm_split_rope", &rms_norm_split_rope,
|
||||
"RMS norm + split RoPE with optional FP8 output");
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
/******************************************************************************
|
||||
* Copyright (c) 2023, Tri Dao.
|
||||
******************************************************************************/
|
||||
|
||||
// Host entry point for the fused RMS-norm + RoPE kernel.
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <torch/extension.h>
|
||||
#include <vector>
|
||||
|
||||
#include "fast_hadamard_transform.h"
|
||||
|
||||
#define CHECK_SHAPE(x, ...) TORCH_CHECK(x.sizes() == torch::IntArrayRef({__VA_ARGS__}), #x " must have shape (" #__VA_ARGS__ ")")
|
||||
|
||||
template<typename input_t, typename output_t, bool norm_affine>
|
||||
void rms_norm_rope_cuda(NormRopeHadamardParamsBase ¶ms, cudaStream_t stream);
|
||||
|
||||
void set_norm_rope_hadamard_params(NormRopeHadamardParamsBase ¶ms,
|
||||
// sizes
|
||||
const size_t batch,
|
||||
const size_t dim,
|
||||
const size_t multiple,
|
||||
// device pointers
|
||||
const at::Tensor x,
|
||||
const at::Tensor cos_freqs,
|
||||
const at::Tensor sin_freqs,
|
||||
const at::Tensor weights,
|
||||
const at::Tensor out,
|
||||
|
||||
bool norm_affine,
|
||||
float scale
|
||||
) {
|
||||
|
||||
// Reset the parameters
|
||||
memset(¶ms, 0, sizeof(params));
|
||||
|
||||
params.batch = batch;
|
||||
params.dim = dim;
|
||||
params.log_N = int(ceil(std::log2(dim / multiple)));
|
||||
|
||||
// Set the pointers and strides.
|
||||
params.x_ptr = x.data_ptr();
|
||||
params.out_ptr = out.data_ptr();
|
||||
params.cos_freq_ptr = cos_freqs.data_ptr();
|
||||
params.sin_freq_ptr = sin_freqs.data_ptr();
|
||||
if (norm_affine){
|
||||
params.weights_ptr = weights.data_ptr();
|
||||
} else {
|
||||
params.weights_ptr = nullptr;
|
||||
}
|
||||
// All stride are in elements, not bytes.
|
||||
params.x_batch_stride = x.stride(0);
|
||||
params.out_batch_stride = out.stride(0);
|
||||
params.cos_freq_batch_stride = cos_freqs.stride(0);
|
||||
params.sin_freq_batch_stride = sin_freqs.stride(0);
|
||||
|
||||
params.scale = scale;
|
||||
|
||||
}
|
||||
|
||||
at::Tensor rms_norm_rope(at::Tensor &x, c10::optional<at::Tensor>& weights_, at::Tensor &cos_freqs, at::Tensor &sin_freqs, bool out_16bit) {
|
||||
auto input_type = x.scalar_type();
|
||||
float scale = 1.0f; // :D
|
||||
TORCH_CHECK(input_type == at::ScalarType::BFloat16);
|
||||
TORCH_CHECK(x.is_cuda());
|
||||
const auto shapes_og = x.sizes();
|
||||
const int dim_og = x.size(-1);
|
||||
x = x.reshape({-1, dim_og});
|
||||
if (x.stride(-1) != 1) { x = x.contiguous(); }
|
||||
const auto sizes = x.sizes();
|
||||
const int batch_size = sizes[0];
|
||||
cos_freqs = cos_freqs.reshape({-1, dim_og});
|
||||
sin_freqs = sin_freqs.reshape({-1, dim_og});
|
||||
at::Tensor weights;
|
||||
bool norm_affine = false;
|
||||
if(weights_.has_value()){
|
||||
weights = weights_.value();
|
||||
norm_affine = true;
|
||||
}
|
||||
CHECK_SHAPE(x, batch_size, dim_og);
|
||||
TORCH_CHECK(x.stride(1) == 1);
|
||||
if (dim_og % 8 != 0) {
|
||||
x = torch::nn::functional::pad(x, torch::nn::functional::PadFuncOptions({0, 8 - dim_og % 8}));
|
||||
}
|
||||
const int dim = x.size(1);
|
||||
at::Tensor out;
|
||||
if (out_16bit){
|
||||
out = torch::empty(x.sizes(), x.options().dtype(torch::kBFloat16));
|
||||
} else {
|
||||
out = torch::empty(x.sizes(), x.options().dtype(torch::kFloat8_e4m3fn));
|
||||
}
|
||||
at::cuda::CUDAGuard device_guard{(char)x.get_device()};
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
NormRopeHadamardParamsBase params;
|
||||
set_norm_rope_hadamard_params(params, batch_size, dim, 1, x, cos_freqs, sin_freqs, weights, out, norm_affine, scale);
|
||||
TORCH_CHECK(dim % 8 == 0, "fast_hadamard_transform only supports hidden dimension divisible by 8 for now");
|
||||
TORCH_CHECK(dim <= 32768, "fast_hadamard_transform only supports hidden dimension at most 32768 for now");
|
||||
if (norm_affine){
|
||||
if (out_16bit){
|
||||
rms_norm_rope_cuda<at::BFloat16, at::BFloat16, true>(params, stream);
|
||||
} else {
|
||||
rms_norm_rope_cuda<at::BFloat16, at::Float8_e4m3fn, true>(params, stream);
|
||||
}
|
||||
|
||||
} else {
|
||||
if (out_16bit){
|
||||
rms_norm_rope_cuda<at::BFloat16, at::BFloat16, false>(params, stream);
|
||||
} else {
|
||||
rms_norm_rope_cuda<at::BFloat16, at::Float8_e4m3fn, false>(params, stream);
|
||||
}
|
||||
}
|
||||
return out.reshape(shapes_og);
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
/******************************************************************************
|
||||
* Copyright (c) 2023, Tri Dao.
|
||||
******************************************************************************/
|
||||
|
||||
// #pragma once
|
||||
|
||||
#include <c10/util/BFloat16.h>
|
||||
#include <c10/util/Half.h>
|
||||
#include <c10/util/Float8_e4m3fn.h>
|
||||
#include <c10/cuda/CUDAException.h> // For C10_CUDA_CHECK and C10_CUDA_KERNEL_LAUNCH_CHECK
|
||||
|
||||
#include "fast_hadamard_transform.h"
|
||||
#include "fast_hadamard_transform_common.h"
|
||||
#include "fast_hadamard_transform_special.h"
|
||||
#include "static_switch.h"
|
||||
|
||||
|
||||
template<int kNThreads_, int kLogN_, typename input_t_, typename output_t_, bool norm_affine_>
|
||||
struct norm_rope_kernel_traits {
|
||||
using input_t = input_t_;
|
||||
using output_t = output_t_;
|
||||
|
||||
static constexpr int kNThreads = kNThreads_;
|
||||
static constexpr int kLogN = kLogN_;
|
||||
static constexpr int N = 1 << kLogN;
|
||||
static constexpr int kNBytes = sizeof(input_t);
|
||||
static constexpr int OutkNBytes = sizeof(output_t);
|
||||
|
||||
static constexpr bool norm_affine = norm_affine_;
|
||||
|
||||
static_assert(kNBytes == 1 || kNBytes == 2 || kNBytes == 4);
|
||||
static constexpr int kNElts = kNBytes == 4 ? 4 : kNBytes == 2 ? 8 : 8;
|
||||
// It's possible that we need to do 2 rounds of exchange if input_t is 16 bits
|
||||
// (since then we'd have 8 values of float, and each round we can exchange 4 floats).
|
||||
static constexpr int kNExchangePerVec = sizeof(float) / sizeof(input_t);
|
||||
|
||||
using vec_t = typename BytesToType<kNBytes * kNElts>::Type;
|
||||
using vec_t_out = typename BytesToType<OutkNBytes * kNElts>::Type;
|
||||
|
||||
static constexpr int kNChunks = N / (kNElts * kNThreads);
|
||||
// We don't want to use more than 32 KB of shared memory.
|
||||
static constexpr int kSmemExchangeSize = std::min(N * 4, 32 * 1024);
|
||||
static constexpr int kNExchangeRounds = N * 4 / kSmemExchangeSize;
|
||||
static_assert(kNExchangeRounds * kSmemExchangeSize == N * 4);
|
||||
static constexpr int kSmemSize = kSmemExchangeSize;
|
||||
};
|
||||
|
||||
|
||||
template<typename Ktraits>
|
||||
__global__ __launch_bounds__(Ktraits::kNThreads)
|
||||
void norm_rope_cvt_kernel(NormRopeHadamardParamsBase params) {
|
||||
constexpr int kNThreads = Ktraits::kNThreads;
|
||||
constexpr int kNElts = Ktraits::kNElts;
|
||||
constexpr int kNExchangePerVec = Ktraits::kNExchangePerVec;
|
||||
constexpr int kNExchangeRounds = Ktraits::kNExchangeRounds;
|
||||
constexpr int kNChunks = Ktraits::kNChunks;
|
||||
constexpr bool norm_affine = Ktraits::norm_affine;
|
||||
|
||||
using input_t = typename Ktraits::input_t;
|
||||
using output_t = typename Ktraits::output_t;
|
||||
using vec_t = typename Ktraits::vec_t;
|
||||
using out_vec_t = typename Ktraits::vec_t_out;
|
||||
using weights_t = typename Ktraits::input_t;
|
||||
using freqs_t = typename Ktraits::input_t;
|
||||
|
||||
constexpr int kLogNElts = cilog2(Ktraits::kNElts);
|
||||
static_assert(1 << kLogNElts == kNElts, "kNElts must be a power of 2");
|
||||
constexpr int kWarpSize = std::min(kNThreads, 32);
|
||||
constexpr int kLogWarpSize = cilog2(kWarpSize);
|
||||
static_assert(1 << kLogWarpSize == kWarpSize, "Warp size must be a power of 2");
|
||||
constexpr int kNWarps = kNThreads / kWarpSize;
|
||||
constexpr int kLogNWarps = cilog2(kNWarps);
|
||||
static_assert(1 << kLogNWarps == kNWarps, "kNWarps must be a power of 2");
|
||||
constexpr int kLoadsPerExchange = Ktraits::kSmemExchangeSize / (sizeof(vec_t) * kNThreads);
|
||||
static_assert(kLoadsPerExchange * sizeof(vec_t) * kNThreads == Ktraits::kSmemExchangeSize, "kSmemExchangeSize should be a power of 2");
|
||||
static_assert(kNExchangeRounds * kLoadsPerExchange * sizeof(vec_t) == kNChunks * kNElts * sizeof(float));
|
||||
|
||||
constexpr int kChunksPerExchange = Ktraits::kSmemExchangeSize / (sizeof(vec_t) * kNExchangePerVec * kNThreads);
|
||||
static_assert(kChunksPerExchange * sizeof(vec_t) * kNExchangePerVec * kNThreads == Ktraits::kSmemExchangeSize);
|
||||
constexpr int kNExchanges = kNChunks / kChunksPerExchange;
|
||||
static_assert(kNExchanges * kChunksPerExchange == kNChunks);
|
||||
|
||||
// Shared memory.
|
||||
extern __shared__ char smem_[];
|
||||
vec_t *smem_exchange = reinterpret_cast<vec_t *>(smem_);
|
||||
|
||||
const int batch_id = blockIdx.x;
|
||||
const int warp_id = threadIdx.x / 32;
|
||||
|
||||
input_t *x = reinterpret_cast<input_t *>(params.x_ptr) + batch_id * params.x_batch_stride;
|
||||
output_t *out = reinterpret_cast<output_t *>(params.out_ptr) + batch_id * params.out_batch_stride;
|
||||
weights_t *weights = norm_affine ? reinterpret_cast<weights_t*>(params.weights_ptr) : nullptr;
|
||||
|
||||
float x_vals[kNChunks][kNElts];
|
||||
float weights_vals[kNChunks][kNElts];
|
||||
|
||||
load_input<kNChunks, kNElts, input_t>(x, x_vals, params.dim);
|
||||
|
||||
//RMS Norm START
|
||||
float thread_squared_sum = 0.0f;
|
||||
#pragma unroll
|
||||
for (size_t c = 0; c < kNChunks; c++)
|
||||
{
|
||||
#pragma unroll
|
||||
for (size_t i = 0; i < kNElts; i++)
|
||||
{
|
||||
thread_squared_sum += x_vals[c][i] * x_vals[c][i];
|
||||
}
|
||||
|
||||
}
|
||||
SumOp<float> sum_op;
|
||||
float warp_sum = Allreduce<32>::run(thread_squared_sum, sum_op);
|
||||
float *smem_sum = reinterpret_cast<float*>(smem_);
|
||||
if(threadIdx.x % 32 == 0){
|
||||
smem_sum[warp_id] = warp_sum;
|
||||
}
|
||||
__syncthreads();
|
||||
float norm = 0.0f;
|
||||
#pragma unroll
|
||||
for (size_t i = 0; i < kNWarps; i++)
|
||||
{
|
||||
norm += smem_sum[i];
|
||||
}
|
||||
|
||||
norm *= 1.0f/params.dim;
|
||||
norm = rsqrtf(norm);
|
||||
if constexpr (norm_affine){
|
||||
load_input<kNChunks, kNElts, weights_t>(weights, weights_vals, params.dim);
|
||||
}
|
||||
#pragma unroll
|
||||
for (size_t c = 0; c < kNChunks; c++)
|
||||
{
|
||||
#pragma unroll
|
||||
for (size_t i = 0; i < kNElts; i++)
|
||||
{
|
||||
if constexpr (norm_affine){
|
||||
x_vals[c][i] *= (norm * weights_vals[c][i]);
|
||||
} else {
|
||||
x_vals[c][i] *= norm;
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
//RMS NORM END
|
||||
|
||||
//ROPE START
|
||||
float sin_freqs_vals[kNChunks][kNElts];
|
||||
float cos_freqs_vals[kNChunks][kNElts];
|
||||
|
||||
freqs_t *cos_freqs = reinterpret_cast<freqs_t*>(params.cos_freq_ptr) + batch_id * params.cos_freq_batch_stride;
|
||||
freqs_t *sin_freqs = reinterpret_cast<freqs_t*>(params.sin_freq_ptr) + batch_id * params.sin_freq_batch_stride;
|
||||
|
||||
load_input<kNChunks, kNElts, freqs_t>(cos_freqs, cos_freqs_vals, params.dim);
|
||||
load_input<kNChunks, kNElts, freqs_t>(sin_freqs, sin_freqs_vals, params.dim);
|
||||
|
||||
#pragma unroll
|
||||
for (size_t c = 0; c < kNChunks; c++)
|
||||
{
|
||||
#pragma unroll
|
||||
for (size_t i = 0; i < kNElts; i+=2)
|
||||
{
|
||||
float x_1 = x_vals[c][i];
|
||||
float x_2 = x_vals[c][i+1];
|
||||
x_vals[c][i] = -x_2*sin_freqs_vals[c][i] + x_1*cos_freqs_vals[c][i];
|
||||
x_vals[c][i+1] = x_1*sin_freqs_vals[c][i+1] + x_2*cos_freqs_vals[c][i+1];
|
||||
}
|
||||
}
|
||||
//ROPE END
|
||||
|
||||
store_output<kNChunks, kNElts, output_t, false>(out, x_vals, params.dim, params.scale);
|
||||
}
|
||||
|
||||
template<int kNThreads, int kLogN, typename input_t, typename output_t, bool norm_affine>
|
||||
void norm_rope_cvt_launch(NormRopeHadamardParamsBase ¶ms, cudaStream_t stream) {
|
||||
using Ktraits = norm_rope_kernel_traits<kNThreads, kLogN, input_t, output_t, norm_affine>;
|
||||
constexpr int kSmemSize = Ktraits::kSmemSize;
|
||||
dim3 grid(params.batch);
|
||||
auto kernel = &norm_rope_cvt_kernel<Ktraits>;
|
||||
if (kSmemSize >= 48 * 1024) {
|
||||
C10_CUDA_CHECK(cudaFuncSetAttribute(
|
||||
kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, kSmemSize));
|
||||
}
|
||||
kernel<<<grid, Ktraits::kNThreads, kSmemSize, stream>>>(params);
|
||||
C10_CUDA_KERNEL_LAUNCH_CHECK();
|
||||
}
|
||||
|
||||
template<typename input_t, typename output_t, bool norm_affine>
|
||||
void rms_norm_rope_cuda(NormRopeHadamardParamsBase ¶ms, cudaStream_t stream) {
|
||||
if (params.log_N == 3) {
|
||||
norm_rope_cvt_launch<1, 3, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 4) {
|
||||
norm_rope_cvt_launch<2, 4, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 5) {
|
||||
norm_rope_cvt_launch<4, 5, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 6) {
|
||||
norm_rope_cvt_launch<8, 6, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 7) {
|
||||
norm_rope_cvt_launch<16, 7, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 8) {
|
||||
norm_rope_cvt_launch<32, 8, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 9) {
|
||||
norm_rope_cvt_launch<32, 9, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 10) {
|
||||
norm_rope_cvt_launch<128, 10, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 11) {
|
||||
norm_rope_cvt_launch<256, 11, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 12) {
|
||||
norm_rope_cvt_launch<256, 12, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 13) {
|
||||
norm_rope_cvt_launch<256, 13, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 14) {
|
||||
norm_rope_cvt_launch<256, 14, input_t, output_t, norm_affine>(params, stream);
|
||||
} else if (params.log_N == 15) {
|
||||
norm_rope_cvt_launch<256, 15, input_t, output_t, norm_affine>(params, stream);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
template void rms_norm_rope_cuda<at::BFloat16, at::Float8_e4m3fn, false>(NormRopeHadamardParamsBase ¶ms, cudaStream_t stream);
|
||||
template void rms_norm_rope_cuda<at::BFloat16, at::Float8_e4m3fn, true>(NormRopeHadamardParamsBase ¶ms, cudaStream_t stream);
|
||||
|
||||
template void rms_norm_rope_cuda<at::BFloat16, at::BFloat16, false>(NormRopeHadamardParamsBase ¶ms, cudaStream_t stream);
|
||||
template void rms_norm_rope_cuda<at::BFloat16, at::BFloat16, true>(NormRopeHadamardParamsBase ¶ms, cudaStream_t stream);
|
||||
|
||||
// template void fast_hadamard_transform_cuda<at::BFloat16, at::BFloat16>(HadamardParamsBase ¶ms, cudaStream_t stream);
|
||||
|
||||
// template void fast_hadamard_transform_cuda<at::Float8_e4m3fn, at::BFloat16>(HadamardParamsBase ¶ms, cudaStream_t stream);
|
||||
// template void fast_hadamard_transform_cuda<at::Float8_e4m3fn, at::Float8_e4m3fn>(HadamardParamsBase ¶ms, cudaStream_t stream);
|
||||
@@ -0,0 +1,111 @@
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <torch/extension.h>
|
||||
#include <torch/python.h>
|
||||
|
||||
#include <vector>
|
||||
|
||||
// Forward declaration of CUDA kernel template
|
||||
template<typename out_t>
|
||||
void rms_norm_split_rope_cuda(
|
||||
void* x,
|
||||
void* sin_freqs,
|
||||
void* cos_freqs,
|
||||
void* weights,
|
||||
int b,
|
||||
int s,
|
||||
int n,
|
||||
int h,
|
||||
long cos_sb, long cos_sn, long cos_ss,
|
||||
long sin_sb, long sin_sn, long sin_ss,
|
||||
void* out,
|
||||
cudaStream_t stream
|
||||
);
|
||||
|
||||
at::Tensor rms_norm_split_rope(
|
||||
at::Tensor &x,
|
||||
at::Tensor &sin_freqs,
|
||||
at::Tensor &cos_freqs,
|
||||
at::Tensor &weights,
|
||||
bool out_fp8
|
||||
) {
|
||||
TORCH_CHECK(x.scalar_type() == at::ScalarType::BFloat16, "Input must be BFloat16");
|
||||
TORCH_CHECK(sin_freqs.scalar_type() == at::ScalarType::BFloat16, "sin_freqs must be BFloat16");
|
||||
TORCH_CHECK(cos_freqs.scalar_type() == at::ScalarType::BFloat16, "cos_freqs must be BFloat16");
|
||||
TORCH_CHECK(x.is_cuda(), "Input must be on CUDA");
|
||||
TORCH_CHECK(sin_freqs.is_cuda(), "sin_freqs must be on CUDA");
|
||||
TORCH_CHECK(cos_freqs.is_cuda(), "cos_freqs must be on CUDA");
|
||||
|
||||
// Get dimensions
|
||||
// x: [b, s, h]
|
||||
// cos, sin: [b, n, s, d] where n*d = h/2
|
||||
int b = x.size(0);
|
||||
int s = x.size(1);
|
||||
int h = x.size(2);
|
||||
|
||||
TORCH_CHECK(cos_freqs.dim() == 4, "cos_freqs must be 4D");
|
||||
TORCH_CHECK(sin_freqs.dim() == 4, "sin_freqs must be 4D");
|
||||
|
||||
int n = cos_freqs.size(1);
|
||||
int d = h / n;
|
||||
|
||||
|
||||
// Require a contiguous innermost (d/2) dim for the vectorized int4 freq load,
|
||||
// but keep the outer (b, n, s) strides: apply_split_rotary_emb hands us a
|
||||
// swapaxes view (logical [b, n, s, d/2], physical [b, s, n, d/2]) whose inner
|
||||
// stride is already 1, so this never copies it. The strides are forwarded to
|
||||
// the kernel so the read is correct regardless of the physical layout.
|
||||
if (x.stride(-1) != 1) { x = x.contiguous(); }
|
||||
if (cos_freqs.stride(-1) != 1) { cos_freqs = cos_freqs.contiguous(); }
|
||||
if (sin_freqs.stride(-1) != 1) { sin_freqs = sin_freqs.contiguous(); }
|
||||
|
||||
long cos_sb = cos_freqs.stride(0), cos_sn = cos_freqs.stride(1), cos_ss = cos_freqs.stride(2);
|
||||
long sin_sb = sin_freqs.stride(0), sin_sn = sin_freqs.stride(1), sin_ss = sin_freqs.stride(2);
|
||||
|
||||
// Create output tensor
|
||||
at::Tensor out;
|
||||
if (out_fp8) {
|
||||
out = torch::empty(x.sizes(), x.options().dtype(torch::kFloat8_e4m3fn));
|
||||
} else {
|
||||
out = torch::empty(x.sizes(), x.options().dtype(torch::kBFloat16));
|
||||
}
|
||||
|
||||
// Setup CUDA
|
||||
at::cuda::CUDAGuard device_guard{(char)x.get_device()};
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
// Launch kernel
|
||||
if (out_fp8) {
|
||||
rms_norm_split_rope_cuda<at::Float8_e4m3fn>(
|
||||
x.data_ptr(),
|
||||
sin_freqs.data_ptr(),
|
||||
cos_freqs.data_ptr(),
|
||||
weights.data_ptr(), // weights (optional, not used yet)
|
||||
b,
|
||||
s,
|
||||
n,
|
||||
h,
|
||||
cos_sb, cos_sn, cos_ss,
|
||||
sin_sb, sin_sn, sin_ss,
|
||||
(void*)out.data_ptr(),
|
||||
stream
|
||||
);
|
||||
} else {
|
||||
rms_norm_split_rope_cuda<at::BFloat16>(
|
||||
x.data_ptr(),
|
||||
sin_freqs.data_ptr(),
|
||||
cos_freqs.data_ptr(),
|
||||
weights.data_ptr(), // weights (optional, not used yet)
|
||||
b,
|
||||
s,
|
||||
n,
|
||||
h,
|
||||
cos_sb, cos_sn, cos_ss,
|
||||
sin_sb, sin_sn, sin_ss,
|
||||
(void*)out.data_ptr(),
|
||||
stream
|
||||
);
|
||||
}
|
||||
|
||||
return out;
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
#include <c10/util/BFloat16.h>
|
||||
#include <c10/util/Float8_e4m3fn.h>
|
||||
#include <c10/cuda/CUDAException.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp8.h>
|
||||
|
||||
// CUDA kernel template for RMS norm + split RoPE
|
||||
// out_t can be at::Float8_e4m3fn or at::BFloat16
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
using fp8 = __nv_fp8_e4m3;
|
||||
__device__ __forceinline__ void _load_x(bf16* x, float x_vals[8], int h){
|
||||
bf16 x_tmp[8];
|
||||
*reinterpret_cast<int4*>(x_tmp) = reinterpret_cast<int4*>(x + blockIdx.x * h)[threadIdx.x];
|
||||
#pragma unroll
|
||||
for(int i = 0; i < 8; i++){
|
||||
x_vals[i] = float(x_tmp[i]);
|
||||
}
|
||||
}
|
||||
// Load 8 freq values for this thread from a table laid out logically as
|
||||
// [b, n, s, d/2] (what apply_split_rotary_emb produces -- a swapaxes view whose
|
||||
// physical layout is [b, s, n, d/2]). The strides (sb, sn, ss; inner d/2 stride
|
||||
// is 1) are forwarded from the host so the read is correct for both that
|
||||
// non-contiguous view and a genuinely contiguous [b, n, s, d/2] tensor. Both
|
||||
// head-halves map to the same freq element (mirrors the eager cos.unsqueeze(-2)).
|
||||
__device__ __forceinline__ void _load_freqs(
|
||||
const bf16* freqs, float x_vals[8], int s, int d, long sb, long sn, long ss
|
||||
){
|
||||
bf16 x_tmp[8];
|
||||
int threads_per_head = d / 8;
|
||||
int head_idx = threadIdx.x / threads_per_head;
|
||||
int lane = threadIdx.x % (threads_per_head / 2);
|
||||
int b_idx = blockIdx.x / s;
|
||||
int t_idx = blockIdx.x % s;
|
||||
long off = b_idx * sb + head_idx * sn + t_idx * ss + (long)lane * 8;
|
||||
*reinterpret_cast<int4*>(x_tmp) = *reinterpret_cast<const int4*>(freqs + off);
|
||||
#pragma unroll
|
||||
for(int i = 0; i < 8; i++){
|
||||
x_vals[i] = float(x_tmp[i]);
|
||||
}
|
||||
}
|
||||
|
||||
template<typename out_t>
|
||||
__global__ void _rms_norm_split_rope_kernel(bf16* x, bf16* sin_freqs, bf16* cos_freqs, void* out, bf16* weights, int b, int s, int n, int h,
|
||||
long cos_sb, long cos_sn, long cos_ss, long sin_sb, long sin_sn, long sin_ss){
|
||||
int token_idx = blockIdx.x;
|
||||
int tid = threadIdx.x;
|
||||
int lane_id = tid % 32;
|
||||
// freqs have shape [b, s, h/2]
|
||||
// each thread block calculate one row
|
||||
// there are h/8 threads in thread block, each thread processes 8 values
|
||||
// num_of_rows = b * s
|
||||
// freqs have h/2 dim
|
||||
// gridDim is (num_of_rows, 1, 1)
|
||||
// calculate rms norm x_normed = x/x_norm * weights. x_norm is calculated across row, it means thread block wide sum reduction
|
||||
|
||||
extern __shared__ float smem[];
|
||||
|
||||
// Step 1: Load input values (8 per thread)
|
||||
float x_vals[8];
|
||||
_load_x(x, x_vals, h);
|
||||
|
||||
float sum_sq = 0.0f;
|
||||
#pragma unroll
|
||||
for(int i = 0; i < 8; i++){
|
||||
sum_sq += x_vals[i] * x_vals[i];
|
||||
}
|
||||
|
||||
// Warp-level reduction
|
||||
#pragma unroll
|
||||
for(int offset = 16; offset > 0; offset >>= 1){
|
||||
sum_sq += __shfl_xor_sync(0xffffffff, sum_sq, offset);
|
||||
}
|
||||
|
||||
if(tid % 32 == 0){
|
||||
smem[tid / 32] = sum_sq;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Final reduction across warps
|
||||
if(tid == 0){
|
||||
float total_sum = 0.0f;
|
||||
int num_warps = blockDim.x / 32;
|
||||
for(int i = 0; i < num_warps; i++){
|
||||
total_sum += smem[i];
|
||||
}
|
||||
// RMS: sqrt(mean(x^2))
|
||||
float rms = rsqrtf(total_sum / h + 1e-6f); // Add epsilon for numerical stability
|
||||
smem[0] = rms;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
float inv_rms = smem[0];
|
||||
|
||||
// Step 3: Apply RMS normalization (and weights if provided)
|
||||
#pragma unroll
|
||||
for(int i = 0; i < 8; i++){
|
||||
x_vals[i] *= inv_rms;
|
||||
// TODO: Apply weights if provided
|
||||
if(weights != nullptr) x_vals[i] *= float(weights[tid * 8 + i]);
|
||||
}
|
||||
|
||||
// Step 4: Calculate dimensions for split RoPE
|
||||
// Conceptually: [b, s, h] -> [b, s, n, 2*d] -> [b, s, n, 2, d]
|
||||
// where h = n * 2 * d
|
||||
int d = h / n;
|
||||
float x_other_vals[8];
|
||||
|
||||
int threads_per_head = d / 8;
|
||||
int head_idx = tid / threads_per_head;
|
||||
int idx_in_head = tid % threads_per_head;
|
||||
bool is_first_half = idx_in_head < (threads_per_head / 2);
|
||||
|
||||
// LT-PATCH: full-warp mask. The original (1u << threads_per_head) - 1 only marks
|
||||
// the first head's lanes active, so lanes belonging to heads beyond the first are
|
||||
// not in the mask -> __shfl_xor_sync result is undefined and can corrupt RoPE. The
|
||||
// XOR pattern keeps data within each power-of-two head group, so a full-warp mask
|
||||
// is correct for every lane.
|
||||
const unsigned mask = 0xffffffffu;
|
||||
const int laneMask = threads_per_head / 2; // 4, 8, or 16
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 8; i++) {
|
||||
x_other_vals[i] = __shfl_xor_sync(mask, x_vals[i], laneMask);
|
||||
}
|
||||
|
||||
float cos_vals[8], sin_vals[8];
|
||||
_load_freqs(cos_freqs, cos_vals, s, d, cos_sb, cos_sn, cos_ss);
|
||||
_load_freqs(sin_freqs, sin_vals, s, d, sin_sb, sin_sn, sin_ss);
|
||||
#pragma unroll
|
||||
for(int i = 0; i < 8; i++){
|
||||
x_vals[i] = cos_vals[i]*x_vals[i];
|
||||
}
|
||||
|
||||
|
||||
float sign = is_first_half ? -1.0f : 1.0f;
|
||||
for(int i = 0; i < 8; i++){
|
||||
x_vals[i] += sign*sin_vals[i]*x_other_vals[i];
|
||||
}
|
||||
|
||||
// Step 6: Convert and store output
|
||||
if constexpr (std::is_same_v<out_t, at::Float8_e4m3fn>){
|
||||
fp8 out_tmp[8];
|
||||
#pragma unroll
|
||||
for(int i = 0; i < 8; i++){
|
||||
out_tmp[i] = fp8(x_vals[i]);
|
||||
}
|
||||
*reinterpret_cast<int64_t*>((fp8*)out + token_idx * h + tid * 8) = *reinterpret_cast<int64_t*>(out_tmp);
|
||||
} else {
|
||||
bf16 out_tmp[8];
|
||||
#pragma unroll
|
||||
for(int i = 0; i < 8; i++){
|
||||
out_tmp[i] = __float2bfloat16(x_vals[i]);
|
||||
}
|
||||
*reinterpret_cast<int4*>((bf16*)out + token_idx * h + tid * 8) = *reinterpret_cast<int4*>(out_tmp);
|
||||
}
|
||||
}
|
||||
|
||||
template<typename out_t>
|
||||
void rms_norm_split_rope_cuda(
|
||||
void* x, // Input: [b, s, h]
|
||||
void* sin_freqs, // Sin frequencies: [b, n, s, d]
|
||||
void* cos_freqs, // Cos frequencies: [b, n, s, d]
|
||||
void* weights,
|
||||
int b, // Batch size
|
||||
int s, // Sequence length
|
||||
int n, // Number of heads (32)
|
||||
int h, // Hidden dimension (2048, 4096, or 8192)
|
||||
long cos_sb, long cos_sn, long cos_ss, // cos_freqs strides (b, n, s)
|
||||
long sin_sb, long sin_sn, long sin_ss, // sin_freqs strides (b, n, s)
|
||||
void* out, // Output: [b, s, h]
|
||||
cudaStream_t stream
|
||||
) {
|
||||
int num_tokens = b * s;
|
||||
int num_threads = h / 8; // Each thread processes 8 elements
|
||||
int smem_size = (num_threads / 32 + 1) * sizeof(float); // Shared memory for reductions
|
||||
|
||||
dim3 grid(num_tokens);
|
||||
dim3 block(num_threads);
|
||||
|
||||
_rms_norm_split_rope_kernel<out_t><<<grid, block, smem_size, stream>>>(
|
||||
reinterpret_cast<bf16*>(x),
|
||||
reinterpret_cast<bf16*>(sin_freqs),
|
||||
reinterpret_cast<bf16*>(cos_freqs),
|
||||
out,
|
||||
reinterpret_cast<bf16*>(weights),
|
||||
b, s, n, h,
|
||||
cos_sb, cos_sn, cos_ss,
|
||||
sin_sb, sin_sn, sin_ss
|
||||
);
|
||||
|
||||
C10_CUDA_KERNEL_LAUNCH_CHECK();
|
||||
}
|
||||
|
||||
// Explicit template instantiations
|
||||
template void rms_norm_split_rope_cuda<at::BFloat16>(
|
||||
void*, void*, void*, void*, int, int, int, int, long, long, long, long, long, long, void*, cudaStream_t
|
||||
);
|
||||
|
||||
template void rms_norm_split_rope_cuda<at::Float8_e4m3fn>(
|
||||
void*, void*, void*, void*, int, int, int, int, long, long, long, long, long, long, void*, cudaStream_t
|
||||
);
|
||||
Reference in New Issue
Block a user