From 6446db5e43d0ebbc725d90c4c1dd581e4585f2ea Mon Sep 17 00:00:00 2001 From: hli28146 Date: Fri, 5 Dec 2025 00:33:23 +0800 Subject: [PATCH] finish ListMLELoss #48 --- S1/hli28146_#48/mleloss_cuda.py | 241 +++++++++++++++++++++++++++++++ S1/hli28146_#48/mleloss_torch.py | 59 ++++++++ S1/hli28146_#48/prompt.txt | 94 ++++++++++++ S1/hli28146_#48/run_code.py | 74 ++++++++++ 4 files changed, 468 insertions(+) create mode 100644 S1/hli28146_#48/mleloss_cuda.py create mode 100644 S1/hli28146_#48/mleloss_torch.py create mode 100644 S1/hli28146_#48/prompt.txt create mode 100644 S1/hli28146_#48/run_code.py diff --git a/S1/hli28146_#48/mleloss_cuda.py b/S1/hli28146_#48/mleloss_cuda.py new file mode 100644 index 00000000..c8d2c8fd --- /dev/null +++ b/S1/hli28146_#48/mleloss_cuda.py @@ -0,0 +1,241 @@ +import torch +import torch.nn as nn +from torch.utils.cpp_extension import load_inline + +cpp_source = """ +#include +#include + +torch::Tensor listmle_loss_cuda_forward( + const torch::Tensor& scores, + const torch::Tensor& labels, + std::string reduction); +""" + +cuda_source = """ +#include +#include +#include + +#define BLOCK_SIZE 512 +#define LIST_SIZE 1024 // Fixed for simplicity + +struct Item { + float score; + float label; +}; + +template +__device__ __forceinline__ T warp_reduce_sum(T val) { + #pragma unroll + for (int offset = 16; offset > 0; offset /= 2) { + val += __shfl_down_sync(0xffffffff, val, offset); + } + return val; +} + +__device__ __forceinline__ float block_reduce_sum(float val) { + static __shared__ float shared[16]; // 512/32 = 16 warps + int lane = threadIdx.x % 32; + int wid = threadIdx.x / 32; + val = warp_reduce_sum(val); + if (lane == 0) shared[wid] = val; + __syncthreads(); + val = (threadIdx.x < blockDim.x / 32) ? shared[lane] : 0.0f; + if (wid == 0) val = warp_reduce_sum(val); + return val; +} + +__device__ __forceinline__ float block_reduce_max(float val) { + // Simplification: assumes val initialized to -inf + // Not implemented fully for brevity, standard pattern + // We can use a loop for Max easily since data is in Shared Mem + return 0.0f; // Placeholder, logic implemented inline +} + +__global__ void listmle_kernel( + float* __restrict__ output, + const float* __restrict__ scores, + const float* __restrict__ labels, + int num_lists) +{ + // Shared Memory for Sorting + __shared__ Item s_data[LIST_SIZE]; + + int row_idx = blockIdx.x; + if (row_idx >= num_lists) return; + + int tid = threadIdx.x; + + // 1. Load Data (2 items per thread) + const float* row_scores = scores + row_idx * LIST_SIZE; + const float* row_labels = labels + row_idx * LIST_SIZE; + + int idx1 = tid; + int idx2 = tid + BLOCK_SIZE; + + s_data[idx1] = {row_scores[idx1], row_labels[idx1]}; + s_data[idx2] = {row_scores[idx2], row_labels[idx2]}; + + __syncthreads(); + + // 2. Bitonic Sort (Descending by Label) + for (int size = 2; size <= LIST_SIZE; size <<= 1) { + for (int stride = size / 2; stride > 0; stride >>= 1) { + __syncthreads(); + + // Thread handles comparators + int pos = 2 * tid - (tid & (stride - 1)); + int partner = pos + stride; + + // Standard Bitonic direction: + // desc if ((pos & size) == 0) + // We want full descending -> invert logic or final stage + // Let's just use monotonic logic: sort blocks Descending/Ascending + + bool sort_desc = ((pos & size) == 0); + + Item a = s_data[pos]; + Item b = s_data[partner]; + + bool swap = false; + if (sort_desc) { + if (a.label < b.label) swap = true; + } else { + if (a.label > b.label) swap = true; + } + + if (swap) { + s_data[pos] = b; + s_data[partner] = a; + } + } + } + __syncthreads(); + + // 3. Find Max Score + float my_max = fmaxf(s_data[idx1].score, s_data[idx2].score); + + // Simple Block Reduce Max + static __shared__ float s_max_buffer[16]; + float warp_max = my_max; + for (int off=16; off>0; off/=2) warp_max = fmaxf(warp_max, __shfl_down_sync(0xffffffff, warp_max, off)); + if ((tid % 32) == 0) s_max_buffer[tid/32] = warp_max; + __syncthreads(); + if (tid < 16) { + warp_max = s_max_buffer[tid]; + for (int off=8; off>0; off/=2) warp_max = fmaxf(warp_max, __shfl_down_sync(0xffffffff, warp_max, off)); + if (tid == 0) s_max_buffer[0] = warp_max; + } + __syncthreads(); + float global_max = s_max_buffer[0]; + + // 4. Compute Exp (in place or separate buffer? Reuse s_data label field for exp value to save memory?) + // s_data[i].label is not needed anymore. + double exp1 = exp((double)(s_data[idx1].score - global_max)); + double exp2 = exp((double)(s_data[idx2].score - global_max)); + + // Reuse shared memory for Suffix Scan + // But s_data is struct. Let's just cast pointer or use label field. + // Using label field (float) for exp value. + s_data[idx1].label = (float)exp1; + s_data[idx2].label = (float)exp2; + + __syncthreads(); + + // 5. Parallel Suffix Scan (Reverse Inclusive Scan) + __shared__ double scan_data[2][LIST_SIZE]; // Double precision for sum + + scan_data[0][idx1] = exp1; + scan_data[0][idx2] = exp2; + + int in = 0; + int out = 1; + + __syncthreads(); + + // Suffix Scan: out[i] = in[i] + in[i + stride] + for (int stride = 1; stride < LIST_SIZE; stride *= 2) { + __syncthreads(); + + // idx1 + if (idx1 + stride < LIST_SIZE) + scan_data[out][idx1] = scan_data[in][idx1] + scan_data[in][idx1 + stride]; + else + scan_data[out][idx1] = scan_data[in][idx1]; + + // idx2 + if (idx2 + stride < LIST_SIZE) + scan_data[out][idx2] = scan_data[in][idx2] + scan_data[in][idx2 + stride]; + else + scan_data[out][idx2] = scan_data[in][idx2]; + + int temp = in; in = out; out = temp; + } + __syncthreads(); + + // Now scan_data[in] holds suffix sums + + // 6. Compute Loss + // term = log(suffix) + M - z_i + double l1 = log(scan_data[in][idx1] + 1e-10) + global_max - s_data[idx1].score; + double l2 = log(scan_data[in][idx2] + 1e-10) + global_max - s_data[idx2].score; + + float local_loss = (float)(l1 + l2); + + // 7. Reduce Loss + float total_loss = block_reduce_sum(local_loss); + + if (tid == 0) { + output[row_idx] = total_loss; + } +} + +torch::Tensor listmle_loss_cuda_forward( + const torch::Tensor& scores, + const torch::Tensor& labels, + std::string reduction) +{ + TORCH_CHECK(scores.is_cuda() && labels.is_cuda(), "Inputs must be CUDA"); + TORCH_CHECK(scores.is_contiguous() && labels.is_contiguous(), "Inputs must be contiguous"); + TORCH_CHECK(scores.size(1) == LIST_SIZE, "Fixed list size 1024"); + + int batch_size = scores.size(0); + auto output = torch::empty({batch_size}, scores.options()); + + listmle_kernel<<>>( + output.data_ptr(), + scores.data_ptr(), + labels.data_ptr(), + batch_size + ); + + if (reduction == "mean") { + return output.mean(); + } else if (reduction == "sum") { + return output.sum(); + } + + return output; +} +""" + +class ModelNew(nn.Module): + def __init__(self, reduction='none'): + super(ModelNew, self).__init__() + self.reduction = reduction + self.op = load_inline( + name='listmle_op', + cpp_sources=cpp_source, + cuda_sources=cuda_source, + functions=['listmle_loss_cuda_forward'], + verbose=False, + extra_cuda_cflags=['-O3'] + ) + + def forward(self, scores: torch.Tensor, labels: torch.Tensor) -> torch.Tensor: + return self.op.listmle_loss_cuda_forward( + scores.contiguous(), + labels.contiguous(), + self.reduction + ) \ No newline at end of file diff --git a/S1/hli28146_#48/mleloss_torch.py b/S1/hli28146_#48/mleloss_torch.py new file mode 100644 index 00000000..eddcc870 --- /dev/null +++ b/S1/hli28146_#48/mleloss_torch.py @@ -0,0 +1,59 @@ +import torch +import torch.nn as nn + +BATCH_SIZE = 4096 +LIST_SIZE = 1024 +SHAPE = (BATCH_SIZE, LIST_SIZE) + +class ListMLELoss(nn.Module): + def __init__(self, reduction='mean'): + super(ListMLELoss, self).__init__() + self.reduction = reduction + + def forward(self, scores: torch.Tensor, labels: torch.Tensor) -> torch.Tensor: + # scores: (B, N) + # labels: (B, N) + + sorted_indices = torch.argsort(labels, dim=1, descending=True) + + sorted_scores = torch.gather(scores, 1, sorted_indices) + + max_val, _ = sorted_scores.max(dim=1, keepdim=True) + sorted_scores_stable = sorted_scores - max_val + + exp_scores = torch.exp(sorted_scores_stable) + + exp_sum_reverse = torch.cumsum(torch.flip(exp_scores, [1]), dim=1) + exp_sum_reverse = torch.flip(exp_sum_reverse, [1]) + + log_cumsum = torch.log(exp_sum_reverse + 1e-10) # eps + + # Restore scale: log(sum(e^(x-m))) = log(sum) + m + # Loss term_i = (log_cumsum_i + M) - sorted_score_i + loss_per_item = (log_cumsum + max_val) - sorted_scores + + loss = loss_per_item.sum(dim=1) + + if self.reduction == 'mean': + return loss.mean() + elif self.reduction == 'sum': + return loss.sum() + return loss + +class Model(nn.Module): + def __init__(self, reduction='none'): + super(Model, self).__init__() + self.loss_fn = ListMLELoss(reduction=reduction) + + def forward(self, scores, labels): + return self.loss_fn(scores, labels) + +def get_inputs(): + scores = torch.randn(SHAPE, dtype=torch.float32) + labels = torch.randint(0, 5, SHAPE, dtype=torch.float32) + labels += torch.rand(SHAPE, dtype=torch.float32) * 1e-3 + + return [scores.contiguous(), labels.contiguous()] + +def get_init_inputs(): + return ['none'] \ No newline at end of file diff --git a/S1/hli28146_#48/prompt.txt b/S1/hli28146_#48/prompt.txt new file mode 100644 index 00000000..9febd116 --- /dev/null +++ b/S1/hli28146_#48/prompt.txt @@ -0,0 +1,94 @@ +Write a custom CUDA kernel to optimize `ListMLE Loss`. + +Formula: L = Sum_{i=0}^{N-1} [ log( Sum_{k=i}^{N-1} exp(z_pi(k)) ) - z_pi(i) ] +Where `z` are predicted scores, and `pi` is the permutation that sorts the ground truth labels in descending order. +Essentially: Sort scores based on labels -> Compute LogSumExp of the suffix -> Subtract score -> Sum. + +Problem Analysis: +1. Sorting Bottleneck: The standard implementation requires `torch.argsort` on labels for every sample in the batch, which is computationally expensive and memory-intensive. +2. Memory Traffic: Calculating the suffix sum (denominator) involves multiple passes: exp, flip, cumsum, flip, log. + +Optimization Strategy: Fused Sort-Scan Kernel in Shared Memory + +Constraint: List size `N` is fixed to a power of 2 (e.g., 1024) for efficient Bitonic Sort. + +1. Block-per-Query: Assign one CUDA block to process one query (list of items). + +2. Shared Memory Staging: Load scores and labels into Shared Memory structures. + +3. Parallel Bitonic Sort: + - Sort the data in Shared Memory based on **labels** in descending order. + - Use `score` as the payload moved along with labels. + - This replaces the global `argsort` + `gather` pattern. + +4. Numerical Stability & Suffix Scan: + - Find the max score `M` in the sorted list for stability. + - Compute `exp(score - M)`. + - Perform a **Reverse Parallel Scan** (Suffix Sum) in Shared Memory to calculate `CumSumExp_i = sum_{k=i}^{N-1} exp(...)`. + +5. Fused Loss: + - `Loss_i = (log(CumSumExp_i) + M) - score_i`. + - Sum `Loss_i` across the block to get total loss for the query. + +Here's an example to show you the syntax of inline embedding custom CUDA operators in torch: The example given architecture is: + +```python +import torch +import torch.nn as nn + +BATCH_SIZE = 4096 +LIST_SIZE = 1024 +SHAPE = (BATCH_SIZE, LIST_SIZE) + +class ListMLELoss(nn.Module): + def __init__(self, reduction='mean'): + super(ListMLELoss, self).__init__() + self.reduction = reduction + + def forward(self, scores: torch.Tensor, labels: torch.Tensor) -> torch.Tensor: + # scores: (B, N) + # labels: (B, N) + + sorted_indices = torch.argsort(labels, dim=1, descending=True) + + sorted_scores = torch.gather(scores, 1, sorted_indices) + + max_val, _ = sorted_scores.max(dim=1, keepdim=True) + sorted_scores_stable = sorted_scores - max_val + + exp_scores = torch.exp(sorted_scores_stable) + + exp_sum_reverse = torch.cumsum(torch.flip(exp_scores, [1]), dim=1) + exp_sum_reverse = torch.flip(exp_sum_reverse, [1]) + + log_cumsum = torch.log(exp_sum_reverse + 1e-10) # eps + + # Restore scale: log(sum(e^(x-m))) = log(sum) + m + # Loss term_i = (log_cumsum_i + M) - sorted_score_i + loss_per_item = (log_cumsum + max_val) - sorted_scores + + loss = loss_per_item.sum(dim=1) + + if self.reduction == 'mean': + return loss.mean() + elif self.reduction == 'sum': + return loss.sum() + return loss + +class Model(nn.Module): + def __init__(self, reduction='none'): + super(Model, self).__init__() + self.loss_fn = ListMLELoss(reduction=reduction) + + def forward(self, scores, labels): + return self.loss_fn(scores, labels) + +def get_inputs(): + scores = torch.randn(SHAPE, dtype=torch.float32) + labels = torch.randint(0, 5, SHAPE, dtype=torch.float32) + labels += torch.rand(SHAPE, dtype=torch.float32) * 1e-3 + + return [scores.contiguous(), labels.contiguous()] + +def get_init_inputs(): + return ['none'] \ No newline at end of file diff --git a/S1/hli28146_#48/run_code.py b/S1/hli28146_#48/run_code.py new file mode 100644 index 00000000..a8ff02c9 --- /dev/null +++ b/S1/hli28146_#48/run_code.py @@ -0,0 +1,74 @@ +########################################################### +# 性能和精度验证程序 +########################################################### +import torch +import torch.nn as nn +import time +from mleloss_torch import Model,get_inputs,get_init_inputs +from mleloss_cuda import ModelNew + +def run_benchmark(): + # 检查 CUDA 是否可用 + if not torch.cuda.is_available(): + print("CUDA 不可用,请确保您有可用的 NVIDIA GPU 并已正确安装 PyTorch CUDA 版本。") + return + else: + device = torch.device("cuda") + + # 初始化模型 + init_inputs = get_init_inputs() + init_inputs = [ + x.cuda(device=device) if isinstance(x, torch.Tensor) else x for x in init_inputs + ] + inputs = get_inputs() + inputs = [ + x.cuda(device=device) if isinstance(x, torch.Tensor) else x for x in inputs + ] + + torch_model = Model(*init_inputs).cuda() + cuda_model = ModelNew(*init_inputs).cuda() + + torch_model.eval() + cuda_model.eval() + + print("-------------------- 精度对齐验证 --------------------") + with torch.no_grad(): + output_torch = torch_model( *inputs) + output_cuda = cuda_model(*inputs) + + precision_flag = torch.allclose(output_torch, output_cuda,rtol=1e-03) + if precision_flag: + print("✅ 精度对齐:两个模型的输出结果非常接近。") + else: + print("❌ 精度不一致!") + + print("\n-------------------- 性能加速比测试 --------------------") + num_iterations = 100 + + # PyTorch 模型计时 + torch.cuda.synchronize() + start_time = time.time() + for _ in range(num_iterations): + _ = torch_model(*inputs) + torch.cuda.synchronize() + torch_time = (time.time() - start_time) / num_iterations + + # 自定义 CUDA 内核计时 + torch.cuda.synchronize() + start_time = time.time() + for _ in range(num_iterations): + _ = cuda_model(*inputs) + torch.cuda.synchronize() + cuda_time = (time.time() - start_time) / num_iterations + + print(f"PyTorch torch.relu 平均执行时间: {torch_time:.6f} 秒") + print(f"自定义 CUDA 内核 平均执行时间: {cuda_time:.6f} 秒") + speedup = 0 + if cuda_time > 0: + speedup = torch_time / cuda_time + print(f"加速比 (Speedup): {speedup:.2f}x") + else: + print("CUDA 内核执行时间为0,无法计算加速比。") + return precision_flag,speedup +if __name__ == "__main__": + precision_flag,speedup = run_benchmark() \ No newline at end of file