yum/gpu_fft
A GPU-friendly FFT
git clone https://git.yummers.dev/yum/gpu_fft
3c2c603
master
1use num:: complex:: Complex ; 2use num:: traits::{ Float , FloatConst }; 3 4fn is_power_of_k ( n : usize , k : usize ) ->bool { 5match n{ 60 =>false , 71 =>true , 8 _ => n % k ==0 &&is_power_of_k ( n / k, k), 9} 10} 11 12fn usize_to_float < T : Float >( value : usize ) ->T { 13 num:: cast ( value). unwrap () 14} 15 16# [ rustfmt :: skip ] 17# [ allow ( dead_code )] 18pub fn naive_dft < T : Float +FloatConst >( data : & mut [ Complex < T >]) { 19let big_n = data. len (); 20let mut result =vec! [ Complex :: new ( T :: zero (), T :: zero ()); big_n]; 21for kin 0 ..big_n{ 22for nin 0 ..big_n{ 23let k_t =usize_to_float ::< T >( k); 24let n_t =usize_to_float ::< T >( n); 25let big_n_t =usize_to_float ::< T >( big_n); 26let phase = -T :: TAU () * k_t* n_t / big_n_t; 27let factor =Complex ::< T >:: cis ( phase); 28 result[ k] = result[ k] + data[ n] * factor; 29} 30} 31 data. copy_from_slice ( & result); 32} 33 34// Helper to naive_fft. Takes `data` along with 3 numbers that let us recreate an even-odd subset: 35// - `start_idx` tells us where the subset begins; 36// - `big_n` is the number of elements in the subset; 37// - `stride` is the distance between elements. 38// We also use a double buffer, `scratch`, to avoid clobbering data while merging results. 39# [ rustfmt :: skip ] 40fn _naive_fft < T : Float +FloatConst >( data : & mut [ Complex < T >], start_idx : usize , big_n : usize , stride : usize , scratch : & mut [ Complex < T >]) { 41if big_n ==1 { 42return ; 43} 44// Compute DFT of even elements. 45_naive_fft ( data, start_idx, big_n/2 , stride* 2 , scratch); 46// Odd elements. 47_naive_fft ( data, start_idx+stride, big_n/2 , stride* 2 , scratch); 48for kin 0 ..( big_n/2 ) { 49let p = data[ start_idx +2 * k* stride]; 50let q = data[ start_idx +( 2 * k +1 ) * stride]; 51let k_t =usize_to_float ::< T >( k); 52let big_n_t =usize_to_float ::< T >( big_n); 53let phase = -T :: TAU () * k_t / big_n_t; 54let factor =Complex ::< T >:: cis ( phase); 55 scratch[ start_idx + k* stride] = p + q* factor; 56 scratch[ start_idx +( k + big_n /2 ) * stride] = p - q* factor; 57} 58 data. copy_from_slice ( scratch); 59} 60 61// Naive implementation of Cooley-Tukey FFT. Modifies `data`in place. Panics if data.len() is not a power of two. 62# [ allow ( dead_code )] 63pub fn naive_fft < T : Float +FloatConst >( data : & mut [ Complex < T >]) { 64assert! ( is_power_of_k ( data. len (), 2 )); 65let mut scratch =Vec :: from ( data. as_ref ()); 66_naive_fft ( data, 0 , data. len (), 1 , & mut scratch); 67} 68 69fn _fft_v1_hoist < T : Float +FloatConst >( 70data : & mut [ Complex < T >], 71start_idx : usize , 72big_n : usize , 73stride : usize , 74scratch : & mut [ Complex < T >], 75twiddles : & [ Complex < T >], 76) { 77if big_n ==1 { 78return ; 79} 80// Compute DFT of even elements. 81_fft_v1_hoist ( data, start_idx, big_n /2 , stride* 2 , scratch, twiddles); 82// Odd elements. 83_fft_v1_hoist ( 84 data, 85 start_idx + stride, 86 big_n /2 , 87 stride* 2 , 88 scratch, 89 twiddles, 90); 91for kin 0 ..( big_n /2 ) { 92let p = data[ start_idx +2 * k* stride]; 93let q = data[ start_idx +( 2 * k +1 ) * stride]; 94let factor = twiddles[ k* stride]; 95 scratch[ start_idx + k* stride] = p + q* factor; 96 scratch[ start_idx +( k + big_n /2 ) * stride] = p - q* factor; 97} 98 data. copy_from_slice ( scratch); 99} 100 101// Modification of fft_naive: hoist out and precompute twiddles. 102pub fn fft_v1_hoist < T : Float +FloatConst >( data : & mut [ Complex < T >], twiddles : & [ Complex < T >]) { 103assert! ( is_power_of_k ( data. len (), 2 )); 104let mut scratch =Vec :: from ( data. as_ref ()); 105_fft_v1_hoist ( data, 0 , data. len (), 1 , & mut scratch, & twiddles); 106} 107 108# [ inline ( always )] 109fn _fft_v2_double_buffer < T : Float +FloatConst >( 110src : & mut [ Complex < T >], 111dst : & mut [ Complex < T >], 112start_idx : usize , 113big_n : usize , 114stride : usize , 115twiddles : & [ Complex < T >], 116) { 117if big_n ==1 { 118return ; 119} 120// Compute DFT of even elements. 121_fft_v2_double_buffer ( dst, src, start_idx, big_n /2 , stride* 2 , twiddles); 122// Odd elements. 123_fft_v2_double_buffer ( 124 dst, 125 src, 126 start_idx + stride, 127 big_n /2 , 128 stride* 2 , 129 twiddles, 130); 131for kin 0 ..( big_n /2 ) { 132let p = src[ start_idx +2 * k* stride]; 133let q = src[ start_idx +( 2 * k +1 ) * stride]; 134let factor = twiddles[ k* stride]; 135 dst[ start_idx + k* stride] = p + q* factor; 136 dst[ start_idx +( k + big_n /2 ) * stride] = p - q* factor; 137} 138} 139 140pub fn fft_v2_double_buffer < T : Float +FloatConst >( 141src : & mut [ Complex < T >], 142dst : & mut [ Complex < T >], 143twiddles : & [ Complex < T >], 144) { 145assert! ( is_power_of_k ( src. len (), 2 )); 146 dst. copy_from_slice ( src); 147// Switching `src` and `dst` means that at the end, the result is in `src` - which is actually 148// what we want! We will be hiding `dst` and `twiddles` in a struct later on :) 149_fft_v2_double_buffer ( dst, src, 0 , src. len (), 1 , twiddles); 150} 151 152// Evaluates the base-`k` logarithm of `n`. 153// Precondition: `n` is a power of `k`. 154const fn log_k_of < const K : usize >( mut n : usize ) ->usize { 155let mut res =0 ; 156while n >1 { 157 n /=K ; 158 res +=1 ; 159} 160 res 161} 162 163pub fn fft_v3_iterative < T : Float +FloatConst >( 164src : & mut [ Complex < T >], 165dst : & mut [ Complex < T >], 166twiddles : & [ Complex < T >], 167) { 168assert! ( is_power_of_k ( src. len (), 2 )); 169let n_iter =log_k_of ::< 2 >( src. len ()); 170 171 dst. copy_from_slice ( src); 172 173let ( mut input, mut output) =if n_iter %2 ==0 { 174( dst, src) 175} else { 176( src, dst) 177}; 178let mut stride = input. len (); 179let mut big_n =1 ; 180for _in 0 ..n_iter{ 181 stride /=2 ; 182 big_n *=2 ; 183 std:: mem:: swap ( & mut input, & mut output); 184 185for start_idxin 0 ..stride{ 186for kin 0 ..big_n /2 { 187// Get odd and even elements. 188let p = input[ start_idx +2 * k* stride]; 189let q = input[ start_idx +( 2 * k +1 ) * stride]; 190// Combine. 191let factor = twiddles[ k* stride]; 192 output[ start_idx + k* stride] = p + q* factor; 193 output[ start_idx +( k + big_n /2 ) * stride] = p - q* factor; 194} 195} 196} 197} 198 199# [ inline ( always )] 200fn mul_ni < T : Float +FloatConst >( x : Complex < T >) ->Complex < T > { 201Complex :: new ( x. im , -x. re ) 202} 203 204fn fft_butterfly_radix_4 < T : Float +FloatConst >( 205input : & mut [ Complex < T >], 206output : & mut [ Complex < T >], 207stride : usize , 208big_n : usize , 209twiddles : & [ Complex < T >], 210) { 211for start_idxin 0 ..stride{ 212for kin 0 ..big_n /4 { 213// Collect inputs. 214let i0 = input[ start_idx +4 * k* stride]; 215let i1 = input[ start_idx +( 4 * k +1 ) * stride]; 216let i2 = input[ start_idx +( 4 * k +2 ) * stride]; 217let i3 = input[ start_idx +( 4 * k +3 ) * stride]; 218// Collect relevant twiddles. 219let ot1 = twiddles[ 1 * k* stride]; 220let ot2 = twiddles[ 2 * k* stride]; 221let ot3 = twiddles[ 3 * k* stride]; 222 223let a = i0; 224let b = ot1* i1; 225let c = ot2* i2; 226let d = ot3* i3; 227 228// To derive this, write the output assignments in terms of 229// a/b/c/d, then factor out! 230let ac_sum = a + c; 231let ac_diff = a - c; 232let bd_sum = b + d; 233let bd_diff_ni =mul_ni ( b - d); 234 235 output[ start_idx + k* stride] = ac_sum + bd_sum; 236 output[ start_idx +( k + big_n /4 ) * stride] = ac_diff + bd_diff_ni; 237 output[ start_idx +( k + big_n /2 ) * stride] = ac_sum - bd_sum; 238 output[ start_idx +( k +3 * big_n /4 ) * stride] = ac_diff - bd_diff_ni; 239} 240} 241} 242 243fn fft_butterfly_radix_4_s0 < T : Float +FloatConst >( 244input : & mut [ Complex < T >], 245output : & mut [ Complex < T >], 246) { 247let stride = input. len () /4 ; 248let big_n =4 ; 249 250for start_idxin 0 ..stride{ 251for kin 0 ..big_n /4 { 252// Collect inputs. 253let i0 = input[ start_idx +4 * k* stride]; 254let i1 = input[ start_idx +( 4 * k +1 ) * stride]; 255let i2 = input[ start_idx +( 4 * k +2 ) * stride]; 256let i3 = input[ start_idx +( 4 * k +3 ) * stride]; 257 258let a = i0; 259let b = i1; 260let c = i2; 261let d = i3; 262 263// To derive this, write the output assignments in terms of 264// a/b/c/d, then factor out! 265let ac_sum = a + c; 266let ac_diff = a - c; 267let bd_sum = b + d; 268let bd_diff_ni =mul_ni ( b - d); 269 270 output[ start_idx + k* stride] = ac_sum + bd_sum; 271 output[ start_idx +( k + big_n /4 ) * stride] = ac_diff + bd_diff_ni; 272 output[ start_idx +( k + big_n /2 ) * stride] = ac_sum - bd_sum; 273 output[ start_idx +( k +3 * big_n /4 ) * stride] = ac_diff - bd_diff_ni; 274} 275} 276} 277 278pub fn fft_v4_radix_4 < T : Float +FloatConst >( 279src : & mut [ Complex < T >], 280dst : & mut [ Complex < T >], 281twiddles : & [ Complex < T >], 282) { 283assert! ( is_power_of_k ( src. len (), 4 )); 284let n_iter =log_k_of ::< 4 >( src. len ()); 285 286 dst. copy_from_slice ( src); 287 288let ( mut input, mut output) =if n_iter %2 ==0 { 289( dst, src) 290} else { 291( src, dst) 292}; 293let big_n = input. len (); 294let mut stride = big_n; 295let mut big_n =1 ; 296for _in 0 ..n_iter{ 297 stride /=4 ; 298 big_n *=4 ; 299 std:: mem:: swap ( & mut input, & mut output); 300 301fft_butterfly_radix_4 ( input, output, stride, big_n, twiddles); 302} 303} 304 305pub fn fft_v5_s0_opt < T : Float +FloatConst >( 306src : & mut [ Complex < T >], 307dst : & mut [ Complex < T >], 308twiddles : & [ Complex < T >], 309) { 310assert! ( is_power_of_k ( src. len (), 4 )); 311let n_iter =log_k_of ::< 4 >( src. len ()); 312 313 dst. copy_from_slice ( src); 314 315let ( mut input, mut output) =if n_iter %2 ==0 { 316( dst, src) 317} else { 318( src, dst) 319}; 320let big_n = input. len (); 321let mut stride = big_n; 322let mut big_n =1 ; 323for stagein 0 ..n_iter{ 324 stride /=4 ; 325 big_n *=4 ; 326 std:: mem:: swap ( & mut input, & mut output); 327 328if stage ==0 { 329fft_butterfly_radix_4_s0 ( input, output); 330} else { 331fft_butterfly_radix_4 ( input, output, stride, big_n, twiddles); 332} 333} 334} 335 336fn fft_butterfly_radix_4_unsafe < T : Float +FloatConst >( 337input : & mut [ Complex < T >], 338output : & mut [ Complex < T >], 339stride : usize , 340big_n : usize , 341twiddles : & [ Complex < T >], 342) { 343let input_ptr = input. as_ptr (); 344let output_ptr = output. as_mut_ptr (); 345for start_idxin 0 ..stride{ 346for kin 0 ..big_n /4 { 347unsafe { 348// Collect inputs. 349let i0 =* input_ptr. add ( start_idx +4 * k* stride); 350let i1 =* input_ptr. add ( start_idx +( 4 * k +1 ) * stride); 351let i2 =* input_ptr. add ( start_idx +( 4 * k +2 ) * stride); 352let i3 =* input_ptr. add ( start_idx +( 4 * k +3 ) * stride); 353// Collect relevant twiddles. 354let ot1 = twiddles. get_unchecked ( 1 * k* stride); 355let ot2 = twiddles. get_unchecked ( 2 * k* stride); 356let ot3 = twiddles. get_unchecked ( 3 * k* stride); 357 358let a = i0; 359let b = ot1* i1; 360let c = ot2* i2; 361let d = ot3* i3; 362 363// To derive this, write the output assignments in terms of 364// a/b/c/d, then factor out! 365let ac_sum = a + c; 366let ac_diff = a - c; 367let bd_sum = b + d; 368let bd_diff_ni =mul_ni ( b - d); 369 370* output_ptr. add ( start_idx + k* stride) = ac_sum + bd_sum; 371* output_ptr. add ( start_idx +( k + big_n /4 ) * stride) = ac_diff + bd_diff_ni; 372* output_ptr. add ( start_idx +( k + big_n /2 ) * stride) = ac_sum - bd_sum; 373* output_ptr. add ( start_idx +( k +3 * big_n /4 ) * stride) = ac_diff - bd_diff_ni; 374} 375} 376} 377} 378 379fn fft_butterfly_radix_4_s0_unsafe < T : Float +FloatConst >( 380input : & mut [ Complex < T >], 381output : & mut [ Complex < T >], 382) { 383let stride = input. len () /4 ; 384let big_n =4 ; 385let input_ptr = input. as_ptr (); 386let output_ptr = output. as_mut_ptr (); 387for start_idxin 0 ..stride{ 388for kin 0 ..big_n /4 { 389unsafe { 390// Collect inputs. 391let i0 =* input_ptr. add ( start_idx +4 * k* stride); 392let i1 =* input_ptr. add ( start_idx +( 4 * k +1 ) * stride); 393let i2 =* input_ptr. add ( start_idx +( 4 * k +2 ) * stride); 394let i3 =* input_ptr. add ( start_idx +( 4 * k +3 ) * stride); 395 396let a = i0; 397let b = i1; 398let c = i2; 399let d = i3; 400 401// To derive this, write the output assignments in terms of 402// a/b/c/d, then factor out! 403let ac_sum = a + c; 404let ac_diff = a - c; 405let bd_sum = b + d; 406let bd_diff_ni =mul_ni ( b - d); 407 408* output_ptr. add ( start_idx + k* stride) = ac_sum + bd_sum; 409* output_ptr. add ( start_idx +( k + big_n /4 ) * stride) = ac_diff + bd_diff_ni; 410* output_ptr. add ( start_idx +( k + big_n /2 ) * stride) = ac_sum - bd_sum; 411* output_ptr. add ( start_idx +( k +3 * big_n /4 ) * stride) = ac_diff - bd_diff_ni; 412} 413} 414} 415} 416 417pub fn fft_v6_unsafe < T : Float +FloatConst >( 418src : & mut [ Complex < T >], 419dst : & mut [ Complex < T >], 420twiddles : & [ Complex < T >], 421) { 422assert! ( is_power_of_k ( src. len (), 4 )); 423assert_eq! ( src. len (), dst. len ()); 424assert_eq! ( twiddles. len (), src. len ()); 425let n_iter =log_k_of ::< 4 >( src. len ()); 426 427 dst. copy_from_slice ( src); 428 429let ( mut input, mut output) =if n_iter %2 ==0 { 430( dst, src) 431} else { 432( src, dst) 433}; 434let big_n = input. len (); 435let mut stride = big_n; 436let mut big_n =1 ; 437for stagein 0 ..n_iter{ 438 stride /=4 ; 439 big_n *=4 ; 440 std:: mem:: swap ( & mut input, & mut output); 441 442if stage ==0 { 443fft_butterfly_radix_4_s0_unsafe ( input, output); 444} else { 445fft_butterfly_radix_4_unsafe ( input, output, stride, big_n, twiddles); 446} 447} 448} 449 450# [ inline ( always )] 451fn rot_45 < T : Float +FloatConst >( c : Complex < T >) ->Complex < T > { 452let s =T :: FRAC_1_SQRT_2 (); 453// The standard 2D rotation matrix gives: 454// [ cos(pi/4) -sin(pi/4)] [ s -s ] 455// [ sin(pi/4) cos(pi/4)] = [ s s ] 456Complex ::< T >:: new ( c. re - c. im , c. re + c. im ) * s 457} 458 459# [ inline ( always )] 460fn rot_90 < T : Float +FloatConst >( c : Complex < T >) ->Complex < T > { 461// The standard 2D rotation matrix gives: 462// [ cos(pi/2) -sin(pi/2)] [ 0 -1 ] 463// [ sin(pi/2) cos(pi/2)] = [ 1 0 ] 464Complex ::< T >:: new ( -c. im , c. re ) 465} 466 467# [ inline ( always )] 468fn rot_180 < T : Float +FloatConst >( c : Complex < T >) ->Complex < T > { 469// The standard 2D rotation matrix gives: 470// [ cos(pi) -sin(pi)] [ -1 0 ] 471// [ sin(pi) cos(pi)] = [ 0 -1 ] 472Complex ::< T >:: new ( -c. re , -c. im ) 473} 474 475# [ inline ( always )] 476fn rot_270 < T : Float +FloatConst >( c : Complex < T >) ->Complex < T > { 477// The standard 2D rotation matrix gives: 478// [ cos(3pi/2) -sin(3pi/2)] [ 0 1 ] 479// [ sin(3pi/2) cos(3pi/2)] = [ -1 0 ] 480Complex ::< T >:: new ( c. im , -c. re ) 481} 482 483fn fft_butterfly_radix_8_unsafe < T : Float +FloatConst >( 484input : & mut [ Complex < T >], 485output : & mut [ Complex < T >], 486stride : usize , 487big_n : usize , 488twiddles : & [ Complex < T >], 489) { 490let input_ptr = input. as_ptr (); 491let output_ptr = output. as_mut_ptr (); 492for start_idxin 0 ..stride{ 493for kin 0 ..big_n /8 { 494unsafe { 495// Collect inputs. 496let i0 =* input_ptr. add ( start_idx +8 * k* stride); 497let i1 =* input_ptr. add ( start_idx +( 8 * k +1 ) * stride); 498let i2 =* input_ptr. add ( start_idx +( 8 * k +2 ) * stride); 499let i3 =* input_ptr. add ( start_idx +( 8 * k +3 ) * stride); 500let i4 =* input_ptr. add ( start_idx +( 8 * k +4 ) * stride); 501let i5 =* input_ptr. add ( start_idx +( 8 * k +5 ) * stride); 502let i6 =* input_ptr. add ( start_idx +( 8 * k +6 ) * stride); 503let i7 =* input_ptr. add ( start_idx +( 8 * k +7 ) * stride); 504 505// Collect relevant twiddles. 506let ot1 = twiddles. get_unchecked ( 1 * k* stride); 507let ot2 = twiddles. get_unchecked ( 2 * k* stride); 508let ot3 = twiddles. get_unchecked ( 3 * k* stride); 509let ot4 = twiddles. get_unchecked ( 4 * k* stride); 510let ot5 = twiddles. get_unchecked ( 5 * k* stride); 511let ot6 = twiddles. get_unchecked ( 6 * k* stride); 512let ot7 = twiddles. get_unchecked ( 7 * k* stride); 513 514let a = i0; 515let b = ot1* i1; 516let c = ot2* i2; 517let d = ot3* i3; 518let e = ot4* i4; 519let f = ot5* i5; 520let g = ot6* i6; 521let h = ot7* i7; 522 523let ae_sum = a + e; 524let ae_diff = a - e; 525let bf_sum = b + f; 526let bf_diff = b - f; 527let cg_sum = c + g; 528let cg_diff = c - g; 529let dh_sum = d + h; 530let dh_diff = d - h; 531 532let w00 = ae_sum + cg_sum; 533let w01 = ae_sum - cg_sum; 534let w10 = ae_diff +rot_270 ( cg_diff); 535let w11 = ae_diff -rot_270 ( cg_diff); 536let x00 = bf_sum + dh_sum; 537let x01 =rot_270 ( bf_sum) +rot_90 ( dh_sum); 538let x10 =rot_45 ( rot_270 ( bf_diff) +rot_180 ( dh_diff)); 539let x11 =rot_45 ( rot_180 ( bf_diff) +rot_270 ( dh_diff)); 540 541* output_ptr. add ( start_idx + k* stride) = w00 + x00; 542* output_ptr. add ( start_idx +( k + big_n /8 ) * stride) = w10 + x10; 543* output_ptr. add ( start_idx +( k + big_n /4 ) * stride) = w01 + x01; 544* output_ptr. add ( start_idx +( k +3 * big_n /8 ) * stride) = w11 + x11; 545* output_ptr. add ( start_idx +( k + big_n /2 ) * stride) = w00 - x00; 546* output_ptr. add ( start_idx +( k +5 * big_n /8 ) * stride) = w10 - x10; 547* output_ptr. add ( start_idx +( k +3 * big_n /4 ) * stride) = w01 - x01; 548* output_ptr. add ( start_idx +( k +7 * big_n /8 ) * stride) = w11 - x11; 549} 550} 551} 552} 553 554fn fft_butterfly_radix_8_s0_unsafe < T : Float +FloatConst >( 555input : & mut [ Complex < T >], 556output : & mut [ Complex < T >], 557) { 558let stride = input. len () /8 ; 559let big_n =8 ; 560let input_ptr = input. as_ptr (); 561let output_ptr = output. as_mut_ptr (); 562for start_idxin 0 ..stride{ 563for kin 0 ..big_n /8 { 564unsafe { 565// Collect inputs. 566let i0 =* input_ptr. add ( start_idx +8 * k* stride); 567let i1 =* input_ptr. add ( start_idx +( 8 * k +1 ) * stride); 568let i2 =* input_ptr. add ( start_idx +( 8 * k +2 ) * stride); 569let i3 =* input_ptr. add ( start_idx +( 8 * k +3 ) * stride); 570let i4 =* input_ptr. add ( start_idx +( 8 * k +4 ) * stride); 571let i5 =* input_ptr. add ( start_idx +( 8 * k +5 ) * stride); 572let i6 =* input_ptr. add ( start_idx +( 8 * k +6 ) * stride); 573let i7 =* input_ptr. add ( start_idx +( 8 * k +7 ) * stride); 574 575let a = i0; 576let b = i1; 577let c = i2; 578let d = i3; 579let e = i4; 580let f = i5; 581let g = i6; 582let h = i7; 583 584let ae_sum = a + e; 585let ae_diff = a - e; 586let bf_sum = b + f; 587let bf_diff = b - f; 588let cg_sum = c + g; 589let cg_diff = c - g; 590let dh_sum = d + h; 591let dh_diff = d - h; 592 593let w00 = ae_sum + cg_sum; 594let w01 = ae_sum - cg_sum; 595let w10 = ae_diff +rot_270 ( cg_diff); 596let w11 = ae_diff -rot_270 ( cg_diff); 597let x00 = bf_sum + dh_sum; 598let x01 =rot_270 ( bf_sum) +rot_90 ( dh_sum); 599let x10 =rot_45 ( rot_270 ( bf_diff) +rot_180 ( dh_diff)); 600let x11 =rot_45 ( rot_180 ( bf_diff) +rot_270 ( dh_diff)); 601 602* output_ptr. add ( start_idx + k* stride) = w00 + x00; 603* output_ptr. add ( start_idx +( k + big_n /8 ) * stride) = w10 + x10; 604* output_ptr. add ( start_idx +( k + big_n /4 ) * stride) = w01 + x01; 605* output_ptr. add ( start_idx +( k +3 * big_n /8 ) * stride) = w11 + x11; 606* output_ptr. add ( start_idx +( k + big_n /2 ) * stride) = w00 - x00; 607* output_ptr. add ( start_idx +( k +5 * big_n /8 ) * stride) = w10 - x10; 608* output_ptr. add ( start_idx +( k +3 * big_n /4 ) * stride) = w01 - x01; 609* output_ptr. add ( start_idx +( k +7 * big_n /8 ) * stride) = w11 - x11; 610} 611} 612} 613} 614 615pub fn fft_v7_radix_8 < T : Float +FloatConst >( 616src : & mut [ Complex < T >], 617dst : & mut [ Complex < T >], 618twiddles : & [ Complex < T >], 619) { 620assert! ( is_power_of_k ( src. len (), 8 )); 621assert_eq! ( src. len (), dst. len ()); 622assert_eq! ( twiddles. len (), src. len ()); 623let n_iter =log_k_of ::< 8 >( src. len ()); 624 625 dst. copy_from_slice ( src); 626 627let ( mut input, mut output) =if n_iter %2 ==0 { 628( dst, src) 629} else { 630( src, dst) 631}; 632let big_n = input. len (); 633let mut stride = big_n; 634let mut big_n =1 ; 635for stagein 0 ..n_iter{ 636 stride /=8 ; 637 big_n *=8 ; 638 std:: mem:: swap ( & mut input, & mut output); 639 640if stage ==0 { 641fft_butterfly_radix_8_s0_unsafe ( input, output); 642} else { 643fft_butterfly_radix_8_unsafe ( input, output, stride, big_n, twiddles); 644} 645} 646} 647 648// Calculates the "twiddle factors" for an n-element FFT, aka all of the nth roots of unity. 649pub fn precompute_twiddles < T : Float +FloatConst >( n : usize ) ->Vec < Complex < T >> { 650let mut result =vec! [ Complex ::< T >:: new ( T :: zero (), T :: zero ()); n]; 651 652let n_f64 =usize_to_float ::< f64 >( n); 653for iin 0 ..n{ 654let tw_f64 =Complex ::< f64 >:: cis ( -f64:: TAU () * usize_to_float ::< f64 >( i) /( n_f64)); 655 result[ i] =Complex :: new ( T :: from ( tw_f64. re ). unwrap (), T :: from ( tw_f64. im ). unwrap ()); 656} 657 658 result 659} 660 661// Precompute all twiddles factors for a given radix `r` and input size `n`. 662pub fn precompute_all_twiddles < T : Float +FloatConst >( r : usize , n : usize ) ->Vec < Vec < Complex < T >>> { 663assert! ( is_power_of_k ( n, r)); 664 665let mut result =Vec :: new (); 666let mut n_cur = n; 667while n_cur >1 { 668 result. push ( precompute_twiddles ( n_cur)); 669 n_cur /= r; 670} 671 result 672} 673 674# [ cfg ( test )] 675mod tests{ 676use rand:: RngExt ; 677use rand:: SeedableRng ; 678use rand:: rngs:: StdRng ; 679 680use super :: * ; 681use num:: complex:: Complex ; 682use std:: time:: Instant ; 683 684// Cast one float type to another, truncating if needed. 685fn t0_to_t1 < T0 : Float , T1 : Float >( val : T0 ) ->T1 { 686return num:: cast ( val). unwrap (); 687} 688 689// Casts one Complex<T> to another. 690fn complex_to_t < T0 : Float , T1 : Float >( data : & [ Complex < T0 >]) ->Vec < Complex < T1 >> { 691 data. iter () 692. map ( |c|Complex :: new ( t0_to_t1 ( c. re ), t0_to_t1 ( c. im ))) 693. collect () 694} 695 696fn evaluate_results ( 697result_ref_64 : & Vec < Complex < f64 >>, 698result_cur_64 : & Vec < Complex < f64 >>, 699duration_dft : std:: time:: Duration , 700duration_cur : std:: time:: Duration , 701algo_name : & str , 702) { 703assert_eq! ( result_ref_64. len (), result_cur_64. len ()); 704 705println! ( " Algorithm: {}" , algo_name); 706println! ( " Duration: {:?}" , duration_cur); 707println! ( 708" Speedup: {:?}" , 709 duration_dft. as_nanos () as f32 / duration_cur. as_nanos () as f32 710); 711 712let mut max_err: f64 =0.0 ; 713let mut sum_err: f64 =0.0 ; 714// TODO: median, 90p, 99p 715 716let result_64: Vec < Complex < f64 >> =complex_to_t ( & result_cur_64); 717 718for iin 0 ..result_ref_64. len () { 719 max_err = max_err. max (( result_ref_64[ i] - result_64[ i]). norm ()); 720 sum_err +=( result_ref_64[ i] - result_64[ i]). norm (); 721} 722println! ( " Max err: {}" , max_err); 723println! ( " Avg err: {}" , sum_err /( result_ref_64. len () as f64 )); 724} 725 726# [ test ] 727fn naive_dft_matches_dft () { 728println! ( "Testing naive DFT against naive FFT" ); 729{ 730let mut data: Vec < Complex < f64 >> =Vec :: new (); 731 data. resize (( 2 as usize ). pow ( 12 ), Complex :: ZERO ); 732// Randomize data. 733let mut rng =StdRng :: seed_from_u64 ( 42 ); 734for iin 0 ..data. len () { 735 data[ i] =Complex :: new ( rng. random (), rng. random ()); 736} 737 738let mut result_ref: Vec < rustfft:: num_complex:: Complex < f64 >> =complex_to_t ( & data); 739let mut ref_planner = rustfft:: FftPlanner ::< f64 >:: new (); 740let ref_fft = ref_planner. plan_fft_forward ( result_ref. len ()); 741 ref_fft. process ( & mut result_ref); 742let result_ref_64: Vec < Complex < f64 >> =complex_to_t ( & result_ref); 743 744println! ( "Input size {}" , data. len ()); 745 746let mut result_dft: Vec < Complex < f32 >> =complex_to_t ( & data); 747let mut begin_instant =Instant :: now (); 748naive_dft ( & mut result_dft); 749let duration_dft = begin_instant. elapsed (); 750let result_dft_64: Vec < Complex < f64 >> =complex_to_t ( & result_dft); 751evaluate_results ( 752& result_ref_64, 753& result_dft_64, 754 duration_dft, 755 duration_dft, 756"Naive DFT" , 757); 758 759{ 760let mut result_fft: Vec < Complex < f32 >> =complex_to_t ( & data); 761 begin_instant =Instant :: now (); 762naive_fft ( & mut result_fft); 763let duration_fft = begin_instant. elapsed (); 764let result_fft_64: Vec < Complex < f64 >> =complex_to_t ( & result_fft); 765evaluate_results ( 766& result_ref_64, 767& result_fft_64, 768 duration_dft, 769 duration_fft, 770"Naive FFT" , 771); 772} 773 774{ 775let mut result_fft: Vec < Complex < f32 >> =complex_to_t ( & data); 776let twiddles =precompute_twiddles ( data. len ()); 777 begin_instant =Instant :: now (); 778fft_v1_hoist ( & mut result_fft, & twiddles); 779let duration_fft = begin_instant. elapsed (); 780let result_fft_64: Vec < Complex < f64 >> =complex_to_t ( & result_fft); 781evaluate_results ( 782& result_ref_64, 783& result_fft_64, 784 duration_dft, 785 duration_fft, 786"FFT v1 (hoist twiddles)" , 787); 788} 789 790{ 791let mut result_fft: Vec < Complex < f32 >> =complex_to_t ( & data); 792let twiddles =precompute_twiddles ( data. len ()); 793let mut scratch: Vec < Complex < f32 >> =vec! [ Complex :: new ( 0.0 , 0.0 ); data. len ()]; 794 begin_instant =Instant :: now (); 795fft_v2_double_buffer ( & mut result_fft, & mut scratch, & twiddles); 796let duration_fft = begin_instant. elapsed (); 797let result_fft_64: Vec < Complex < f64 >> =complex_to_t ( & result_fft); 798evaluate_results ( 799& result_ref_64, 800& result_fft_64, 801 duration_dft, 802 duration_fft, 803"FFT v2 (double buffer)" , 804); 805} 806 807{ 808let mut result_fft: Vec < Complex < f32 >> =complex_to_t ( & data); 809let twiddles =precompute_twiddles ( data. len ()); 810let mut scratch: Vec < Complex < f32 >> =vec! [ Complex :: new ( 0.0 , 0.0 ); data. len ()]; 811 begin_instant =Instant :: now (); 812fft_v3_iterative ( & mut result_fft, & mut scratch, & twiddles); 813let duration_fft = begin_instant. elapsed (); 814let result_fft_64: Vec < Complex < f64 >> =complex_to_t ( & result_fft); 815evaluate_results ( 816& result_ref_64, 817& result_fft_64, 818 duration_dft, 819 duration_fft, 820"FFT v3 (iterative)" , 821); 822} 823 824{ 825let mut result_fft: Vec < Complex < f32 >> =complex_to_t ( & data); 826let twiddles =precompute_twiddles ( data. len ()); 827let mut scratch: Vec < Complex < f32 >> =vec! [ Complex :: new ( 0.0 , 0.0 ); data. len ()]; 828 begin_instant =Instant :: now (); 829fft_v4_radix_4 ( & mut result_fft, & mut scratch, & twiddles); 830let duration_fft = begin_instant. elapsed (); 831let result_fft_64: Vec < Complex < f64 >> =complex_to_t ( & result_fft); 832evaluate_results ( 833& result_ref_64, 834& result_fft_64, 835 duration_dft, 836 duration_fft, 837"FFT v4 (radix 4)" , 838); 839} 840 841{ 842let mut result_fft: Vec < Complex < f32 >> =complex_to_t ( & data); 843let twiddles =precompute_twiddles ( data. len ()); 844let mut scratch: Vec < Complex < f32 >> =vec! [ Complex :: new ( 0.0 , 0.0 ); data. len ()]; 845 begin_instant =Instant :: now (); 846fft_v5_s0_opt ( & mut result_fft, & mut scratch, & twiddles); 847let duration_fft = begin_instant. elapsed (); 848let result_fft_64: Vec < Complex < f64 >> =complex_to_t ( & result_fft); 849evaluate_results ( 850& result_ref_64, 851& result_fft_64, 852 duration_dft, 853 duration_fft, 854"FFT v5 (stage 0 opt)" , 855); 856} 857 858{ 859let mut result_fft: Vec < Complex < f32 >> =complex_to_t ( & data); 860let twiddles =precompute_twiddles ( data. len ()); 861let mut scratch: Vec < Complex < f32 >> =vec! [ Complex :: new ( 0.0 , 0.0 ); data. len ()]; 862 begin_instant =Instant :: now (); 863fft_v6_unsafe ( & mut result_fft, & mut scratch, & twiddles); 864let duration_fft = begin_instant. elapsed (); 865let result_fft_64: Vec < Complex < f64 >> =complex_to_t ( & result_fft); 866evaluate_results ( 867& result_ref_64, 868& result_fft_64, 869 duration_dft, 870 duration_fft, 871"FFT v6 (unsafe)" , 872); 873} 874 875{ 876let mut result_fft: Vec < Complex < f32 >> =complex_to_t ( & data); 877let twiddles =precompute_twiddles ( data. len ()); 878let mut scratch: Vec < Complex < f32 >> =vec! [ Complex :: new ( 0.0 , 0.0 ); data. len ()]; 879 begin_instant =Instant :: now (); 880fft_v7_radix_8 ( & mut result_fft, & mut scratch, & twiddles); 881let duration_fft = begin_instant. elapsed (); 882let result_fft_64: Vec < Complex < f64 >> =complex_to_t ( & result_fft); 883evaluate_results ( 884& result_ref_64, 885& result_fft_64, 886 duration_dft, 887 duration_fft, 888"FFT v7 (radix-8)" , 889); 890} 891 892let mut result_rfft: Vec < rustfft:: num_complex:: Complex < f32 >> =complex_to_t ( & data); 893let mut rfft_planner = rustfft:: FftPlanner ::< f32 >:: new (); 894let rfft = rfft_planner. plan_fft_forward ( result_ref. len ()); 895 begin_instant =Instant :: now (); 896 rfft. process ( & mut result_rfft); 897let duration_rfft = begin_instant. elapsed (); 898let result_rfft_64: Vec < Complex < f64 >> =complex_to_t ( & result_rfft); 899evaluate_results ( 900& result_ref_64, 901& result_rfft_64, 902 duration_dft, 903 duration_rfft, 904"RFFT" , 905); 906} 907} 908}