qwerty787788

Fast FFT

Jul 17th, 2015
439
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Java 5.21 KB | None | 0 0
  1. import java.io.*;
  2. import java.util.*;
  3.  
  4. public class CF {
  5.  
  6.     void FFT(double[] re, double[] im, boolean invert) {
  7.         int n = re.length;
  8.         if (im.length != n) {
  9.             throw new AssertionError("Sizes of arrays differ");
  10.         }
  11.         if (Integer.bitCount(n) != 1) {
  12.             throw new AssertionError("N is not power of 2");
  13.         }
  14.         int shift = 32 - Integer.numberOfTrailingZeros(n);
  15.         for (int i = 1; i < n; i++) {
  16.             int j = Integer.reverse(i << shift);
  17.             if (i < j) {
  18.                 double temp = re[i];
  19.                 re[i] = re[j];
  20.                 re[j] = temp;
  21.                 temp = im[i];
  22.                 im[i] = im[j];
  23.                 im[j] = temp;
  24.             }
  25.         }
  26.         for (int len = 2; len <= n; len *= 2) {
  27.             int half = len / 2;
  28.             double alpha = 2 * Math.PI / len;
  29.             double cosAlpha = Math.cos(alpha);
  30.             double sinAlpha = (invert ? -1 : 1) * Math.sin(alpha);
  31.             for (int start = 0; start < n; start += len) {
  32.                 double curRe = 1;
  33.                 double curIm = 0;
  34.                 for (int j = 0; j < half; j++) {
  35.                     double uRe = re[start + j];
  36.                     double uIm = im[start + j];
  37.                     double vRe = re[start + j + half] * curRe
  38.                             - im[start + j + half] * curIm;
  39.                     double vIm = re[start + j + half] * curIm
  40.                             + im[start + j + half] * curRe;
  41.                     re[start + j] = uRe + vRe;
  42.                     im[start + j] = uIm + vIm;
  43.                     re[start + j + half] = uRe - vRe;
  44.                     im[start + j + half] = uIm - vIm;
  45.                     double newRe = curRe * cosAlpha - curIm * sinAlpha;
  46.                     curIm = curRe * sinAlpha + curIm * cosAlpha;
  47.                     curRe = newRe;
  48.                 }
  49.             }
  50.         }
  51.         if (invert) {
  52.             for (int i = 0; i < n; i++) {
  53.                 re[i] /= n;
  54.                 im[i] /= n;
  55.             }
  56.         }
  57.     }
  58.  
  59.     long[] mul(long[] a, long[] b) {
  60.         int len = Math.max(a.length, b.length) * 2;
  61.         int mLen = 1;
  62.         while (mLen < len)
  63.             mLen *= 2;
  64.         double[] r1 = new double[mLen];
  65.         double[] i1 = new double[mLen];
  66.         for (int i = 0; i < a.length; i++)
  67.             r1[i] = a[i];
  68.         double[] r2 = new double[mLen];
  69.         double[] i2 = new double[mLen];
  70.         for (int i = 0; i < b.length; i++)
  71.             r2[i] = b[i];
  72.         FFT(r1, i1, false);
  73.         FFT(r2, i2, false);
  74.         double[] rNew = new double[mLen];
  75.         double[] iNew = new double[mLen];
  76.         for (int i = 0; i < mLen; i++) {
  77.             rNew[i] = r1[i] * r2[i] - i1[i] * i2[i];
  78.             iNew[i] = r1[i] * i2[i] + r2[i] * i1[i];
  79.         }
  80.         FFT(rNew, iNew, true);
  81.         long[] res = new long[mLen];
  82.         for (int i = 0; i < mLen; i++)
  83.             res[i] = (long) Math.round(rNew[i]);
  84.         return res;
  85.     }
  86.  
  87.     long[] slowMul(long[] a, long[] b) {
  88.         long[] res = new long[a.length + b.length + 2];
  89.         for (int i = 0; i < a.length; i++)
  90.             for (int j = 0; j < b.length; j++)
  91.                 res[i + j] += a[i] * b[j];
  92.         return res;
  93.     }
  94.  
  95.     boolean same(long[] a, long[] b) {
  96.         for (int i = 0; i < Math.max(a.length, b.length); i++) {
  97.             if (i < a.length && i < b.length) {
  98.                 if (a[i] != b[i])
  99.                     return false;
  100.             } else {
  101.                 if (i < a.length) {
  102.                     if (a[i] != 0)
  103.                         return false;
  104.                 } else {
  105.                     if (b[i] != 0)
  106.                         return false;
  107.                 }
  108.             }
  109.         }
  110.         return true;
  111.     }
  112.  
  113.     void ACTest() {
  114.         for (int test = 0;; test++) {
  115.             System.err.println(test);
  116.             int n = 1 + rnd.nextInt(100);
  117.             long[] a = new long[n];
  118.             long[] b = new long[n];
  119.             for (int i = 0; i < n; i++) {
  120.                 a[i] = rnd.nextInt(100);
  121.                 b[i] = rnd.nextInt(100);
  122.             }
  123.             long[] r1 = mul(a, b);
  124.             long[] r2 = slowMul(a, b);
  125.             if (!same(r1, r2)) {
  126.                 System.err.println(Arrays.toString(r1));
  127.                 System.err.println(Arrays.toString(r2));
  128.                 throw new AssertionError();
  129.             }
  130.         }
  131.     }
  132.  
  133.     void TLTest() {
  134.         for (int test = 0;; test++) {
  135.             System.err.println(test);
  136.             long st = System.currentTimeMillis();
  137.             int n = 1000000;
  138.             long[] a = new long[n];
  139.             long[] b = new long[n];
  140.             for (int i = 0; i < n; i++) {
  141.                 a[i] = rnd.nextInt(10000);
  142.                 b[i] = rnd.nextInt(10000);
  143.             }
  144.             long[] r1 = mul(a, b);
  145.             System.err.println(System.currentTimeMillis() - st + " ms");
  146.         }
  147.     }
  148.  
  149.     Random rnd = new Random(77);
  150.  
  151.     void realSolve() throws IOException {
  152.         TLTest();
  153.     }
  154.  
  155.     private class InputReader {
  156.         StringTokenizer st;
  157.         BufferedReader br;
  158.  
  159.         public InputReader(File f) {
  160.             try {
  161.                 br = new BufferedReader(new FileReader(f));
  162.             } catch (FileNotFoundException e) {
  163.                 e.printStackTrace();
  164.             }
  165.         }
  166.  
  167.         public InputReader(InputStream f) {
  168.             br = new BufferedReader(new InputStreamReader(f));
  169.         }
  170.  
  171.         String next() {
  172.             while (st == null || !st.hasMoreElements()) {
  173.                 String s;
  174.                 try {
  175.                     s = br.readLine();
  176.                 } catch (IOException e) {
  177.                     return null;
  178.                 }
  179.                 if (s == null)
  180.                     return null;
  181.                 st = new StringTokenizer(s);
  182.             }
  183.             return st.nextToken();
  184.         }
  185.  
  186.         int nextInt() {
  187.             return Integer.parseInt(next());
  188.         }
  189.  
  190.         double nextDouble() {
  191.             return Double.parseDouble(next());
  192.         }
  193.  
  194.         boolean hasMoreElements() {
  195.             while (st == null || !st.hasMoreElements()) {
  196.                 String s;
  197.                 try {
  198.                     s = br.readLine();
  199.                 } catch (IOException e) {
  200.                     return false;
  201.                 }
  202.                 st = new StringTokenizer(s);
  203.             }
  204.             return st.hasMoreElements();
  205.         }
  206.  
  207.         long nextLong() {
  208.             return Long.parseLong(next());
  209.         }
  210.     }
  211.  
  212.     InputReader in;
  213.     PrintWriter out;
  214.  
  215.     void solveIO() throws IOException {
  216.         in = new InputReader(System.in);
  217.         out = new PrintWriter(System.out);
  218.  
  219.         realSolve();
  220.  
  221.         out.close();
  222.  
  223.     }
  224.  
  225.     public static void main(String[] args) throws IOException {
  226.         new CF().solveIO();
  227.     }
  228. }
Advertisement
Add Comment
Please, Sign In to add comment