MCPcopy Create free account
hub / github.com/ModelTC/LightX2V / sdp_torch

Function sdp_torch

lightx2v_kernel_xpu/csrc/sdp.cpp:36–84  ·  view source on GitHub ↗

────────────────────────────────────────────────────────────────────────────── sdp: unified Flash Attention SDP Q: [1, L_q, H_q, 128] fp16 or bf16 K: [1, L_kv, H_kv, 128] same dtype as Q V: [1, L_kv, H_kv, 128] same dtype as Q returns: [1, L_q, H_q, 128] same dtype as Q fp16 → sdp_fp16 (optimised FP16 ESIMD kernel) bf16 → sdp_bf16io (BF16 I/O, FP16 internal ESIMD kernel) ───────────────

Source from the content-addressed store, hash-verified

34// bf16 → sdp_bf16io (BF16 I/O, FP16 internal ESIMD kernel)
35// ──────────────────────────────────────────────────────────────────────────────
36torch::Tensor sdp_torch(
37 torch::Tensor Q,
38 torch::Tensor K,
39 torch::Tensor V)
40{
41 check_sdp_tensor(Q, "Q");
42 check_sdp_tensor(K, "K");
43 check_sdp_tensor(V, "V");
44 TORCH_CHECK(Q.scalar_type() == K.scalar_type() && Q.scalar_type() == V.scalar_type(),
45 "Q, K, V must have the same dtype");
46
47 const int q_len = (int)Q.size(1);
48 const int kv_len = (int)K.size(1);
49 const int headQ = (int)Q.size(2);
50 const int headKv = (int)K.size(2);
51
52 auto out = torch::empty_like(Q);
53
54 // Cache normAlpha: only reallocate when headQ changes (avoids per-call
55 // XPU malloc + fill-kernel on every sdp() invocation).
56 static torch::Tensor s_normAlpha;
57 static int s_headQ = -1;
58 if (headQ != s_headQ) {
59 s_normAlpha = torch::ones({headQ * 128},
60 torch::dtype(torch::kFloat).device(Q.device()));
61 s_headQ = headQ;
62 }
63 const auto& normAlpha = s_normAlpha;
64
65 sycl::queue& sq = utils::get_queue(Q.device());
66
67 auto dispatch = [&](auto kernel) {
68 kernel(Q.data_ptr(), K.data_ptr(), V.data_ptr(),
69 normAlpha.data_ptr(), out.data_ptr(),
70 q_len, kv_len, headQ, headKv, &sq);
71 };
72
73 switch (Q.scalar_type()) {
74 case ST::Half:
75 dispatch(sdp_fp16); break;
76 case ST::BFloat16:
77 dispatch(sdp_bf16io); break;
78 default:
79 TORCH_CHECK(false,
80 "sdp: unsupported dtype, only FP16 and BF16 are supported");
81 }
82
83 return out;
84}

Callers

nothing calls this directly

Calls 2

check_sdp_tensorFunction · 0.85
deviceMethod · 0.45

Tested by

no test coverage detected