/home/liu/actions-runner/_work/ccv/ccv/lib/nnc/cmd/gated_delta/ccv_nnc_gated_delta.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_gated_delta_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 | 0 | { |
7 | | // Inputs: q, k, v, decay / log_decay, beta, state_in. |
8 | | // Outputs: y, state_out. |
9 | 0 | if (input_size == 6 && output_size == 2 && (input_bitmasks[0] & 63u) == 63u && (output_bitmasks[0] & 3u) == 3u) |
10 | 0 | return 1; |
11 | 0 | return 0; |
12 | 0 | } |
13 | | |
14 | | static int _ccv_nnc_gated_delta_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) |
15 | 0 | { |
16 | 0 | return 0; |
17 | 0 | } |
18 | | |
19 | | static int _ccv_nnc_gated_delta_allow_inplace(const ccv_nnc_cmd_param_t cmd, const int input_idx, const int input_size, const int output_idx, const int output_size) |
20 | 0 | { |
21 | 0 | if (input_idx == 5 && output_idx == 1) |
22 | 0 | return 1; |
23 | 0 | return 0; |
24 | 0 | } |
25 | | |
26 | | static void _ccv_nnc_gated_delta_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) |
27 | 0 | { |
28 | 0 | assert(input_size == 6); |
29 | 0 | assert(output_size == 2); |
30 | 0 | const int q_nd = ccv_nnc_tensor_nd(inputs[0].dim); |
31 | 0 | const int k_nd = ccv_nnc_tensor_nd(inputs[1].dim); |
32 | 0 | const int v_nd = ccv_nnc_tensor_nd(inputs[2].dim); |
33 | 0 | const int log_decay_nd = ccv_nnc_tensor_nd(inputs[3].dim); |
34 | 0 | const int beta_nd = ccv_nnc_tensor_nd(inputs[4].dim); |
35 | 0 | const int state_nd = ccv_nnc_tensor_nd(inputs[5].dim); |
36 | 0 | assert(q_nd == 4); |
37 | 0 | assert(k_nd == 4); |
38 | 0 | assert(v_nd == 4); |
39 | 0 | assert(log_decay_nd == 3); |
40 | 0 | assert(beta_nd == 3); |
41 | 0 | assert(state_nd == 4); |
42 | 0 | assert(inputs[0].dim[0] == inputs[1].dim[0]); |
43 | 0 | assert(inputs[0].dim[0] == inputs[2].dim[0]); |
44 | 0 | assert(inputs[0].dim[1] == inputs[1].dim[1]); |
45 | 0 | assert(inputs[0].dim[1] == inputs[2].dim[1]); |
46 | 0 | assert(inputs[0].dim[2] == inputs[1].dim[2]); |
47 | 0 | assert(inputs[0].dim[3] == inputs[1].dim[3]); |
48 | 0 | assert(inputs[2].dim[2] % inputs[0].dim[2] == 0); |
49 | 0 | assert(inputs[3].dim[0] == inputs[0].dim[0]); |
50 | 0 | assert(inputs[3].dim[1] == inputs[0].dim[1]); |
51 | 0 | assert(inputs[3].dim[2] == inputs[2].dim[2]); |
52 | 0 | assert(inputs[4].dim[0] == inputs[3].dim[0]); |
53 | 0 | assert(inputs[4].dim[1] == inputs[3].dim[1]); |
54 | 0 | assert(inputs[4].dim[2] == inputs[3].dim[2]); |
55 | 0 | assert(inputs[5].dim[0] == inputs[0].dim[0]); |
56 | 0 | assert(inputs[5].dim[1] == inputs[2].dim[2]); |
57 | 0 | assert(inputs[5].dim[2] == inputs[2].dim[3]); |
58 | 0 | assert(inputs[5].dim[3] == inputs[0].dim[3]); |
59 | 0 | outputs[0] = inputs[2]; |
60 | 0 | outputs[1] = inputs[5]; |
61 | 0 | } |
62 | | |
63 | | REGISTER_COMMAND(CCV_NNC_GATED_DELTA_FORWARD)(ccv_nnc_cmd_registry_t* const registry) |
64 | | FIND_BACKEND(ccv_nnc_gated_delta_cpu_ref.c, mps/ccv_nnc_gated_delta_mps.m) |
65 | 1 | { |
66 | 1 | registry->bitmask = _ccv_nnc_gated_delta_forw_bitmask; |
67 | 1 | registry->tensor_auto = _ccv_nnc_gated_delta_tensor_auto_forw; |
68 | 1 | registry->allow_inplace = _ccv_nnc_gated_delta_allow_inplace; |
69 | 1 | } |
70 | | |
71 | | REGISTER_COMMAND(CCV_NNC_GATED_DELTA_BACKWARD)(ccv_nnc_cmd_registry_t* const registry) |
72 | 1 | { |
73 | 1 | registry->bitmask = _ccv_nnc_gated_delta_back_bitmask; |
74 | 1 | } |
75 | | |
76 | | //@REGISTER_EASY_COMMAND_MACRO(CCV_NNC_GATED_DELTA_FORWARD) |
77 | | #define CMD_GATED_DELTA_FORWARD_X_F(...) ("This should not be used, you should have either 0 or 1 parameter for CMD_GATED_DELTA_FORWARD") |
78 | | #define CMD_GATED_DELTA_FORWARD_X_0() ccv_nnc_cmd(CCV_NNC_GATED_DELTA_FORWARD, 0, ((ccv_nnc_cmd_param_t){.size={.dim={1,1,1}},.gated_delta={.log_decay=1}}), 0) |
79 | | #define CMD_GATED_DELTA_FORWARD_X_1(_log_decay) ccv_nnc_cmd(CCV_NNC_GATED_DELTA_FORWARD, 0, ((ccv_nnc_cmd_param_t){.size={.dim={1,1,1}},.gated_delta={.log_decay=(_log_decay)}}), 0) |
80 | | #define CMD_GATED_DELTA_FORWARD_X_SEL(_0, _1, _FX, ...) _FX |
81 | | #define CMD_GATED_DELTA_FORWARD(...) CMD_GATED_DELTA_FORWARD_X_SEL(CMD_GATED_DELTA_FORWARD_X_F, ##__VA_ARGS__, CMD_GATED_DELTA_FORWARD_X_1, CMD_GATED_DELTA_FORWARD_X_0)(__VA_ARGS__) |