MCPcopy Create free account
hub / github.com/cactus-compute/cactus / test_stft_kernel_correctness

Function test_stft_kernel_correctness

tests/test_kernel.cpp:432–477  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

430}
431
432bool 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
479bool test_fast_tanh_f32x4_correctness() {
480 constexpr float TOL = 1e-5f;

Callers 1

mainFunction · 0.85

Calls 2

cactus_stft_f16Function · 0.85
dataMethod · 0.80

Tested by

no test coverage detected