/* * Online Softmax CUDA Kernel * * Implements numerically stable online softmax computation for Flash Attention. * This is used as a utility kernel for incremental softmax updates. */ #include #include #include // Constants #define WARP_SIZE 32 #define MAX_THREADS_PER_BLOCK 1024 namespace cg = cooperative_groups; // Online softmax state structure struct OnlineSoftmaxState { float m; // running max float l; // running sum __device__ OnlineSoftmaxState() : m(-INFINITY), l(0.0f) {} __device__ OnlineSoftmaxState(float max_val, float sum_val) : m(max_val), l(sum_val) {} __device__ void update(float x) { float m_new = fmaxf(m, x); float l_new = l * expf(m - m_new) + expf(x - m_new); m = m_new; l = l_new; } __device__ void merge(const OnlineSoftmaxState& other) { float m_new = fmaxf(m, other.m); float l_new = l * expf(m - m_new) + other.l * expf(other.m - m_new); m = m_new; l = l_new; } __device__ float get_log_sum_exp() const { return m + logf(l); } __device__ float get_normalization() const { return l; } }; // Warp-level reduction for online softmax __device__ OnlineSoftmaxState warp_reduce_softmax(OnlineSoftmaxState state) { for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) { float other_m = __shfl_down_sync(0xffffffff, state.m, offset); float other_l = __shfl_down_sync(0xffffffff, state.l, offset); OnlineSoftmaxState other(other_m, other_l); state.merge(other); } return state; } // Block-level reduction for online softmax __device__ OnlineSoftmaxState block_reduce_softmax(OnlineSoftmaxState state) { __shared__ OnlineSoftmaxState shared_states[32]; // Max warps per block int warp_id = threadIdx.x / WARP_SIZE; int lane_id = threadIdx.x % WARP_SIZE; // Warp-level reduce state = warp_reduce_softmax(state); // Store warp result if (lane_id == 0) { shared_states[warp_id] = state; } __syncthreads(); // Block-level reduce if (warp_id == 0) { OnlineSoftmaxState final_state; if (lane_id < (blockDim.x + WARP_SIZE - 1) / WARP_SIZE) { final_state = shared_states[lane_id]; } final_state = warp_reduce_softmax(final_state); // Broadcast result shared_states[0] = final_state; } __syncthreads(); return shared_states[0]; } // Online softmax kernel for a batch of sequences extern "C" __global__ void online_softmax_kernel( const float* __restrict__ scores, // Input scores [batch, seq_len, seq_len] float* __restrict__ probs, // Output probabilities [batch, seq_len, seq_len] float* __restrict__ lse, // Log-sum-exp [batch, seq_len] int batch_size, int seq_len, unsigned int causal ) { int batch_idx = blockIdx.x; int seq_idx = blockIdx.y; int tid = threadIdx.x; if (batch_idx >= batch_size || seq_idx >= seq_len) return; // Calculate offsets int batch_offset = batch_idx * seq_len * seq_len; int row_offset = batch_offset + seq_idx * seq_len; int lse_offset = batch_idx * seq_len + seq_idx; // Compute softmax for this sequence position OnlineSoftmaxState state; // First pass: compute max and sum for (int i = tid; i < seq_len; i += blockDim.x) { if (!causal || i <= seq_idx) { float score = scores[row_offset + i]; state.update(score); } } // Reduce across block state = block_reduce_softmax(state); // Second pass: compute probabilities for (int i = tid; i < seq_len; i += blockDim.x) { float prob = 0.0f; if (!causal || i <= seq_idx) { float score = scores[row_offset + i]; prob = expf(score - state.m) / state.l; } probs[row_offset + i] = prob; } // Store log-sum-exp if (tid == 0) { lse[lse_offset] = state.get_log_sum_exp(); } } // Incremental online softmax update kernel extern "C" __global__ void online_softmax_update_kernel( const float* __restrict__ new_scores, // New scores to add [batch, seq_len, new_len] const float* __restrict__ old_lse, // Previous log-sum-exp [batch, seq_len] float* __restrict__ updated_probs, // Updated probabilities [batch, seq_len, total_len] float* __restrict__ new_lse, // Updated log-sum-exp [batch, seq_len] int batch_size, int seq_len, int old_len, int new_len, unsigned int causal ) { int batch_idx = blockIdx.x; int seq_idx = blockIdx.y; int tid = threadIdx.x; if (batch_idx >= batch_size || seq_idx >= seq_len) return; // Calculate offsets int old_offset = batch_idx * seq_len * old_len + seq_idx * old_len; int new_offset = batch_idx * seq_len * new_len + seq_idx * new_len; int updated_offset = batch_idx * seq_len * (old_len + new_len) + seq_idx * (old_len + new_len); int lse_offset = batch_idx * seq_len + seq_idx; // Get old state float old_lse_val = old_lse[lse_offset]; OnlineSoftmaxState old_state(old_lse_val - logf(expf(old_lse_val)), expf(old_lse_val)); // Compute state for new scores OnlineSoftmaxState new_state; for (int i = tid; i < new_len; i += blockDim.x) { int global_pos = old_len + i; if (!causal || global_pos <= seq_idx) { float score = new_scores[new_offset + i]; new_state.update(score); } } // Reduce new state across block new_state = block_reduce_softmax(new_state); // Merge old and new states OnlineSoftmaxState merged_state = old_state; merged_state.merge(new_state); // Update old probabilities float scale_factor = old_state.l * expf(old_state.m - merged_state.m) / merged_state.l; for (int i = tid; i < old_len; i += blockDim.x) { if (!causal || i <= seq_idx) { updated_probs[updated_offset + i] = updated_probs[old_offset + i] * scale_factor; } else { updated_probs[updated_offset + i] = 0.0f; } } // Compute new probabilities for (int i = tid; i < new_len; i += blockDim.x) { int global_pos = old_len + i; float prob = 0.0f; if (!causal || global_pos <= seq_idx) { float score = new_scores[new_offset + i]; prob = expf(score - merged_state.m) / merged_state.l; } updated_probs[updated_offset + old_len + i] = prob; } // Store updated log-sum-exp if (tid == 0) { new_lse[lse_offset] = merged_state.get_log_sum_exp(); } } // Fused online softmax and value computation kernel extern "C" __global__ void online_softmax_value_kernel( const float* __restrict__ scores, // Input scores [batch, heads, seq_len, seq_len] const half* __restrict__ values, // Value tensor [batch, heads, seq_len, head_dim] half* __restrict__ output, // Output tensor [batch, heads, seq_len, head_dim] float* __restrict__ lse, // Log-sum-exp [batch, heads, seq_len] int batch_size, int num_heads, int seq_len, int head_dim, unsigned int causal ) { int batch_head_idx = blockIdx.x; int seq_idx = blockIdx.y; int tid = threadIdx.x; int batch_idx = batch_head_idx / num_heads; int head_idx = batch_head_idx % num_heads; if (batch_idx >= batch_size || head_idx >= num_heads || seq_idx >= seq_len) return; // Calculate offsets int scores_offset = batch_head_idx * seq_len * seq_len + seq_idx * seq_len; int values_offset = batch_head_idx * seq_len * head_dim; int output_offset = batch_head_idx * seq_len * head_dim + seq_idx * head_dim; int lse_offset = batch_head_idx * seq_len + seq_idx; // Shared memory for values extern __shared__ half shared_values[]; // Load values into shared memory for (int i = tid; i < seq_len * head_dim; i += blockDim.x) { shared_values[i] = values[values_offset + i]; } __syncthreads(); // Compute online softmax OnlineSoftmaxState state; for (int i = 0; i < seq_len; i++) { if (!causal || i <= seq_idx) { float score = scores[scores_offset + i]; state.update(score); } } // Compute output values for (int d = tid; d < head_dim; d += blockDim.x) { float output_val = 0.0f; for (int i = 0; i < seq_len; i++) { if (!causal || i <= seq_idx) { float score = scores[scores_offset + i]; float prob = expf(score - state.m) / state.l; float value = __half2float(shared_values[i * head_dim + d]); output_val += prob * value; } } output[output_offset + d] = __float2half(output_val); } // Store log-sum-exp if (tid == 0) { lse[lse_offset] = state.get_log_sum_exp(); } } // Vectorized online softmax kernel using half2 extern "C" __global__ void online_softmax_half2_kernel( const half* __restrict__ scores, // Input scores [batch, seq_len, seq_len] half* __restrict__ probs, // Output probabilities [batch, seq_len, seq_len] float* __restrict__ lse, // Log-sum-exp [batch, seq_len] int batch_size, int seq_len, unsigned int causal ) { int batch_idx = blockIdx.x; int seq_idx = blockIdx.y; int tid = threadIdx.x; if (batch_idx >= batch_size || seq_idx >= seq_len) return; // Calculate offsets int batch_offset = batch_idx * seq_len * seq_len; int row_offset = batch_offset + seq_idx * seq_len; int lse_offset = batch_idx * seq_len + seq_idx; // Compute softmax using vectorized operations OnlineSoftmaxState state; // First pass: compute max and sum using half2 for (int i = tid * 2; i < seq_len; i += blockDim.x * 2) { half2 scores_vec = *reinterpret_cast(&scores[row_offset + i]); if (!causal || i <= seq_idx) { float score1 = __half2float(scores_vec.x); state.update(score1); } if (i + 1 < seq_len && (!causal || i + 1 <= seq_idx)) { float score2 = __half2float(scores_vec.y); state.update(score2); } } // Reduce across block state = block_reduce_softmax(state); // Second pass: compute probabilities using half2 for (int i = tid * 2; i < seq_len; i += blockDim.x * 2) { half2 scores_vec = *reinterpret_cast(&scores[row_offset + i]); half2 probs_vec; if (!causal || i <= seq_idx) { float score1 = __half2float(scores_vec.x); float prob1 = expf(score1 - state.m) / state.l; probs_vec.x = __float2half(prob1); } else { probs_vec.x = __float2half(0.0f); } if (i + 1 < seq_len) { if (!causal || i + 1 <= seq_idx) { float score2 = __half2float(scores_vec.y); float prob2 = expf(score2 - state.m) / state.l; probs_vec.y = __float2half(prob2); } else { probs_vec.y = __float2half(0.0f); } } else { probs_vec.y = __float2half(0.0f); } *reinterpret_cast(&probs[row_offset + i]) = probs_vec; } // Store log-sum-exp if (tid == 0) { lse[lse_offset] = state.get_log_sum_exp(); } }