refactor: template AttentionParams, rename .cuh to .h

- Convert AttentionParams to a template struct supporting arbitrary types
- Rename attn_common.cuh -> attn_common.h (no CUDA-specific code remains)
- Include standard headers explicitly in each .cuh instead of via attn_common.cuh
- Allow .h files in csrc/ via .gitignore
This commit is contained in:
2026-07-11 11:03:14 +08:00
parent b8b439b713
commit 8a8550184f
13 changed files with 80 additions and 64 deletions
+6 -5
View File
@@ -1,7 +1,8 @@
#pragma once
#include <torch/extension.h>
#include "attn_common.cuh"
#include "attn_common.h"
template<typename T>
inline void attn_pack_params(
torch::Tensor q,
torch::Tensor k,
@@ -10,7 +11,7 @@ inline void attn_pack_params(
bool is_causal,
int64_t causal_offset,
c10::optional<double> scale,
AttentionParams& p
AttentionParams<T>& p
) {
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
TORCH_CHECK(q.dtype() == torch::kBFloat16);
@@ -27,9 +28,9 @@ inline void attn_pack_params(
p.is_causal = is_causal ? 1 : 0;
p.causal_offset = (int)causal_offset;
p.scale = scale.has_value() ? (float)scale.value() : 1.0f / sqrtf((float)p.head_dim);
p.q = (const bf16*)q.data_ptr();
p.k = (const bf16*)k.data_ptr();
p.v = (const bf16*)v.data_ptr();
p.q = (const T*)q.data_ptr();
p.k = (const T*)k.data_ptr();
p.v = (const T*)v.data_ptr();
if (p.use_mask) {
TORCH_CHECK(mask.value().dtype() == torch::kBool);
TORCH_CHECK(mask.value().dim() == 2);