qwerty787788

1.5x Fast FFT

Jul 10th, 2017
590
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Java 7.74 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[] mulFast(long[] a, long[] b) {
  88.         int len = Math.max(a.length, b.length) * 2;
  89.         int mLen = 1;
  90.         while (mLen < len)
  91.             mLen *= 2;
  92.         double[] r1 = new double[mLen];
  93.         double[] i1 = new double[mLen];
  94.         for (int i = 0; i < a.length; i++) {
  95.             r1[i] = a[i];
  96.             i1[i] = b[i];
  97.         }
  98.         FFT(r1, i1, false);
  99.         double[] rNew = new double[mLen];
  100.         double[] iNew = new double[mLen];
  101.         for (int i = 0; i < mLen; i++) {
  102.             rNew[i] = r1[i] * r1[i] - i1[i] * i1[i];
  103.             iNew[i] = r1[i] * i1[i] + r1[i] * i1[i];
  104.         }
  105.         FFT(rNew, iNew, true);
  106.         long[] res = new long[mLen];
  107.         for (int i = 0; i < mLen; i++)
  108.             res[i] = (long) Math.round(iNew[i] / 2);
  109.         return res;
  110.     }
  111.  
  112.     long[] slowMul(long[] a, long[] b) {
  113.         long[] res = new long[a.length + b.length + 2];
  114.         for (int i = 0; i < a.length; i++)
  115.             for (int j = 0; j < b.length; j++)
  116.                 res[i + j] += a[i] * b[j];
  117.         return res;
  118.     }
  119.  
  120.     boolean same(long[] a, long[] b) {
  121.         for (int i = 0; i < Math.max(a.length, b.length); i++) {
  122.             if (i < a.length && i < b.length) {
  123.                 if (a[i] != b[i])
  124.                     return false;
  125.             } else {
  126.                 if (i < a.length) {
  127.                     if (a[i] != 0)
  128.                         return false;
  129.                 } else {
  130.                     if (b[i] != 0)
  131.                         return false;
  132.                 }
  133.             }
  134.         }
  135.         return true;
  136.     }
  137.  
  138.     void ACTest() {
  139.         for (int test = 0;; test++) {
  140.             System.err.println(test);
  141.             int n = 1 + rnd.nextInt(100);
  142.             long[] a = new long[n];
  143.             long[] b = new long[n];
  144.             for (int i = 0; i < n; i++) {
  145.                 a[i] = rnd.nextInt(100);
  146.                 b[i] = rnd.nextInt(100);
  147.             }
  148.             long[] r1 = mulFast(a, b);
  149.             long[] r2 = slowMul(a, b);
  150.             if (!same(r1, r2)) {
  151.                 System.err.println(Arrays.toString(r1));
  152.                 System.err.println(Arrays.toString(r2));
  153.                 throw new AssertionError();
  154.             }
  155.         }
  156.     }
  157.  
  158.     void TLTest() {
  159.         for (int test = 0;; test++) {
  160.             System.err.println(test);
  161.             long st = System.currentTimeMillis();
  162.             int n = 1000000;
  163.             long[] a = new long[n];
  164.             long[] b = new long[n];
  165.             for (int i = 0; i < n; i++) {
  166.                 a[i] = rnd.nextInt(10000);
  167.                 b[i] = rnd.nextInt(10000);
  168.             }
  169.             long[] r1 = mulFast(a, b);
  170.             System.err.println(System.currentTimeMillis() - st + " ms");
  171.         }
  172.     }
  173.  
  174.     Random rnd = new Random(77);
  175.  
  176.     void realSolve() throws IOException {
  177.         TLTest();
  178.     }
  179.  
  180.     private class InputReader {
  181.         StringTokenizer st;
  182.         BufferedReader br;
  183.  
  184.         public InputReader(File f) {
  185.             try {
  186.                 br = new BufferedReader(new FileReader(f));
  187.             } catch (FileNotFoundException e) {
  188.                 e.printStackTrace();
  189.             }
  190.         }
  191.  
  192.         public InputReader(InputStream f) {
  193.             br = new BufferedReader(new InputStreamReader(f));
  194.         }
  195.  
  196.         String next() {
  197.             while (st == null || !st.hasMoreElements()) {
  198.                 String s;
  199.                 try {
  200.                     s = br.readLine();
  201.                 } catch (IOException e) {
  202.                     return null;
  203.                 }
  204.                 if (s == null)
  205.                     return null;
  206.                 st = new StringTokenizer(s);
  207.             }
  208.             return st.nextToken();
  209.         }
  210.  
  211.         int nextInt() {
  212.             return Integer.parseInt(next());
  213.         }
  214.  
  215.         double nextDouble() {
  216.             return Double.parseDouble(next());
  217.         }
  218.  
  219.         boolean hasMoreElements() {
  220.             while (st == null || !st.hasMoreElements()) {
  221.                 String s;
  222.                 try {
  223.                     s = br.readLine();
  224.                 } catch (IOException e) {
  225.                     return false;
  226.                 }
  227.                 st = new StringTokenizer(s);
  228.             }
  229.             return st.hasMoreElements();
  230.         }
  231.  
  232.         long nextLong() {
  233.             return Long.parseLong(next());
  234.         }
  235.     }
  236.  
  237.     InputReader in;
  238.     PrintWriter out;
  239.  
  240.     void solveIO() throws IOException {
  241.         in = new InputReader(System.in);
  242.         out = new PrintWriter(System.out);
  243.  
  244.         realSolve();
  245.  
  246.         out.close();
  247.  
  248.     }
  249.  
  250.     public static void main(String[] args) throws IOException {
  251.         new CF().solveIO();
  252.     }
  253. }
Advertisement
Add Comment
Please, Sign In to add comment