File size: 8,582 Bytes
b9ecbf8
 
 
 
 
 
 
 
 
 
5974ffe
b9ecbf8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5974ffe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b9ecbf8
 
 
 
5974ffe
 
 
b9ecbf8
5974ffe
 
 
 
 
 
b9ecbf8
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
#include <torch/all.h>
#include <torch/library.h>

#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>

#include <limits>

#include "attention_mha_masked.cuh"
#include "attention_seqused_fused.cuh"
#include "registration.h"

namespace {

int checked_int(int64_t value, const char* name) {
  TORCH_CHECK(value > 0 && value <= std::numeric_limits<int>::max(),
              name, " must fit in a positive int");
  return static_cast<int>(value);
}

void check_qkv(torch::Tensor const& tensor, const char* name,
               c10::ScalarType dtype) {
  TORCH_CHECK(tensor.is_cuda(), name, " must be CUDA");
  TORCH_CHECK(tensor.scalar_type() == dtype, name, " has the wrong dtype");
  TORCH_CHECK(tensor.dim() == 3, name, " must have shape (S, H, D)");
  TORCH_CHECK(tensor.stride(2) == 1 && tensor.stride(1) == tensor.size(2),
              name, " must be contiguous within each token");
}

void masked_mha_forward_static(
    torch::Tensor const& q, torch::Tensor const& k, torch::Tensor const& v,
    torch::Tensor& logits, torch::Tensor& out, double scale) {
  TORCH_CHECK(q.scalar_type() == torch::kFloat16 ||
                  q.scalar_type() == torch::kBFloat16,
              "q must be FP16 or BF16");
  check_qkv(q, "q", q.scalar_type());
  check_qkv(k, "k", q.scalar_type());
  check_qkv(v, "v", q.scalar_type());
  TORCH_CHECK(q.size(1) == k.size(1) && q.size(1) == v.size(1) &&
                  q.size(2) == k.size(2) && q.size(2) == v.size(2) &&
                  k.size(0) == v.size(0),
              "q/k/v head shapes must match");
  TORCH_CHECK(q.get_device() == k.get_device() &&
                  q.get_device() == v.get_device(),
              "q/k/v must be on the same device");
  TORCH_CHECK(out.is_cuda() && out.is_contiguous() &&
                  out.scalar_type() == q.scalar_type() &&
                  out.sizes() == q.sizes(),
              "out must be contiguous and match q");
  TORCH_CHECK(logits.is_cuda() && logits.scalar_type() == q.scalar_type() &&
                  logits.dim() == 3 && logits.size(0) == q.size(1) &&
                  logits.size(1) == q.size(0) &&
                  logits.size(2) >= k.size(0) && logits.stride(2) == 1,
              "logits must have shape (H, S_q, stride >= S_kv)");
  TORCH_CHECK(logits.get_device() == q.get_device() &&
                  out.get_device() == q.get_device(),
              "outputs must be on the q device");
  TORCH_CHECK(logits.stride(1) == logits.size(2) &&
                  logits.stride(0) == logits.size(1) * logits.size(2),
              "logits must use a dense padded row stride");

  c10::cuda::CUDAGuard guard(q.device());
  auto stream = at::cuda::getCurrentCUDAStream(q.get_device()).stream();
  auto handle = at::cuda::getCurrentCUDABlasHandle();
  const int sq = checked_int(q.size(0), "S_q");
  const int sk = checked_int(k.size(0), "S_kv");
  const int heads = checked_int(q.size(1), "heads");
  const int dim = checked_int(q.size(2), "head_dim");

  if (q.scalar_type() == torch::kFloat16) {
    TORCH_CHECK(q.stride(0) == heads * dim &&
                    k.stride(0) == heads * dim &&
                    v.stride(0) == heads * dim,
                "FP16 q/k/v must be contiguous across tokens");
    attention_mha_fp16_masked(
        handle, static_cast<const __half*>(q.data_ptr()),
        static_cast<const __half*>(k.data_ptr()),
        static_cast<const __half*>(v.data_ptr()),
        static_cast<__half*>(logits.data_ptr()),
        static_cast<__half*>(out.data_ptr()), sq, sk, heads, dim,
        static_cast<float>(scale), stream);
  } else {
    TORCH_CHECK(q.stride(0) == k.stride(0) && q.stride(0) == v.stride(0),
                "BF16 q/k/v must share one token stride");
    attention_mha_bf16_masked(
        handle, static_cast<const __nv_bfloat16*>(q.data_ptr()),
        static_cast<const __nv_bfloat16*>(k.data_ptr()),
        static_cast<const __nv_bfloat16*>(v.data_ptr()),
        static_cast<__nv_bfloat16*>(logits.data_ptr()),
        static_cast<__nv_bfloat16*>(out.data_ptr()), sq, sk, heads, dim,
        static_cast<float>(scale), checked_int(logits.size(2), "logits stride"),
        checked_int(q.stride(0), "qkv token stride"), stream);
  }
  C10_CUDA_KERNEL_LAUNCH_CHECK();
}

void attention_mha_fp16_masked_op(
    torch::Tensor const& q, torch::Tensor const& k, torch::Tensor const& v,
    torch::Tensor& logits, torch::Tensor& out, double scale) {
  TORCH_CHECK(q.scalar_type() == torch::kFloat16,
              "attention_mha_fp16_masked requires FP16 q/k/v");
  masked_mha_forward_static(q, k, v, logits, out, scale);
}

void attention_mha_bf16_masked_op(
    torch::Tensor const& q, torch::Tensor const& k, torch::Tensor const& v,
    torch::Tensor& logits, torch::Tensor& out, double scale,
    int64_t qkv_token_stride) {
  TORCH_CHECK(q.scalar_type() == torch::kBFloat16,
              "attention_mha_bf16_masked requires BF16 q/k/v");
  TORCH_CHECK(qkv_token_stride == q.stride(0),
              "qkv_token_stride must match q.stride(0)");
  masked_mha_forward_static(q, k, v, logits, out, scale);
}

void masked_mha_forward_seqused_static(
    torch::Tensor const& q, torch::Tensor const& k, torch::Tensor const& v,
    torch::Tensor const& valid_k, torch::Tensor& logits, torch::Tensor& out,
    double scale) {
  check_qkv(q, "q", torch::kFloat16);
  TORCH_CHECK(k.is_cuda() && v.is_cuda() && k.scalar_type() == torch::kFloat16 &&
                  v.scalar_type() == torch::kFloat16 && k.dim() == 2 &&
                  v.dim() == 2 && k.is_contiguous() && v.is_contiguous() &&
                  k.sizes() == v.sizes(),
              "k/v must be contiguous FP16 tensors with shape (S_kv_max, D)");
  TORCH_CHECK(q.is_contiguous() && q.size(2) == k.size(1),
              "q must be contiguous and share head_dim with k/v");
  TORCH_CHECK(k.size(0) <= 1024,
              "forward_seqused_static supports S_kv_max <= 1024");
  TORCH_CHECK(valid_k.is_cuda() && valid_k.scalar_type() == torch::kInt &&
                  valid_k.numel() == 1 && valid_k.is_contiguous(),
              "valid_k must be a contiguous CUDA int32 scalar tensor");
  TORCH_CHECK(q.get_device() == k.get_device() && q.get_device() == v.get_device() &&
                  q.get_device() == valid_k.get_device(),
              "q/k/v/valid_k must be on the same device");
  TORCH_CHECK(logits.is_cuda() && logits.scalar_type() == torch::kFloat16 &&
                  logits.is_contiguous() && logits.dim() == 2 &&
                  logits.size(0) == q.size(0) * q.size(1) &&
                  logits.size(1) == k.size(0),
              "logits must be contiguous FP16 with shape (S_q * H, S_kv_max)");
  TORCH_CHECK(out.is_cuda() && out.scalar_type() == torch::kFloat16 &&
                  out.is_contiguous() && out.sizes() == q.sizes(),
              "out must be contiguous FP16 and match q");

  c10::cuda::CUDAGuard guard(q.device());
  auto stream = at::cuda::getCurrentCUDAStream(q.get_device()).stream();
  auto handle = at::cuda::getCurrentCUDABlasHandle();
  attention_qkv_fp16_seqused_v2(
      handle, static_cast<const __half*>(q.data_ptr()),
      static_cast<const __half*>(k.data_ptr()),
      static_cast<const __half*>(v.data_ptr()),
      static_cast<__half*>(logits.data_ptr()),
      static_cast<__half*>(out.data_ptr()), checked_int(q.size(0), "S_q"),
      checked_int(k.size(0), "S_kv_max"), checked_int(q.size(1), "heads"),
      checked_int(q.size(2), "head_dim"),
      static_cast<const int*>(valid_k.data_ptr()), static_cast<float>(scale),
      stream);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
}

}  // namespace

TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
  ops.def("forward_static(Tensor q, Tensor k, Tensor v, Tensor! logits, Tensor! out, float scale) -> ()");
  ops.def("attention_mha_fp16_masked(Tensor q, Tensor k, Tensor v, Tensor! logits, Tensor! out, float scale) -> ()");
  ops.def("attention_mha_bf16_masked(Tensor q, Tensor k, Tensor v, Tensor! logits, Tensor! out, float scale, int qkv_token_stride) -> ()");
  ops.def("forward_seqused_static(Tensor q, Tensor k, Tensor v, Tensor valid_k, Tensor! logits, Tensor! out, float scale) -> ()");
  ops.impl("forward_static", torch::kCUDA, &masked_mha_forward_static);
  ops.impl("attention_mha_fp16_masked", torch::kCUDA,
           &attention_mha_fp16_masked_op);
  ops.impl("attention_mha_bf16_masked", torch::kCUDA,
           &attention_mha_bf16_masked_op);
  ops.impl("forward_seqused_static", torch::kCUDA,
           &masked_mha_forward_seqused_static);
}

REGISTER_EXTENSION(TORCH_EXTENSION_NAME)