Skip to main content

intt_batch

Function intt_batch 

Source
fn intt_batch<M>(a: &mut [MInt<M>], width: usize)
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform.rs (line 1388)
1335    fn power_projection_step(
1336        p_flat: Vec<MInt<M>>,
1337        q_flat: Vec<MInt<M>>,
1338        n: usize,
1339        py: usize,
1340        qy: usize,
1341    ) -> (Vec<MInt<M>>, Vec<MInt<M>>) {
1342        let high_degree = (qy - 1) * 2;
1343        let rows = (py + qy - 1).max(high_degree).next_power_of_two();
1344        let cols = n * 2;
1345        let size = rows * cols;
1346        let mut p = p_flat;
1347        p.resize_with(size, MInt::<M>::zero);
1348        ntt_rows(&mut p, cols);
1349        ntt_batch(&mut p, cols);
1350
1351        let mut q = q_flat;
1352        q.resize_with(size, MInt::<M>::zero);
1353        ntt_rows(&mut q, cols);
1354        let q_high = (rows == high_degree).then(|| q[(qy - 1) * cols..qy * cols].to_vec());
1355        ntt_batch(&mut q, cols);
1356
1357        let half = cols / 2;
1358        let mut odd_factor = vec![MInt::<M>::zero(); half];
1359        let mut factor = MInt::<M>::from(2).inv();
1360        let k = cols.trailing_zeros() as usize;
1361        let w = MInt::<M>::new_unchecked(M::INFO.inv_root[k]);
1362        BIT_REVERSE.with(|br| {
1363            let br = unsafe { &mut *br.get() };
1364            if br.len() < k {
1365                br.resize_with(k, Default::default);
1366            }
1367            let k = k - 1;
1368            if br[k].is_empty() {
1369                let mut v = vec![0; 1 << k];
1370                for i in 0..1 << k {
1371                    v[i] = (v[i >> 1] >> 1) | ((i & 1) << k.saturating_sub(1));
1372                }
1373                br[k] = v;
1374            }
1375            for &i in &br[k] {
1376                odd_factor[i] = factor;
1377                factor *= w;
1378            }
1379        });
1380
1381        let mut pr = vec![MInt::<M>::zero(); rows * half];
1382        let mut qr = vec![MInt::<M>::zero(); rows * half];
1383        for i in 0..pr.len() {
1384            pr[i] = (p[i << 1] * q[i << 1 | 1] - p[i << 1 | 1] * q[i << 1])
1385                * odd_factor[i & (half - 1)];
1386            qr[i] = q[i << 1] * q[i << 1 | 1];
1387        }
1388        intt_batch(&mut pr, half);
1389        intt_rows(&mut pr, half);
1390        intt_batch(&mut qr, half);
1391        intt_rows(&mut qr, half);
1392
1393        if let Some(q_high) = q_high {
1394            let mut q_high_even = vec![MInt::<M>::zero(); half];
1395            for i in 0..half {
1396                q_high_even[i] = q_high[i << 1] * q_high[i << 1 | 1];
1397            }
1398            intt(&mut q_high_even);
1399            for (value, high) in qr.iter_mut().zip(&q_high_even) {
1400                *value -= *high;
1401            }
1402            qr.extend_from_slice(&q_high_even);
1403        }
1404        (pr, qr)
1405    }