From 56c684281092ced0c8e654c2a335b84e7371f757 Mon Sep 17 00:00:00 2001 From: jwu773 Date: Thu, 13 Aug 2026 08:48:07 +0800 Subject: [PATCH] jwu773 submit 1 --- README.md | 159 +++++++++++++++++++++++++++++ src/kernels.cu | 257 ++++++++++++++++++++++++++++++++++++++++++++++- src/kernels.maca | 28 ++++++ src/kernels.mu | 28 ++++++ 4 files changed, 471 insertions(+), 1 deletion(-) diff --git a/README.md b/README.md index ec69e0ee..879903d4 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,6 @@ # Learning-CUDA +<<<<<<< HEAD 本项目为 2026 年夏季 InfiniTensor 大模型与人工智能系统训练营 CUDA 方向专业阶段的作业与项目系统。 ## 项目结构 @@ -58,10 +59,78 @@ output[i, j] = input[i, j] * rsqrt(mean_square + eps) * weight[j] 输入和输出均按 row-major 方式展平存储。该函数需支持 `float` 和 `half` 两种类型。 +======= +本项目为 2025 年冬季 InfiniTensor 大模型与人工智能系统训练营 CUDA 方向专业阶段的作业与项目系统。 + +## 📁 项目结构 + +```text +learning-CUDA/ +├── Makefile +├── README +├── src +│ └── kernels.cu +└── tester + ├── tester.o + └── utils.h +``` + +## 环境配置 + +### > 英伟达(NVIDIA) + +- 如果你使用的是训练营所提供的服务器,遵照英伟达算力文档中的步骤配置好环境即可。 + +- 如果为本地或其他环境,请确保系统已安装以下工具: + + 1. **CUDA Toolkit**(版本11.0及以上): + - 验证安装:运行`nvcc --version`。 + - 安装:从[NVIDIA CUDA Toolkit下载页](https://developer.nvidia.com/cuda-downloads)获取。 + 2. **GNU Make**: + - 验证安装:运行`make --version`(大多数Linux/macOS已预装)。 + +### > 天数智芯(Iluvatar CoreX) + +- 如果你使用的是训练营所提供的服务器,遵照天数 BI-100 算力文档中的步骤配置好环境即可。 + +- 对于非训练营所提供的天数算力,请配置标准的天数 GPU 开放环境。**本次作业的配置不保证能在所有其他天数环境上无修改直接运行**。 + +### > 沐曦集成电路(Metax) + +- 如果你使用的是训练营所提供的服务器,遵照沐曦 (C500) 算力文档中的步骤配置好环境即可。 + +- 对于非训练营所提供的沐曦算力,请配置标准的沐曦 GPU 开放环境。**本次作业的配置不保证能在所有其他沐曦环境上无修改直接运行**。 + +### > 摩尔线程(Moore Threads) + +- 如果你使用的是训练营所提供的服务器,请先遵照摩尔 (S5000) 算力文档中的步骤配置环境。 + + 在此基础上,确保在 `.bashrc` 中添加以下环境变量: + + ```bash + export MUSA_ROOT=/usr/local/musa + export PATH="$MUSA_ROOT/bin:$PATH" + export LD_LIBRARY_PATH="$MUSA_ROOT/lib:$LD_LIBRARY_PATH" + export CPLUS_INCLUDE_PATH=/usr/include/c++/11:/usr/include/x86_64-linux-gnu/c++/11 + ``` + +- 对于非训练营所提供的摩尔算力,请配置标准的摩尔 GPU 开放环境。**本次作业的配置不保证能在所有其他摩尔环境上无修改直接运行**。 + + +## 🧠 作业 + +作业一共有两题。需实现 `src/kernels.cu` 中给定的 **2 个 CUDA 函数** 。 + +1. **trace** + +实现 CUDA 的 trace 函数。给定一个逻辑上 2D 的输入矩阵,返回该矩阵的迹。该函数需支持 `int` 和 `float` 两种类型的输入。具体边界处理和一些条件可见文件中的注释。 + +>>>>>>> fcdc822 (feat.: add 2025_winter assignment) 2. **flashAttention** 实现 Flash Attention 算子。需支持 causal masking 和 GQA。具体行为与 [torch.nn.functional.scaled_dot_product_attention](https://docs.pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html) 保持一致。接口未提供的参数所代表的功能无需支持和实现。具体参数要求请参考文件中的注释。该函数需支持 `float` 和 `half` 两种类型。 +<<<<<<< HEAD ### 国产平台适配 在完成英伟达的基础上,可以将实现适配至天数、沐曦和/或摩尔这三款 GPU 平台上。 @@ -144,3 +213,93 @@ make PLATFORM=moore ## 有疑问? 可以在群里直接询问助教。 +======= +### 注意事项 + +1. **禁止抄袭与舞弊**,包括抄袭其他学员的代码和开源实现。可以讨论和参考思路,但禁止直接看/抄代码。一经发现,成绩作废并失去进入项目阶段和后续实习与推荐等资格; +2. 两个题目都**禁止使用任何库函数**来直接实现关键功能; +3. 主要计算均需在 GPU 上实现;如有一些信息和程序准备性质的(例如元信息计算/转换、资源准备等)则可以在 CPU/Host 上进行; +4. 代码风格不限,但需保持一致; +5. 需进行**适当**的代码注释解释重要部分; + +### 提交方式 +在网站 [InfiniTensor 开源社区](https://www.infinitensor.com/camp/winter2025/homework) 上提交 GitHub 链接,以最新提交为准。 + +## 🛠️ 编译与运行 + + 代码编译与运行可以使用提供的 `Makefile` 十分简便的实现。 + +### 构建与运行指令 + +使用 `Makefile` 简化构建流程,以下命令需在**项目根目录**(即 `Makefile` 所在的目录)执行: + +#### 1. 默认:构建并运行测试(非 verbose 模式) + +- 直接在命令行使用 `make` 指令编译代码并执行测试,输出简洁结果。 + +#### 2. 构建并运行测试(verbose 模式) + +- 直接在命令行使用 `make VERBOSE=true` 指令编译代码并执行测试,输出包括执行时间等更多信息在内的结果。 + +#### 3. 选择性测试算子 + +如果只想调试/测试某个算子,可以通过设置环境变量的方式来实现。比如只想测试第一题的 trace,则可以: + +1. 如果只是临时跳过,可以直接在命令行使用 `SKIP_ATTENTION=1 make` 编译代码并**跳过第二题,只测试第一题**。 + +2. 如果较长时间都想跳过,可以一开始使用一次 `export SKIP_ATTENTION=1`, 随后在同一个命令行中照常使用 `make` 相关命令。如果想撤销,则可以使用 `export SKIP_ATTENTION=0` 或 `unset SKIP_ATTENTION`。 + + +#### 4. 选择编译平台 + +可以通过在命令行使用 `make PLATFORM=` 指令来指定编译平台。**默认的编译平台为英伟达平台**,即如果不指定 `PLATFORM` 直接 `make`,则是编译英伟达平台。具体平台选项: + +1. 编译并在英伟达平台运行:`make` 或 `make PLATFORM=nvidia`; + +2. 编译并在天数平台运行:`make PLATFORM=iluvatar`; + +3. 编译并在沐曦平台运行:`make PLATFORM=metax`; + +4. 编译并在摩尔平台运行:`make PLATFORM=moore`; + + +**以上提及的编译选项与环境变量均可根据需求组合。例如:`SKIP_TRACE=1 make PLATFORM=nvidia VERBOSE=true`** + + +#### 环境变量: +1. `SKIP_TRACE`: 跳过第一题的 trace 测试。 + +2. `SKIP_ATTENTION`: 跳过第二题的 Flash Attention 测试。 + + +## 📊 评分规则 + +本次作业的评分标准如下: + +1. **正确性优先** + - 所有提交首先以正确性为前提,需在提供的测试用例中正确输出结果; + - 正确性提供基础分:每通过一个测例,获得相应的基础得分; + - 未通过的测试用例,不计入性能排名; + - 不符合**注意事项**中要求的,不得分。 + +2. **性能加分** + - 在正确性的基础上,会对各实现的性能进行排名; + - 性能越优,获得的额外分数越多; + - **性能评判将在提供的服务器上进行**,因此请在服务器上进行性能评估。 + +3. **平台适配加分** + - 每道题在英伟达上测例正确的基础上,每多适配一个国产平台可以获得固定得分乘算系数(20%); + - 每个平台适配完成的标准为该题在该平台上可以通过全部测例。全部通过则获得该平台的 20% 加成,无法全部通过则无法获得加成(0%); + - 题目分开计算,即只有在该平台上通过全部测例的题目可以获得该题目部分的乘算系数加成; + - 国产平台不进行性能测试,故不参与性能得分计算(与性能加分正交)。 + +4. **最终成绩** + - 总体得分由「通过的测试用例数量」、「性能排名加分」和「平台适配加分」共同决定。 + - 各测试用例的分数相加,形成最终成绩。 + +## 📬 有疑问? + +可以在群里直接询问助教! + +Good luck and happy coding! 🚀 +>>>>>>> fcdc822 (feat.: add 2025_winter assignment) diff --git a/src/kernels.cu b/src/kernels.cu index 2cc53e7e..bf46f699 100644 --- a/src/kernels.cu +++ b/src/kernels.cu @@ -3,6 +3,13 @@ #include "../tester/utils.h" +#include +#include +#include +#include +#include +#define tileHeight 4 + /** * @brief Computes RMSNorm over the last dimension of a 2D tensor. * @@ -21,13 +28,121 @@ * @param[in] hidden_dim Size of the normalized dimension. * @param[in] eps Numerical stability epsilon. */ + +template +__device__ __forceinline__ T castFloatToT(float v) { + return static_cast(v); +} + +// float 特化 +template <> +__device__ __forceinline__ float castFloatToT(float v) { + return v; +} + +// half 特化 +template <> +__device__ __forceinline__ half castFloatToT(float v) { + return __float2half(v); +} + +// 转float模板 +template +__device__ __forceinline__ float castToFloat(T v); + +template <> +__device__ __forceinline__ float castToFloat(float v) { return v; } + +template <> +__device__ __forceinline__ float castToFloat(half v) { return __half2float(v); } + + +extern __shared__ char smem_raw[]; + template + __global__ void RMSNormKernel(T* d_input, T* d_weight, T* d_output, size_t hidden_dim, size_t iterRounds, float eps, float* d_res){ + + T *smem = reinterpret_cast(smem_raw); + //load data from global mem to share mem + int smemIdx = threadIdx.x; + int gmemIdx = blockIdx.x * hidden_dim + smemIdx; + for(int r = 0; r < iterRounds - 1; r++){ + smem[smemIdx] = d_input[gmemIdx]; + smemIdx += blockDim.x; + gmemIdx += blockDim.x; + } + if(smemIdx < hidden_dim){ + smem[smemIdx] = d_input[gmemIdx]; + } + __syncthreads(); + + //let the first thread compute the sum, mean, rsqrt + if(threadIdx.x == 0){ + float tempRes = 0; + for(int i = 0; i < hidden_dim; i++){ + float v = castToFloat(smem[i]); + tempRes += v * v; + } + tempRes /= (hidden_dim * 1.0f); + tempRes = rsqrt(tempRes + eps); + *d_res = tempRes; + } + __syncthreads(); + + //让线程块的所有线程进行权重乘法,并写回全局内存 + float res = *d_res; + smemIdx = threadIdx.x; + gmemIdx = blockIdx.x * hidden_dim + smemIdx; + + for(int r = 0; r < iterRounds - 1; r++){ + float v = castToFloat(smem[smemIdx]) * res * castToFloat(d_weight[smemIdx]); + d_output[gmemIdx] = castFloatToT(v); + smemIdx += blockDim.x; + gmemIdx += blockDim.x; + } + if(smemIdx < hidden_dim){ + float v = castToFloat(smem[smemIdx]) * res * castToFloat(d_weight[smemIdx]); + d_output[gmemIdx] = castFloatToT(v); + } +} + + + template void rmsNorm(const std::vector& h_input, const std::vector& h_weight, std::vector& h_output, size_t rows, size_t hidden_dim, float eps) { // TODO: Implement the rmsNorm function + + T* d_input = nullptr; + cudaMalloc(&d_input, rows * hidden_dim * sizeof(T)); + cudaMemcpyAsync(d_input, h_input.data(), rows * hidden_dim * sizeof(T), cudaMemcpyHostToDevice); + + int blksPerGrid = (int)rows; + int threadsPerBlk = std::min(1024, (int)hidden_dim); + + T* d_weight = nullptr; + cudaMalloc(&d_weight, hidden_dim * sizeof(T)); + cudaMemcpyAsync(d_weight, h_weight.data(), hidden_dim * sizeof(T), cudaMemcpyHostToDevice); + + T* d_output = nullptr; + cudaMalloc(&d_output, rows * hidden_dim * sizeof(T)); + + float* d_res = nullptr; + cudaMalloc(&d_res, sizeof(float)); + + size_t iterRounds = (hidden_dim - 1) / threadsPerBlk + 1; + + RMSNormKernel<<>>(d_input, d_weight, d_output, hidden_dim, iterRounds, eps, d_res); + cudaDeviceSynchronize(); + cudaMemcpyAsync(h_output.data(), d_output, rows * hidden_dim * sizeof(T), cudaMemcpyDeviceToHost); + + cudaFree(d_input); + cudaFree(d_weight); + cudaFree(d_output); } + + /** * @brief Computes flash attention for given query, key, and value tensors. * @@ -44,14 +159,154 @@ void rmsNorm(const std::vector& h_input, const std::vector& h_weight, * @param[in] head_dim Dimension size of each attention head * @param[in] is_causal Whether to apply causal masking */ + +template +__global__ void flashAttentionKernel( + T* q, T* k, T* v, T* o, int batch_size, + int target_seq_len, int src_seq_len, int query_heads, int kv_heads, int head_dim, bool is_causal, bool isHalfType) { + int b_idx = blockIdx.z; + int qh_head_idx = blockIdx.y; + int q_row_idx = blockIdx.x * tileHeight + threadIdx.y; + + int kv_head_idx = qh_head_idx * kv_heads / query_heads; + //int valid_seq_len = is_causal? min(q_row_idx + 1, src_seq_len) : src_seq_len; + int valid_seq_len = src_seq_len; + + //assume head_dim <= 128 + __shared__ T q_patch[tileHeight][128]; + __shared__ T k_patch[tileHeight][128]; + __shared__ float intermediate[tileHeight][2048]; + __shared__ float softmaxSum[tileHeight]; + + //load Q tile of current block to smem (smem:share memory) + if(q_row_idx < target_seq_len && threadIdx.x < head_dim){ + int q_offset = b_idx * (target_seq_len * query_heads * head_dim) + \ + q_row_idx * (query_heads * head_dim) + qh_head_idx * head_dim + threadIdx.x; + q_patch[threadIdx.y][threadIdx.x] = q[q_offset]; + } + + //load K to smem tile by tile + int valid_row_curTile; + for(int i = 0; i < (valid_seq_len - 1) / tileHeight + 1; i++){ + int k_row_idx = i * tileHeight + threadIdx.y; + + //load tile i to smem + if(k_row_idx < valid_seq_len && threadIdx.x < head_dim){ + int k_offset = ((b_idx * src_seq_len + k_row_idx) * kv_heads + kv_head_idx) * head_dim + threadIdx.x; + k_patch[threadIdx.y][threadIdx.x] = k[k_offset]; + } + __syncthreads(); + + //dot product: q patch * k patch + valid_row_curTile = (i == (valid_seq_len - 1) / tileHeight)? valid_seq_len % tileHeight : tileHeight; + if(valid_row_curTile == 0) + valid_row_curTile = tileHeight; + + if(q_row_idx < target_seq_len && threadIdx.x < valid_row_curTile){ + if(is_causal && q_row_idx < threadIdx.x + i * tileHeight) + intermediate[threadIdx.y][threadIdx.x + i * tileHeight] = -INFINITY; + else{ + float sum = 0.0f; + for(int j = 0; j < head_dim; j++) + sum += castToFloat(q_patch[threadIdx.y][j]) * castToFloat(k_patch[threadIdx.x][j]); + + intermediate[threadIdx.y][threadIdx.x + i * tileHeight] = sum; + } + } + } + __syncthreads(); + + valid_row_curTile = (blockIdx.x == (target_seq_len - 1) / tileHeight)? target_seq_len % tileHeight : tileHeight; + if(valid_row_curTile == 0) + valid_row_curTile = tileHeight; + + //find max value + if(threadIdx.y == 0 && threadIdx.x < valid_row_curTile){ + float maxVal = -INFINITY; + for(int i = 0; i < valid_seq_len; i++){ + if(maxVal < intermediate[threadIdx.x][i]) + maxVal = intermediate[threadIdx.x][i]; + } + //scale, exp + float softmax_sum = 0.0f; + float scale_factor = 1.0f / sqrt(static_cast(head_dim) + 1e-8f); + for(int i = 0; i < valid_seq_len; i++){ + if(intermediate[threadIdx.x][i] > -INFINITY){ + intermediate[threadIdx.x][i] = exp((intermediate[threadIdx.x][i] - maxVal) * scale_factor); + softmax_sum += intermediate[threadIdx.x][i]; + } + else{ + intermediate[threadIdx.x][i] = 0.0f; + } + + } + softmaxSum[threadIdx.x] = softmax_sum; + } + __syncthreads(); + + //dot product intermediate with V + if(q_row_idx < target_seq_len && threadIdx.x < head_dim){ + float dotProduct = 0.0f; + for(int i = 0; i < valid_seq_len; i++){ + int v_offset = ((b_idx * src_seq_len + i) * kv_heads + kv_head_idx) * head_dim + threadIdx.x; + dotProduct += intermediate[threadIdx.y][i] * castToFloat(v[v_offset]); + } + //softmax + if(softmaxSum[threadIdx.y] > 0.0f) + dotProduct /= (softmaxSum[threadIdx.y] + 1e-12f); + //store it to o + int o_offset = b_idx * (target_seq_len * query_heads * head_dim) + \ + q_row_idx * (query_heads * head_dim) + qh_head_idx * head_dim + threadIdx.x; + o[o_offset] = castFloatToT(dotProduct); + } +} + + + template void flashAttention(const std::vector& h_q, const std::vector& h_k, const std::vector& h_v, std::vector& h_o, int batch_size, int target_seq_len, int src_seq_len, int query_heads, int kv_heads, int head_dim, bool is_causal) { // TODO: Implement the flash attention function + if (query_heads % kv_heads != 0) { + return; + } + + T *d_q = nullptr, *d_k = nullptr, *d_v = nullptr, *d_o = nullptr; + cudaMalloc(&d_q, batch_size * target_seq_len * query_heads * head_dim * sizeof(T)); + cudaMalloc(&d_k, batch_size * src_seq_len * kv_heads * head_dim * sizeof(T)); + cudaMalloc(&d_v, batch_size * src_seq_len * kv_heads * head_dim * sizeof(T)); + cudaMalloc(&d_o, batch_size * target_seq_len * query_heads * head_dim * sizeof(T)); + + cudaMemcpy(d_q, h_q.data(), batch_size * target_seq_len * query_heads * head_dim * sizeof(T), cudaMemcpyHostToDevice); + cudaMemcpy(d_k, h_k.data(), batch_size * src_seq_len * kv_heads * head_dim * sizeof(T), cudaMemcpyHostToDevice); + cudaMemcpy(d_v, h_v.data(), batch_size * src_seq_len * kv_heads * head_dim * sizeof(T), cudaMemcpyHostToDevice); + + dim3 gridDim((target_seq_len + tileHeight - 1) / tileHeight, query_heads, batch_size); + dim3 blockDim(max(tileHeight,head_dim), tileHeight); + + bool isHalfType = false; + if constexpr (std::is_same::value){ + isHalfType = true; + } + flashAttentionKernel<<>>( + d_q, d_k, d_v, d_o, + batch_size, target_seq_len, src_seq_len, + query_heads, kv_heads, head_dim, + is_causal, isHalfType + ); + + cudaDeviceSynchronize(); + cudaMemcpy(h_o.data(), d_o, batch_size * target_seq_len * query_heads * head_dim * sizeof(T), cudaMemcpyDeviceToHost); + + cudaFree(d_q); + cudaFree(d_k); + cudaFree(d_v); + cudaFree(d_o); } + // ********************************************************************* // Explicit Template Instantiations (REQUIRED FOR LINKING WITH TESTER.O) // DO NOT MODIFY THIS SECTION @@ -65,4 +320,4 @@ template void flashAttention(const std::vector&, const std::vector int, int, int, int, int, int, bool); template void flashAttention(const std::vector&, const std::vector&, const std::vector&, std::vector&, - int, int, int, int, int, int, bool); + int, int, int, int, int, int, bool); \ No newline at end of file diff --git a/src/kernels.maca b/src/kernels.maca index 4c320f21..d595996f 100644 --- a/src/kernels.maca +++ b/src/kernels.maca @@ -4,6 +4,7 @@ #include "../tester/utils.h" /** +<<<<<<< HEAD * @brief Computes RMSNorm over the last dimension of a 2D tensor. * * The input is a row-major matrix with shape [rows, hidden_dim]. For each row @@ -26,6 +27,25 @@ void rmsNorm(const std::vector& h_input, const std::vector& h_weight, std::vector& h_output, size_t rows, size_t hidden_dim, float eps) { // TODO: Implement the rmsNorm function +======= + * @brief Computes the trace of a matrix. + * + * The trace of a matrix is defined as the sum of its diagonal elements. + * This function expects a flattened row-major matrix stored in a + * std::vector. If the matrix is not square, the trace will sum up + * elements along the main diagonal up to the smaller of rows or cols. + * + * @tparam T The numeric type of matrix elements (e.g., float, int). + * @param h_input A flattened matrix of size rows * cols. + * @param rows Number of rows in the matrix. + * @param cols Number of columns in the matrix. + * @return The trace (sum of diagonal values) of the matrix. + */ +template +T trace(const std::vector& h_input, size_t rows, size_t cols) { + // TODO: Implement the trace function + return T(-1); +>>>>>>> fcdc822 (feat.: add 2025_winter assignment) } /** @@ -49,17 +69,25 @@ void flashAttention(const std::vector& h_q, const std::vector& h_k, const std::vector& h_v, std::vector& h_o, int batch_size, int target_seq_len, int src_seq_len, int query_heads, int kv_heads, int head_dim, bool is_causal) { +<<<<<<< HEAD // TODO: Implement the flash attention function +======= +>>>>>>> fcdc822 (feat.: add 2025_winter assignment) } // ********************************************************************* // Explicit Template Instantiations (REQUIRED FOR LINKING WITH TESTER.O) // DO NOT MODIFY THIS SECTION // ********************************************************************* +<<<<<<< HEAD template void rmsNorm(const std::vector&, const std::vector&, std::vector&, size_t, size_t, float); template void rmsNorm(const std::vector&, const std::vector&, std::vector&, size_t, size_t, float); +======= +template int trace(const std::vector&, size_t, size_t); +template float trace(const std::vector&, size_t, size_t); +>>>>>>> fcdc822 (feat.: add 2025_winter assignment) template void flashAttention(const std::vector&, const std::vector&, const std::vector&, std::vector&, int, int, int, int, int, int, bool); diff --git a/src/kernels.mu b/src/kernels.mu index 1ce371eb..355d7293 100644 --- a/src/kernels.mu +++ b/src/kernels.mu @@ -4,6 +4,7 @@ #include "../tester/utils.h" /** +<<<<<<< HEAD * @brief Computes RMSNorm over the last dimension of a 2D tensor. * * The input is a row-major matrix with shape [rows, hidden_dim]. For each row @@ -26,6 +27,25 @@ void rmsNorm(const std::vector& h_input, const std::vector& h_weight, std::vector& h_output, size_t rows, size_t hidden_dim, float eps) { // TODO: Implement the rmsNorm function +======= + * @brief Computes the trace of a matrix. + * + * The trace of a matrix is defined as the sum of its diagonal elements. + * This function expects a flattened row-major matrix stored in a + * std::vector. If the matrix is not square, the trace will sum up + * elements along the main diagonal up to the smaller of rows or cols. + * + * @tparam T The numeric type of matrix elements (e.g., float, int). + * @param h_input A flattened matrix of size rows * cols. + * @param rows Number of rows in the matrix. + * @param cols Number of columns in the matrix. + * @return The trace (sum of diagonal values) of the matrix. + */ +template +T trace(const std::vector& h_input, size_t rows, size_t cols) { + // TODO: Implement the trace function + return T(-1); +>>>>>>> fcdc822 (feat.: add 2025_winter assignment) } /** @@ -49,17 +69,25 @@ void flashAttention(const std::vector& h_q, const std::vector& h_k, const std::vector& h_v, std::vector& h_o, int batch_size, int target_seq_len, int src_seq_len, int query_heads, int kv_heads, int head_dim, bool is_causal) { +<<<<<<< HEAD // TODO: Implement the flash attention function +======= +>>>>>>> fcdc822 (feat.: add 2025_winter assignment) } // ********************************************************************* // Explicit Template Instantiations (REQUIRED FOR LINKING WITH TESTER.O) // DO NOT MODIFY THIS SECTION // ********************************************************************* +<<<<<<< HEAD template void rmsNorm(const std::vector&, const std::vector&, std::vector&, size_t, size_t, float); template void rmsNorm(const std::vector&, const std::vector&, std::vector&, size_t, size_t, float); +======= +template int trace(const std::vector&, size_t, size_t); +template float trace(const std::vector&, size_t, size_t); +>>>>>>> fcdc822 (feat.: add 2025_winter assignment) template void flashAttention(const std::vector&, const std::vector&, const std::vector&, std::vector&, int, int, int, int, int, int, bool);