@@ -12,13 +12,14 @@ namespace llm {
1212using namespace cute ;
1313
1414// AttentionTile specialization for AttentionParams
15- template <typename TileShape, // (BLK_M, BLK_N, BLK_K)
16- typename Element, // Element type
17- typename StrideQ, // (Q, D, ((KH, G), B))
18- typename StrideK, // (K, D, ((KH, _0), B))
19- typename StrideV, // (V, D, ((KH, _0), B))
20- typename StrideO, // (Q, Q, ((KH, G), B))
21- bool kLocal >
15+ template <typename TileShape, // (BLK_M, BLK_N, BLK_K)
16+ typename BlocKCoord, // (m_block_idx, ((kv_head_idx, _0), batch_idx))
17+ typename Element, // Element type
18+ typename StrideQ, // (Q, D, ((KH, G), B))
19+ typename StrideK, // (K, D, ((KH, _0), B))
20+ typename StrideV, // (V, D, ((KH, _0), B))
21+ typename StrideO // (Q, Q, ((KH, G), B))
22+ >
2223struct FmhaBlock {
2324 // Host side parameters
2425 struct Arguments {
@@ -31,8 +32,6 @@ struct FmhaBlock {
3132 StrideK k_stride;
3233 StrideV v_stride;
3334 StrideO o_stride;
34-
35- int sliding_window = -1 ; // -1 means no sliding window
3635 };
3736
3837 // Device side parameters
@@ -47,8 +46,6 @@ struct FmhaBlock {
4746 StrideV v_stride;
4847 StrideO o_stride;
4948
50- int sliding_window;
51-
5249 // Parameters from problem shape
5350 int batch_size;
5451 int q_len;
@@ -59,18 +56,6 @@ struct FmhaBlock {
5956 FastDivmod group_size;
6057 };
6158
62- // using StrideK = ...;
63-
64- // using TMA_K = decltype(make_tma_copy(
65- // GmemTiledCopy{}, // TMA_COPY
66- // make_tensor(static_cast<InternalElementA const*>(nullptr),
67- // repeat_like(StrideK{}, int32_t(0)), StrideK{}),
68- // SmemLayoutK{}(_,_,_0{})));
69-
70- // Tensor tensor_k = make_tensor(ptr_k, make_layout(make_shape(M,K,L),
71- // args.stride_k)); auto tma_load_k = make_tma_copy(SM90_TMA_LOAD{},
72- // gtensor_k, SmemLayoutK{}(_,_,_0{}));
73-
7459 template <class ProblemShape >
7560 static Params to_underlying_arguments (const ProblemShape& problem_shape,
7661 const Arguments& args,
@@ -93,7 +78,6 @@ struct FmhaBlock {
9378 .k_stride = args.k_stride ,
9479 .v_stride = args.v_stride ,
9580 .o_stride = args.o_stride ,
96- .sliding_window = args.sliding_window ,
9781 .batch_size = batch_size,
9882 .q_len = q_len,
9983 .kv_len = kv_len,
@@ -113,9 +97,6 @@ struct FmhaBlock {
11397
11498 // hold a reference to the parameters and block coordination
11599 const Params& params_;
116- // TODO: pass in as parameter is better
117- // (m_block_idx, ((kv_head_idx, _0), batch_idx))
118- using BlocKCoord = Coord<int , Coord<Coord<int , _0>, int >>;
119100 const BlocKCoord& blk_coord_;
120101
121102 // derived parameters to avoid recomputation
@@ -161,9 +142,9 @@ struct FmhaBlock {
161142 // returns (m_block_idx, ((kv_head_idx, _0), batch_idx))
162143 CUTE_HOST_DEVICE const auto & get_block_coord () const { return blk_coord_; }
163144
164- // TODO: pass in kLocal and sliding_window instead
165145 // returns kv block range: (n_block_min, n_block_max]
166- CUTE_HOST_DEVICE auto get_kv_blocks () const {
146+ template <bool kLocal >
147+ CUTE_HOST_DEVICE auto get_kv_blocks (int sliding_window) const {
167148 static constexpr int kBlockM = get<0 >(TileShape{});
168149 static constexpr int kBlockN = get<1 >(TileShape{});
169150
@@ -176,7 +157,7 @@ struct FmhaBlock {
176157 const int n_block_max = cute::ceil_div (kv_idx_max, kBlockN );
177158
178159 if constexpr (kLocal ) {
179- const int kv_idx_min = std::max (0 , diagonal - params_. sliding_window );
160+ const int kv_idx_min = std::max (0 , diagonal - sliding_window);
180161 const int n_block_min = kv_idx_min / kBlockN ;
181162 return make_tuple (n_block_min, n_block_max);
182163 } else {
@@ -300,31 +281,44 @@ struct FmhaBlock {
300281 }
301282
302283 // functions for tma load
284+
285+ // using StrideK = ...;
286+
287+ // using TMA_K = decltype(make_tma_copy(
288+ // GmemTiledCopy{}, // TMA_COPY
289+ // make_tensor(static_cast<InternalElementA const*>(nullptr),
290+ // repeat_like(StrideK{}, int32_t(0)), StrideK{}),
291+ // SmemLayoutK{}(_,_,_0{})));
292+
293+ // Tensor tensor_k = make_tensor(ptr_k, make_layout(make_shape(M,K,L),
294+ // args.stride_k)); auto tma_load_k = make_tma_copy(SM90_TMA_LOAD{},
295+ // gtensor_k, SmemLayoutK{}(_,_,_0{}));
296+
303297 // returns kv tma tile: (BLK_N, BLK_K, n) => (1@0, 1@1, 1@2)
304- template <class TMA_K , class TMA_V >
305- CUTE_HOST_DEVICE auto get_kv_tma_tile (TMA_K tma_k, TMA_V tma_v) const {
306- // 1: make_gather_tma_tensor()
307- // tma_tensor = (seq, dim, kv_head) => (1@0, 1@1, 1@2)
308- // 2: partition into tiles
309- // tma_tile = (BLK_N, BLK_K, n) => (1@0, 1@1, 1@2)
310-
311- // (Q, D, (B, H))
312- // Tensor mQ_qdl_p = tma_k.get_tma_tensor(select<0,2,3>(problem_shape));
313- // Tensor mQ_qdl = domain_offset(make_coord(q_offs_0, _0{}, make_coord(_0{},
314- // q_offs_2_1)), mQ_qdl_p); (BLK_N, BLK_K, m, k, (b)) Tensor gQ_qdl =
315- // local_tile(mQ_qdl, TileShapeQK{}, make_coord(_, _, _), Step<_1, X,
316- // _1>{});
317-
318- // outside in caller part
319- // (BLK_N, BLK_K, n) => (TMA,TMA_M,TMA_N, n)
320- // auto cta_tma = tma.get_slice(Int<0>{}); // CTA slice
321- // (TMA,TMA_M,TMA_N,REST_M,REST_N)
322- // Tensor tAgA_x = cta_tma.partition_S(gA);
323- // (TMA,TMA_M,TMA_N)
324- // Tensor tAsA_x = cta_tma.partition_D(sA);
325-
326- return ;
327- }
298+ // template <class TMA_K, class TMA_V>
299+ // CUTE_HOST_DEVICE auto get_kv_tma_tile(TMA_K tma_k, TMA_V tma_v) const {
300+ // 1: make_gather_tma_tensor()
301+ // tma_tensor = (seq, dim, kv_head) => (1@0, 1@1, 1@2)
302+ // 2: partition into tiles
303+ // tma_tile = (BLK_N, BLK_K, n) => (1@0, 1@1, 1@2)
304+
305+ // (Q, D, (B, H))
306+ // Tensor mQ_qdl_p = tma_k.get_tma_tensor(select<0,2,3>(problem_shape));
307+ // Tensor mQ_qdl = domain_offset(make_coord(q_offs_0, _0{}, make_coord(_0{},
308+ // q_offs_2_1)), mQ_qdl_p); (BLK_N, BLK_K, m, k, (b)) Tensor gQ_qdl =
309+ // local_tile(mQ_qdl, TileShapeQK{}, make_coord(_, _, _), Step<_1, X,
310+ // _1>{});
311+
312+ // outside in caller part
313+ // (BLK_N, BLK_K, n) => (TMA,TMA_M,TMA_N, n)
314+ // auto cta_tma = tma.get_slice(Int<0>{}); // CTA slice
315+ // (TMA,TMA_M,TMA_N,REST_M,REST_N)
316+ // Tensor tAgA_x = cta_tma.partition_S(gA);
317+ // (TMA,TMA_M,TMA_N)
318+ // Tensor tAsA_x = cta_tma.partition_D(sA);
319+
320+ // return;
321+ // }
328322};
329323
330324} // namespace llm
0 commit comments