| 430 | } |
| 431 | |
| 432 | bool test_stft_kernel_correctness() { |
| 433 | const size_t N = 2, C_in = 1, L = 8, K = 4, stride = 2, num_fft_bins = 2; |
| 434 | const size_t C_out = 2 * num_fft_bins; |
| 435 | const size_t out_len = (L - K) / stride + 1; |
| 436 | |
| 437 | const __fp16 bin0_re[] = {(__fp16) 1, (__fp16) 1, (__fp16) 1, (__fp16) 1}; |
| 438 | const __fp16 bin1_re[] = {(__fp16) 1, (__fp16) 0, (__fp16)-1, (__fp16) 0}; |
| 439 | const __fp16 bin0_im[] = {(__fp16) 0, (__fp16) 0, (__fp16) 0, (__fp16) 0}; |
| 440 | const __fp16 bin1_im[] = {(__fp16) 0, (__fp16)-1, (__fp16) 0, (__fp16) 1}; |
| 441 | std::vector<__fp16> weight; |
| 442 | for (const __fp16* row : {bin0_re, bin1_re, bin0_im, bin1_im}) |
| 443 | weight.insert(weight.end(), row, row + K); |
| 444 | |
| 445 | const __fp16 ramp[] = {(__fp16)1, (__fp16)2, (__fp16)3, (__fp16)4, |
| 446 | (__fp16)5, (__fp16)6, (__fp16)7, (__fp16)8}; |
| 447 | const __fp16 cosine[] = {(__fp16)0, (__fp16)1, (__fp16) 0, (__fp16)-1, |
| 448 | (__fp16)0, (__fp16)1, (__fp16) 0, (__fp16)-1}; |
| 449 | std::vector<__fp16> input; |
| 450 | input.insert(input.end(), ramp, ramp + L); |
| 451 | input.insert(input.end(), cosine, cosine + L); |
| 452 | |
| 453 | struct Cplx { float r, i; }; |
| 454 | const Cplx expected[2][2][3] = { |
| 455 | { { {10,0},{18,0},{26,0} }, { {-2,2},{-2,2},{-2,2} } }, |
| 456 | { { { 0,0},{ 0,0},{ 0,0} }, { { 0,-2},{ 0,2},{ 0,-2} } }, |
| 457 | }; |
| 458 | |
| 459 | std::vector<__fp16> cplx(N * C_out * out_len, (__fp16)0); |
| 460 | cactus_stft_f16(input.data(), weight.data(), cplx.data(), |
| 461 | N, L, C_in, C_out, K, stride, num_fft_bins); |
| 462 | |
| 463 | const size_t out_bs = C_out * out_len; |
| 464 | const float tol = 0.1f; |
| 465 | for (size_t n = 0; n < N; ++n) { |
| 466 | for (size_t b = 0; b < num_fft_bins; ++b) { |
| 467 | for (size_t t = 0; t < out_len; ++t) { |
| 468 | float r = (float)cplx[n * out_bs + b * out_len + t]; |
| 469 | float im = (float)cplx[n * out_bs + (b + num_fft_bins) * out_len + t]; |
| 470 | if (std::abs(r - expected[n][b][t].r) > tol) return false; |
| 471 | if (std::abs(im - expected[n][b][t].i) > tol) return false; |
| 472 | } |
| 473 | } |
| 474 | } |
| 475 | |
| 476 | return true; |
| 477 | } |
| 478 | |
| 479 | bool test_fast_tanh_f32x4_correctness() { |
| 480 | constexpr float TOL = 1e-5f; |
no test coverage detected