Coverage Report

Created: 2026-05-04 15:30

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
/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)