Skip to content

Commit 9c18557

Browse files
committed
tests: add ABS and LOG coverage to test-backend-ops
1 parent 192067b commit 9c18557

1 file changed

Lines changed: 21 additions & 115 deletions

File tree

tests/test-backend-ops.cpp

Lines changed: 21 additions & 115 deletions
Original file line numberDiff line numberDiff line change
@@ -2469,13 +2469,8 @@ struct test_set_rows : public test_case {
24692469
// See dicussion here: https://github.com/ggml-org/llama.cpp/pull/23760#issuecomment-4566312209
24702470
double max_nmse_err(ggml_backend_t backend) override {
24712471
ggml_backend_reg_t reg = ggml_backend_dev_backend_reg(ggml_backend_get_device(backend));
2472-
if (type_dst == GGML_TYPE_Q8_0) {
2473-
if (strcmp(ggml_backend_reg_name(reg), "WebGPU") == 0) {
2474-
return std::max(test_case::max_nmse_err(backend), 2e-7);
2475-
}
2476-
if (strcmp(ggml_backend_reg_name(reg), "HTP") == 0) {
2477-
return std::max(test_case::max_nmse_err(backend), 5e-6);
2478-
}
2472+
if (type_dst == GGML_TYPE_Q8_0 && strcmp(ggml_backend_reg_name(reg), "WebGPU") == 0) {
2473+
return std::max(test_case::max_nmse_err(backend), 2e-7);
24792474
}
24802475
return test_case::max_nmse_err(backend);
24812476
}
@@ -4125,10 +4120,9 @@ struct test_ssm_scan : public test_case {
41254120
const int64_t n_seqs;
41264121
const bool xbc_overlap;
41274122
const int64_t K;
4128-
const bool weak_decay;
41294123

41304124
std::string vars() override {
4131-
return VARS_TO_STR10(type, d_state, head_dim, n_head, n_group, n_seq_tokens, n_seqs, xbc_overlap, K, weak_decay);
4125+
return VARS_TO_STR9(type, d_state, head_dim, n_head, n_group, n_seq_tokens, n_seqs, xbc_overlap, K);
41324126
}
41334127

41344128
test_ssm_scan(ggml_type type = GGML_TYPE_F32,
@@ -4139,9 +4133,8 @@ struct test_ssm_scan : public test_case {
41394133
int64_t n_seq_tokens = 32,
41404134
int64_t n_seqs = 32,
41414135
bool xbc_overlap = false,
4142-
int64_t K = 1,
4143-
bool weak_decay = false)
4144-
: type(type), d_state(d_state), head_dim(head_dim), n_head(n_head), n_group(n_group), n_seq_tokens(n_seq_tokens), n_seqs(n_seqs), xbc_overlap(xbc_overlap), K(K), weak_decay(weak_decay) {}
4136+
int64_t K = 1)
4137+
: type(type), d_state(d_state), head_dim(head_dim), n_head(n_head), n_group(n_group), n_seq_tokens(n_seq_tokens), n_seqs(n_seqs), xbc_overlap(xbc_overlap), K(K) {}
41454138

41464139
double max_nmse_err() override {
41474140
// SSD path (head_dim > 1) uses FP16 intermediates (M matrix, X_dt); Mamba-1 is pure FP32.
@@ -4194,7 +4187,7 @@ struct test_ssm_scan : public test_case {
41944187
continue;
41954188
} else if (t->ne[1] == n_head && t->ne[2] == 1) {
41964189
// A {1 or d_state, n_head}: negative decay (2-D tensor, ne[2]==1 distinguishes from 3-D/4-D tensors)
4197-
init_tensor_uniform(t, weak_decay ? -0.02f : -1.0f, weak_decay ? -0.005f : -0.5f);
4190+
init_tensor_uniform(t, -1.0f, -0.5f);
41984191
} else {
41994192
init_tensor_uniform(t);
42004193
}
@@ -9118,10 +9111,6 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
91189111
test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 4, 2, false, /*K=*/4)); // Mamba-2 rollback snapshots
91199112
test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 8, 2, false, /*K=*/3)); // Mamba-2 rollback overflow
91209113
test_cases.emplace_back(new test_ssm_scan_rollback(GGML_TYPE_F32, 128, 64, 16, 2, 8, 2, /*K=*/3)); // rollback snapshots match prefix states
9121-
test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 64, 4)); // Metal SSD one chunk MMA only, no seq tail
9122-
test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 65, 2)); // SSD one chunk + 1-token sequential tail
9123-
test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 128, 2)); // SSD multi-chunk, no tail (exercises the chunk-to-chunk state handoff)
9124-
test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 128, 2, false, /*K=*/1, /*weak_decay=*/true)); // SSD multi-chunk, carried state not numerically negligible
91259114

91269115
test_cases.emplace_back(new test_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 1, 1));
91279116
test_cases.emplace_back(new test_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 32, 1));
@@ -10306,6 +10295,20 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
1030610295
test_cases.emplace_back(new test_cumsum(GGML_TYPE_F32, { 2048, 16, 5, 4 }));
1030710296
test_cases.emplace_back(new test_cumsum(GGML_TYPE_F32, { 20000, 10, 4, 1 }));
1030810297

10298+
// GGML_UNARY_OP_ABS (hexagon HTP kernel)
10299+
test_cases.emplace_back(new test_unary(GGML_UNARY_OP_ABS, GGML_TYPE_F32, { 256, 1, 1, 1 })); // Mamba-130M ssm_a
10300+
test_cases.emplace_back(new test_unary(GGML_UNARY_OP_ABS, GGML_TYPE_F32, { 5120, 1, 1, 1 })); // Mamba-3B ssm_a (large)
10301+
test_cases.emplace_back(new test_unary(GGML_UNARY_OP_ABS, GGML_TYPE_F32, { 4096, 1, 1, 1 })); // aligned vector
10302+
test_cases.emplace_back(new test_unary(GGML_UNARY_OP_ABS, GGML_TYPE_F32, { 96, 32, 1, 1 })); // nloe != 0 boundary (96 % 32 != 0)
10303+
10304+
// GGML_OP_LOG (hexagon HTP kernel)
10305+
test_cases.emplace_back(new test_log(GGML_TYPE_F32, { 32000, 1, 1, 1 })); // LLaMA-3 log_softmax, single token
10306+
test_cases.emplace_back(new test_log(GGML_TYPE_F32, { 32000, 8, 1, 1 })); // LLaMA-3 batch=8
10307+
test_cases.emplace_back(new test_log(GGML_TYPE_F32, { 151936, 1, 1, 1 })); // Qwen2.5 vocab
10308+
test_cases.emplace_back(new test_log(GGML_TYPE_F32, { 16, 5120, 1, 1 })); // Mamba ssm_a log
10309+
test_cases.emplace_back(new test_log(GGML_TYPE_F32, { 288, 16, 1, 1 })); // nloe boundary
10310+
test_cases.emplace_back(new test_log(GGML_TYPE_F32, { 128256, 1, 1, 1 })); // LLaMA-3.1 vocab, sampler top_p/min_p/temp_ext
10311+
1030910312
for (int bs : {1, 2, 3, 4, 5, 8, 512}) {
1031010313
for (ggml_type type_a : all_types) {
1031110314
for (ggml_type type_b : {GGML_TYPE_F32}) {
@@ -10584,101 +10587,6 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_from_file(const c
1058410587
return test_cases;
1058510588
}
1058610589

10587-
// ---- FA vec (Q,NE): forced-config numerical slice (Metal only) ----
10588-
using set_fa_vec_override_t = void (*)(int, int);
10589-
using clear_fa_vec_override_t = void (*)(void);
10590-
10591-
// NL = 32/NE must divide both dk/4 and dv/4.
10592-
static std::vector<int> fa_vec_legal_ne(int dk, int dv) {
10593-
std::vector<int> r;
10594-
for (int ne : {1, 2, 4}) {
10595-
const int nl = 32 / ne;
10596-
if ((dk/4) % nl == 0 && (dv/4) % nl == 0) {
10597-
r.push_back(ne);
10598-
}
10599-
}
10600-
return r;
10601-
}
10602-
10603-
static bool op_names_filter_selects(const char * op_names_filter, const char * op_name) {
10604-
if (!op_names_filter) {
10605-
return true;
10606-
}
10607-
std::string_view filter(op_names_filter);
10608-
while (!filter.empty()) {
10609-
auto comma_pos = filter.find_first_of(',');
10610-
const auto lparen_pos = filter.find_first_of('(');
10611-
std::string_view entry;
10612-
if (lparen_pos < comma_pos) {
10613-
const auto rparen_pos = filter.find_first_of(')');
10614-
comma_pos = filter.find_first_of(',', rparen_pos);
10615-
entry = filter.substr(0, lparen_pos);
10616-
} else {
10617-
entry = filter.substr(0, comma_pos);
10618-
}
10619-
if (entry == op_name) {
10620-
return true;
10621-
}
10622-
filter = comma_pos != std::string_view::npos ? filter.substr(comma_pos + 1) : "";
10623-
}
10624-
return false;
10625-
}
10626-
10627-
// Covers padded rows, sinks, kvpad, multi-SIMDgroup reduction, quantized K/V, and MLA views.
10628-
// The override is backend-global, so this runs after all parallel workers have joined.
10629-
static bool run_fa_vec_slice(ggml_backend_t backend, ggml_backend_t backend_cpu, const char * op_names_filter) {
10630-
if (!op_names_filter_selects(op_names_filter, "FLASH_ATTN_EXT")) {
10631-
return true;
10632-
}
10633-
10634-
auto * reg = ggml_backend_dev_backend_reg(ggml_backend_get_device(backend));
10635-
10636-
auto set_ov = (set_fa_vec_override_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_metal_tuning_set_fa_vec_override");
10637-
auto clear_ov = (clear_fa_vec_override_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_metal_tuning_clear_fa_vec_override");
10638-
if (!set_ov || !clear_ov) {
10639-
return true; // not the Metal backend: nothing to force
10640-
}
10641-
10642-
struct shape_t { int dk, dv; };
10643-
const shape_t shapes[] = { { 128, 128 }, { 576, 512 } }; // mainstream head size + MLA shared K/V view
10644-
const int ne01_pts[] = { 1, 3 }; // decode, and padded rows for Q=2 and Q=4
10645-
const int ne11_pts[] = { 512, 4097 }; // nsg=1, and nsg>=2 together with kvpad
10646-
const ggml_type types[] = { GGML_TYPE_F16, GGML_TYPE_Q4_0 };
10647-
10648-
int n_run = 0, n_fail = 0;
10649-
for (auto s : shapes) {
10650-
for (int ne : fa_vec_legal_ne(s.dk, s.dv)) {
10651-
for (int Q : { 1, 2, 4 }) {
10652-
for (ggml_type type_kv : types) {
10653-
for (bool sinks : { false, true }) {
10654-
for (int ne01 : ne01_pts) {
10655-
for (int ne11 : ne11_pts) {
10656-
set_ov(Q, ne);
10657-
test_flash_attn_ext tc(s.dk, s.dv, /*nh=*/4, { 1, 1 }, /*kv=*/ne11, /*nb=*/ne01,
10658-
/*mask=*/true, sinks, 0.0f, 0.0f, GGML_PREC_F32,
10659-
type_kv, type_kv);
10660-
auto st = tc.eval(backend, backend_cpu, "FLASH_ATTN_EXT", nullptr);
10661-
clear_ov();
10662-
10663-
if (st == test_status_t::FAIL) {
10664-
printf(" FAIL fa_vec slice: dk=%d dv=%d Q=%d ne=%d type=%s ne01=%d ne11=%d sinks=%d\n",
10665-
s.dk, s.dv, Q, ne, ggml_type_name(type_kv), ne01, ne11, (int) sinks);
10666-
n_fail++;
10667-
}
10668-
n_run++;
10669-
}
10670-
}
10671-
}
10672-
}
10673-
}
10674-
}
10675-
}
10676-
10677-
printf(" fa_vec (Q,NE) slice: %d cases run, %d failed\n", n_run, n_fail);
10678-
10679-
return n_fail == 0;
10680-
}
10681-
1068210590
static bool test_backend(ggml_backend_t backend, ggml_backend_dev_t dev, test_mode mode, const char * op_names_filter, const char * params_filter,
1068310591
printer * output_printer, const char * test_file_path, int parallel_workers) {
1068410592
auto filter_test_cases = [](std::vector<std::unique_ptr<test_case>> & test_cases, const char * params_filter) {
@@ -10816,9 +10724,7 @@ static bool test_backend(ggml_backend_t backend, ggml_backend_dev_t dev, test_mo
1081610724
output_printer->print_summary(test_summary_info(n_ok, tests_run, false));
1081710725
output_printer->print_failed_tests(failed_tests);
1081810726

10819-
const bool slice_ok = run_fa_vec_slice(backend, backend_cpu.get(), op_names_filter);
10820-
10821-
return n_ok == tests_run && slice_ok;
10727+
return n_ok == tests_run;
1082210728
}
1082310729

1082410730
if (mode == MODE_GRAD) {

0 commit comments

Comments
 (0)