diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 2c63f8d302cc..df9c0e877d2f 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -1216,6 +1216,9 @@ struct test_case { } virtual bool run_whole_graph() { return false; } + + // evaluate the same graph this many times, new inputs each time (reaches graph capture/replay in backends that record graphs) + virtual int n_eval() { return 1; } virtual std::vector fusion_test_nodes() { return {}; } virtual bool use_weight_context() { return false; } @@ -1490,9 +1493,15 @@ struct test_case { if (fused_nodes_to_verify.size() == 0 && run_whole_graph()) { fused_nodes_to_verify.push_back(out); } - const bool cmp_ok = ggml_backend_compare_graph_backend(backend1, backend2, gf, callback, &ud, - run_whole_graph() ? fused_nodes_to_verify.data() : nullptr, - fused_nodes_to_verify.size()); + bool cmp_ok = true; + for (int i = 0; i < n_eval() && cmp_ok && ud.ok; i++) { + if (i > 0) { + initialize_tensors(ctx.get()); + } + cmp_ok = ggml_backend_compare_graph_backend(backend1, backend2, gf, callback, &ud, + run_whole_graph() ? fused_nodes_to_verify.data() : nullptr, + fused_nodes_to_verify.size()); + } // Create test result bool test_passed = ud.ok && cmp_ok; @@ -4440,6 +4449,66 @@ struct test_gated_delta_net : public test_case { } }; +// GET_ROWS -> RESHAPE -> GATED_DELTA_NET: one recurrent state row gathered from the cache (CUDA folds the gather into the GDN kernel) +struct test_gated_delta_net_gathered_state : public test_case { + const int64_t head_count; + const int64_t head_size; + const int64_t cache_rows; + int n_init = 0; + + std::string vars() override { + return VARS_TO_STR3(head_count, head_size, cache_rows); + } + + test_gated_delta_net_gathered_state(int64_t head_count = 4, int64_t head_size = 16, int64_t cache_rows = 5) + : head_count(head_count), head_size(head_size), cache_rows(cache_rows) {} + + bool run_whole_graph() override { return true; } + + // eager, then CUDA graph capture, then replays with a different row + int n_eval() override { return 4; } + + ggml_tensor * build_graph(ggml_context * ctx) override { + ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, head_size, head_count, 1, 1); + ggml_tensor * k = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, head_size, head_count, 1, 1); + ggml_tensor * v = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, head_size, head_count, 1, 1); + ggml_tensor * g = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, 1, head_count, 1, 1); + ggml_tensor * beta = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, 1, head_count, 1, 1); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, head_size * head_size * head_count, cache_rows); + ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + ggml_set_name(v, "v"); + ggml_set_name(g, "g"); + ggml_set_name(beta, "beta"); + ggml_set_name(ids, "ids"); + q = ggml_l2_norm(ctx, q, 1e-6f); + k = ggml_l2_norm(ctx, k, 1e-6f); + + ggml_tensor * state = ggml_get_rows(ctx, cache, ids); + state = ggml_reshape_4d(ctx, state, head_size, head_size, head_count, 1); + return ggml_gated_delta_net(ctx, q, k, v, g, beta, state, 1); + } + + void initialize_tensors(ggml_context * ctx) override { + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != nullptr; t = ggml_get_next_tensor(ctx, t)) { + if (ggml_is_view_op(t->op)) { continue; } + if (strcmp(t->name, "g") == 0) { + init_tensor_uniform(t, -20.0f, -1e-4f); + } else if (strcmp(t->name, "beta") == 0) { + init_tensor_uniform(t, 0.0f, 1.0f); + } else if (strcmp(t->name, "v") == 0) { + init_tensor_uniform(t, -0.3f, 5.0f); + } else if (strcmp(t->name, "ids") == 0) { + // nonzero, and a different row on every evaluation + const int32_t row = 1 + n_init % (cache_rows - 1); + ggml_backend_tensor_set(t, &row, 0, sizeof(row)); + } else { + init_tensor_uniform(t); + } + } + n_init++; + } +}; + // GGML_OP_GATED_LINEAR_ATTN struct test_gla : public test_case { const ggml_type type; @@ -10283,6 +10352,8 @@ static std::vector> make_test_cases_eval() { } } + test_cases.emplace_back(new test_gated_delta_net_gathered_state()); + test_cases.emplace_back(new test_gated_delta_net_gathered_state(32, 128, 3)); test_cases.emplace_back(new test_gated_delta_net(GGML_TYPE_F32, 32, 128, 1, 1)); test_cases.emplace_back(new test_gated_delta_net(GGML_TYPE_F32, 32, 16, 1, 1)); test_cases.emplace_back(new test_gated_delta_net(GGML_TYPE_F32, 32, 16, 1, 1, 1, true, true));