/home/liu/actions-runner/_work/ccv/ccv/lib/nnc/cmd/scaled_dot_product_attention/ccv_nnc_scaled_dot_product_attention.c
Line | Count | Source |
1 | | #include "ccv.h" |
2 | | #include "nnc/ccv_nnc.h" |
3 | | #include "nnc/ccv_nnc_internal.h" |
4 | | |
5 | | static int _ccv_nnc_scaled_dot_product_attention_forw_bitmask(const ccv_nnc_cmd_param_t cmd, const int input_size, const int output_size, const uint64_t* const input_bitmasks, const int input_bitmask_size, const uint64_t* const output_bitmasks, const int output_bitmask_size) |
6 | 17 | { |
7 | | // 6 inputs (query, key, value, [attn_mask], [unify head weight], [unify head bias]) |
8 | | // 8 inputs with varlen (query, key, value, 0, 0, 0, q_seq_offsets, kv_seq_offsets) |
9 | | // 3 outputs (y, softmax_lse, [qkv]) |
10 | 17 | if (cmd.scaled_dot_product_attention.is_varlen && input_size == 83 && (input_bitmasks[0] & 255u) == ((1u << 0) | (1u << 1) | (1u << 2) | (1u << 6) | (1u << 7))3 && (output_bitmasks[0] & 1u) == 1u3 ) |
11 | 3 | return 1; |
12 | 14 | if (input_size == 6 && (input_bitmasks[0] & 55u) == 55u && (output_bitmasks[0] & 7u) == 7u6 ) |
13 | 4 | return 1; |
14 | 10 | if (input_size == 5 && (input_bitmasks[0] & 23u) == 23u0 && (output_bitmasks[0] & 7u) == 7u0 ) |
15 | 0 | return 1; |
16 | 10 | if ((input_bitmasks[0] & 55u) == 7u && (output_bitmasks[0] & 3u) == 3u8 ) |
17 | 6 | return 1; |
18 | 4 | return 0; |
19 | 10 | } |
20 | | |
21 | | |
22 | | static int _ccv_nnc_allow_query_inplace(const ccv_nnc_cmd_param_t cmd, const int input_idx, const int input_size, const int output_idx, const int output_size) |
23 | 49 | { |
24 | 49 | if (input_idx == 0 && output_idx == 012 ) |
25 | 6 | return 1; |
26 | 43 | return 0; |
27 | 49 | } |
28 | | |
29 | | static int _ccv_nnc_scaled_dot_product_attention_back_bitmask(const ccv_nnc_cmd_param_t cmd, const int input_size, const int output_size, const uint64_t* const input_bitmasks, const int input_bitmask_size, const uint64_t* const output_bitmasks, const int output_bitmask_size) |
30 | 1 | { |
31 | | // 1, 0, 0, 8, 16, 32, 64?, 128?, 256?, 512, 1024, 2048? |
32 | | // 1, 2, 4, 8, 16, 32 |
33 | | // Inputs (gradient, 0, 0, q, k, v, [attn_mask], [head weight], [bias], y, saved softmax_lse, qkv) |
34 | | // Output (dquery, dkey, dvalue, [attn mask], dweight, dbias) [cannot diff against attn_mask] |
35 | 1 | if ((input_bitmasks[0] & 4025u) == 4025u && (output_bitmasks[0] & 63u) == 55u0 ) |
36 | 0 | return 1; |
37 | 1 | if ((input_bitmasks[0] & 3769u) == 3769u && (output_bitmasks[0] & 31u) == 23u0 ) |
38 | 0 | return 1; |
39 | 1 | if ((input_bitmasks[0] & 1593u) == 1593u && (output_bitmasks[0] & 7u) == 7u0 ) |
40 | 0 | return 1; |
41 | 1 | return 0; |
42 | 1 | } |
43 | | |
44 | | static void _ccv_nnc_scaled_dot_product_attention_tensor_auto_forw(const ccv_nnc_cmd_param_t cmd, const ccv_nnc_tensor_param_t* const inputs, const int input_size, const ccv_nnc_hint_t hint, ccv_nnc_tensor_param_t* const outputs, const int output_size) |
45 | 23 | { |
46 | 23 | assert(input_size >= 3); |
47 | 23 | assert(output_size >= 1); |
48 | 23 | const int q_nd = ccv_nnc_tensor_nd(inputs[0].dim); |
49 | 23 | assert(q_nd == 3 || q_nd == 4); |
50 | 23 | const int k_nd = ccv_nnc_tensor_nd(inputs[1].dim); |
51 | 23 | assert(k_nd == 3 || k_nd == 4); |
52 | 23 | const int v_nd = ccv_nnc_tensor_nd(inputs[2].dim); |
53 | 23 | assert(v_nd == 3 || v_nd == 4); |
54 | 23 | assert(q_nd == k_nd && k_nd == v_nd); |
55 | 23 | if (!cmd.scaled_dot_product_attention.is_varlen && input_size > 419 ) |
56 | 12 | { |
57 | 12 | assert(output_size >= 3); |
58 | 12 | outputs[0] = inputs[0]; |
59 | 12 | outputs[0].dim[1] = inputs[0].dim[1]; // sequence length matches query, embedding size matches value * num_head. |
60 | 12 | outputs[0].dim[2] = inputs[2].dim[v_nd - 1] * (q_nd == 4 ? inputs[0].dim[2] : 10 ); |
61 | 12 | outputs[0].dim[3] = 0; |
62 | | // This is saved softmax_lse, which would be in 32F if exists. |
63 | 12 | outputs[1] = inputs[0]; |
64 | 12 | outputs[1].dim[q_nd - 3] = inputs[0].dim[q_nd - 2]; |
65 | 12 | outputs[1].dim[q_nd - 2] = inputs[0].dim[q_nd - 3]; |
66 | 12 | outputs[1].dim[q_nd - 1] = 0; |
67 | 12 | outputs[1].datatype = CCV_32F; |
68 | 12 | outputs[2] = inputs[0]; |
69 | 12 | outputs[2].dim[q_nd - 1] = inputs[2].dim[v_nd - 1]; // sequence length matches query, embedding size matches value. |
70 | 12 | } else { |
71 | 11 | outputs[0] = inputs[0]; |
72 | 11 | outputs[0].dim[q_nd - 1] = inputs[2].dim[v_nd - 1]; // sequence length matches query, embedding size matches value. |
73 | 11 | if (output_size == 1) |
74 | 3 | return; |
75 | 11 | assert(output_size > 1)8 ; |
76 | | // This is saved softmax_lse, which would be in 32F if exists. |
77 | 8 | outputs[1] = inputs[0]; |
78 | 8 | if (cmd.scaled_dot_product_attention.is_varlen) |
79 | 4 | { |
80 | 4 | assert(input_size >= 8); |
81 | 4 | assert(q_nd == 4); |
82 | 4 | assert(inputs[6].dim[0] > 0); |
83 | 4 | assert(cmd.scaled_dot_product_attention.max_seqlen_q > 0); |
84 | 4 | outputs[1].dim[0] = inputs[6].dim[0] - 1; |
85 | 4 | outputs[1].dim[1] = inputs[0].dim[2]; |
86 | 4 | outputs[1].dim[2] = cmd.scaled_dot_product_attention.max_seqlen_q; |
87 | 4 | } else { |
88 | 4 | outputs[1].dim[q_nd - 3] = inputs[0].dim[q_nd - 2]; |
89 | 4 | outputs[1].dim[q_nd - 2] = inputs[0].dim[q_nd - 3]; |
90 | 4 | } |
91 | 8 | outputs[1].dim[q_nd - 1] = 0; |
92 | 8 | outputs[1].datatype = CCV_32F; |
93 | 8 | } |
94 | 23 | } |
95 | | |
96 | | static void _ccv_nnc_scaled_dot_product_attention_tensor_auto_back(const ccv_nnc_cmd_param_t cmd, const ccv_nnc_tensor_param_t* const inputs, const int input_size, const ccv_nnc_hint_t hint, ccv_nnc_tensor_param_t* const outputs, const int output_size) |
97 | 1 | { |
98 | 1 | assert(input_size >= 6); |
99 | 1 | assert(output_size >= 3); |
100 | 1 | int i; |
101 | 4 | for (i = 0; i < output_size; i++3 ) |
102 | 3 | outputs[i] = inputs[3 + i]; |
103 | 1 | } |
104 | | |
105 | | REGISTER_COMMAND(CCV_NNC_SCALED_DOT_PRODUCT_ATTENTION_FORWARD)(ccv_nnc_cmd_registry_t* const registry) |
106 | | FIND_BACKEND(ccv_nnc_scaled_dot_product_attention_cpu_ref.c, mps/ccv_nnc_scaled_dot_product_attention_mps.m, gpu/ccv_nnc_scaled_dot_product_attention_flash_attn.cu) |
107 | 1 | { |
108 | 1 | registry->bitmask = _ccv_nnc_scaled_dot_product_attention_forw_bitmask; |
109 | 1 | registry->tensor_auto = _ccv_nnc_scaled_dot_product_attention_tensor_auto_forw; |
110 | 1 | registry->allow_inplace = _ccv_nnc_allow_query_inplace; |
111 | 1 | } |
112 | | |
113 | | REGISTER_COMMAND(CCV_NNC_SCALED_DOT_PRODUCT_ATTENTION_BACKWARD)(ccv_nnc_cmd_registry_t* const registry) |
114 | | FIND_BACKEND(ccv_nnc_scaled_dot_product_attention_cpu_ref.c, mps/ccv_nnc_scaled_dot_product_attention_mps.m, gpu/ccv_nnc_scaled_dot_product_attention_flash_attn.cu) |
115 | 1 | { |
116 | 1 | registry->bitmask = _ccv_nnc_scaled_dot_product_attention_back_bitmask; |
117 | 1 | registry->tensor_auto = _ccv_nnc_scaled_dot_product_attention_tensor_auto_back; |
118 | 1 | } |
119 | | |
120 | | //@REGISTER_EASY_COMMAND_MACRO(CCV_NNC_SCALED_DOT_PRODUCT_ATTENTION_FORWARD) |
121 | | #define CMD_SCALED_DOT_PRODUCT_ATTENTION_FORWARD(_scale, _is_causal) ccv_nnc_cmd(CCV_NNC_SCALED_DOT_PRODUCT_ATTENTION_FORWARD, 0, ((ccv_nnc_cmd_param_t){.size={.dim={1,1,1}},.scaled_dot_product_attention={.scale=_scale,.is_causal=_is_causal}}), 0) |
122 | | //@REGISTER_EASY_COMMAND_MACRO(CCV_NNC_SCALED_DOT_PRODUCT_ATTENTION_BACKWARD) |
123 | | #define CMD_SCALED_DOT_PRODUCT_ATTENTION_BACKWARD(_scale, _is_causal) ccv_nnc_cmd(CCV_NNC_SCALED_DOT_PRODUCT_ATTENTION_BACKWARD, 0, ((ccv_nnc_cmd_param_t){.size={.dim={1,1,1}},.scaled_dot_product_attention={.scale=_scale,.is_causal=_is_causal}}), 0) |