Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
77 changes: 74 additions & 3 deletions tests/test-backend-ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<ggml_tensor *> fusion_test_nodes() { return {}; }
virtual bool use_weight_context() { return false; }

Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -10283,6 +10352,8 @@ static std::vector<std::unique_ptr<test_case>> 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));
Expand Down