@@ -2469,8 +2469,13 @@ 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 && strcmp (ggml_backend_reg_name (reg), " WebGPU" ) == 0 ) {
2473- return std::max (test_case::max_nmse_err (backend), 2e-7 );
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+ }
24742479 }
24752480 return test_case::max_nmse_err (backend);
24762481 }
@@ -4120,9 +4125,10 @@ struct test_ssm_scan : public test_case {
41204125 const int64_t n_seqs;
41214126 const bool xbc_overlap;
41224127 const int64_t K;
4128+ const bool weak_decay;
41234129
41244130 std::string vars () override {
4125- return VARS_TO_STR9 (type, d_state, head_dim, n_head, n_group, n_seq_tokens, n_seqs, xbc_overlap, K);
4131+ return VARS_TO_STR10 (type, d_state, head_dim, n_head, n_group, n_seq_tokens, n_seqs, xbc_overlap, K, weak_decay );
41264132 }
41274133
41284134 test_ssm_scan (ggml_type type = GGML_TYPE_F32 ,
@@ -4133,8 +4139,9 @@ struct test_ssm_scan : public test_case {
41334139 int64_t n_seq_tokens = 32 ,
41344140 int64_t n_seqs = 32 ,
41354141 bool xbc_overlap = false ,
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) {}
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) {}
41384145
41394146 double max_nmse_err () override {
41404147 // SSD path (head_dim > 1) uses FP16 intermediates (M matrix, X_dt); Mamba-1 is pure FP32.
@@ -4187,7 +4194,7 @@ struct test_ssm_scan : public test_case {
41874194 continue ;
41884195 } else if (t->ne [1 ] == n_head && t->ne [2 ] == 1 ) {
41894196 // A {1 or d_state, n_head}: negative decay (2-D tensor, ne[2]==1 distinguishes from 3-D/4-D tensors)
4190- init_tensor_uniform (t, - 1 .0f , -0 .5f );
4197+ init_tensor_uniform (t, weak_decay ? - 0 . 02f : - 1 .0f , weak_decay ? - 0 . 005f : -0 .5f );
41914198 } else {
41924199 init_tensor_uniform (t);
41934200 }
@@ -9111,6 +9118,10 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
91119118 test_cases.emplace_back (new test_ssm_scan (GGML_TYPE_F32 , 128 , 64 , 16 , 2 , 4 , 2 , false , /* K=*/ 4 )); // Mamba-2 rollback snapshots
91129119 test_cases.emplace_back (new test_ssm_scan (GGML_TYPE_F32 , 128 , 64 , 16 , 2 , 8 , 2 , false , /* K=*/ 3 )); // Mamba-2 rollback overflow
91139120 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
91149125
91159126 test_cases.emplace_back (new test_rwkv_wkv6 (GGML_TYPE_F32 , 32 , 64 , 1 , 1 ));
91169127 test_cases.emplace_back (new test_rwkv_wkv6 (GGML_TYPE_F32 , 32 , 64 , 32 , 1 ));
@@ -10573,6 +10584,101 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_from_file(const c
1057310584 return test_cases;
1057410585}
1057510586
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+
1057610682static bool test_backend (ggml_backend_t backend, ggml_backend_dev_t dev, test_mode mode, const char * op_names_filter, const char * params_filter,
1057710683 printer * output_printer, const char * test_file_path, int parallel_workers) {
1057810684 auto filter_test_cases = [](std::vector<std::unique_ptr<test_case>> & test_cases, const char * params_filter) {
@@ -10710,7 +10816,9 @@ static bool test_backend(ggml_backend_t backend, ggml_backend_dev_t dev, test_mo
1071010816 output_printer->print_summary (test_summary_info (n_ok, tests_run, false ));
1071110817 output_printer->print_failed_tests (failed_tests);
1071210818
10713- return n_ok == tests_run;
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;
1071410822 }
1071510823
1071610824 if (mode == MODE_GRAD ) {
0 commit comments