ecnerwala's competitive programming library
#include "fft/engines/split.hpp"
| Coverage | Exec / Excl / Total | |
|---|---|---|
| Lines | 93.6% | 102 / 0 / 109 |
| Functions | 100.0% | 25 / 0 / 25 |
| Branches | 73.8% | 163 / 0 / 221 |
| Full report |
#pragma once
#include <algorithm>
#include <cassert>
#include <cmath>
#include <cstdint>
#include <span>
#include <utility>
#include <vector>
#include "fft/core.hpp"
#include "fft/engine.hpp"
namespace ecnerwala::fft::engines {
// Multiplies mod `mnum` by splitting values into balanced 15-bit halves (each limb in
// [-2^14, 2^14], from the balanced representative |v| <= MOD/2) packed into one complex
// transform per operand.
template <typename mnum> struct split {
static_assert(sizeof(decltype(mnum::MOD)) <= 4, "limbs must fit 15 bits");
using value_type = mnum;
static constexpr bool commutative = true;
static constexpr int unit_scale = 1;
using cnum = cplx<double>;
using core = fft_core<cnum>;
template <int A = 1> struct transformed_t {
vector<cnum> v;
int size() const { return sz(v); }
transformed_t() = default;
explicit transformed_t(vector<cnum>&& v_) : v(std::move(v_)) {}
template <int A2> requires (A2 != A) explicit(A2 > A) transformed_t(transformed_t<A2>&& o)
: v(std::move(o.v)) {}
};
using transformed = transformed_t<1>;
template <int K> struct product_t {
// After finish's inverse transforms: lo = (lo*lo, hi*lo), hi = (lo*hi, hi*hi).
vector<cnum> lo, hi;
int size() const { return sz(lo); }
product_t() = default;
product_t(vector<cnum>&& lo_, vector<cnum>&& hi_) : lo(std::move(lo_)), hi(std::move(hi_)) {}
template <int K2> requires (K2 != K) explicit(K2 > K) product_t(product_t<K2>&& o)
: lo(std::move(o.lo)), hi(std::move(o.hi)) {}
};
using product = product_t<1>;
static cnum pack(mnum x) {
int64_t v = x.balanced();
int64_t hi = (v + (1 << 14)) >> 15;
return cnum(double(v - (hi << 15)), double(hi));
}
static transformed transform(std::span<const mnum> a, int n) {
assert(sz(a) <= 2 * n);
transformed r;
r.v.assign(n, cnum(0));
for (int i = 0; i < sz(a); i++) {
int j = i < n ? i : i - n;
r.v[j] = r.v[j] + pack(a[i]);
}
core::forward(std::span<cnum>(r.v));
return r;
}
static void extend_to(transformed& t, int m, std::span<const mnum> coeffs) {
assert(!(m & (m-1)) && sz(coeffs) <= 2 * m);
if (t.size() >= m) return;
if (t.size() == 0) { t = transform(coeffs, m); return; }
auto buf = buffer_pool<cnum>::get(sz(coeffs));
for (int i = 0; i < sz(coeffs); i++) buf[i] = pack(coeffs[i]);
while (t.size() < m) {
int s = t.size();
t.v.resize(2 * s);
// coeffs past 2s are zero: they didn't fit in the transform we're a prefix of
core::extend(
std::span<cnum>(t.v),
std::span<const cnum>(buf.span()).first(size_t(min(sz(coeffs), 2 * s)))
);
}
}
static void downsample_core(std::span<const cnum> in, std::span<cnum> out, bool odd) {
if (odd) core::odd_half(in, out);
else core::even_half(in, out);
}
template <int A> static transformed_t<A> downsample(const transformed_t<A>& t, int n, bool odd) {
transformed_t<A> r; r.v.resize(n);
downsample_core(std::span<const cnum>(t.v), std::span<cnum>(r.v), odd);
return r;
}
template <int K> static product_t<K> downsample(const product_t<K>& p, int n, bool odd) {
product_t<K> r; r.lo.resize(n); r.hi.resize(n);
downsample_core(std::span<const cnum>(p.lo), std::span<cnum>(r.lo), odd);
downsample_core(std::span<const cnum>(p.hi), std::span<cnum>(r.hi), odd);
return r;
}
template <int A> static transformed_t<A> negate_arg(const transformed_t<A>& t, int n) {
assert(n >= 2 && t.size() >= n);
transformed_t<A> r; r.v.resize(n);
for (int j = 0; j < n; j++) r.v[j] = t.v[j ^ 1];
return r;
}
template <int A, int B> static transformed_t<A + B> add(transformed_t<A>&& a, const transformed_t<B>& b) {
transformed_t<A + B> r{std::move(a.v)};
add_into(r.v, b.v);
return r;
}
// Unpacks b's transform into transforms of its low/high halves via conjugate
// symmetry, then multiplies both against a's (still packed) transform. The scale
// parameter only affects the bookkeeping, so the body is a shared untyped impl.
static void mul_impl(const vector<cnum>& a, const vector<cnum>& b, vector<cnum>& lo, vector<cnum>& hi, int n, bool acc = false) {
core::init(n);
lo.resize(n); hi.resize(n);
for (int i = 0; i < n; i++) {
int ci = core::conj_index(i);
cnum g0 = (b[i] + conj(b[ci])) * cnum(0.5);
cnum t = (b[i] - conj(b[ci])) * cnum(0.5);
cnum g1 = cnum(t.y, -t.x);
if (acc) {
lo[i] = lo[i] + a[i] * g0;
hi[i] = hi[i] + a[i] * g1;
} else {
lo[i] = a[i] * g0;
hi[i] = a[i] * g1;
}
}
}
template <int A, int B> static product_t<A * B> mul(const transformed_t<A>& a, const transformed_t<B>& b, int n) {
assert(a.size() >= n && b.size() >= n);
product_t<A * B> p;
mul_impl(a.v, b.v, p.lo, p.hi, n);
return p;
}
template <int A> static product_t<A * A> sq(const transformed_t<A>& a, int n) { return mul(a, a, n); }
template <int A1, int B1, int A2, int B2>
static product_t<A1 * B1 + A2 * B2> mul2(
const transformed_t<A1>& a1, const transformed_t<B1>& b1,
const transformed_t<A2>& a2, const transformed_t<B2>& b2,
int n
) {
assert(a1.size() >= n && b1.size() >= n && a2.size() >= n && b2.size() >= n);
product_t<A1 * B1 + A2 * B2> p;
mul_impl(a1.v, b1.v, p.lo, p.hi, n);
mul_impl(a2.v, b2.v, p.lo, p.hi, n, true);
return p;
}
static void add_into(vector<cnum>& a, const vector<cnum>& b) {
assert(sz(a) == sz(b));
for (int i = 0; i < sz(a); i++) a[i] = a[i] + b[i];
}
template <int K1, int K2> static product_t<K1 + K2> add(product_t<K1>&& a, product_t<K2>&& b) {
product_t<K1 + K2> r{std::move(a.lo), std::move(a.hi)};
add_into(r.lo, b.lo);
add_into(r.hi, b.hi);
return r;
}
template <int K = 1, typename Op = assign_op> static void finish(product_t<K>&& p, std::span<mnum> out, Op op = {}) {
// The fp error budget is divided by the accumulated scale; K <= 2 is very
// conservative (balanced limbs already left ~2x headroom at max lengths).
static_assert(K <= 2, "split: accumulated scale too large");
int n = p.size();
assert(sz(out) <= n);
core::inverse(std::span<cnum>(p.lo));
core::inverse(std::span<cnum>(p.hi));
const int64_t m = mnum::MOD;
double d = 1.0 / double(n);
// llround + a final wrap so negative half-products (e.g. from negate_arg'd
// transforms) reconstruct correctly.
for (int i = 0; i < sz(out); i++) {
int64_t v = (llround(p.lo[i].x * d)
+ (llround(p.lo[i].y * d) % m << 15)
+ (llround(p.hi[i].x * d) % m << 15)
+ (llround(p.hi[i].y * d) % m << 30)) % m;
if (v < 0) v += m;
op(out[i], mnum(v));
}
}
};
/* namespace ecnerwala::fft::engines */ }
#include <algorithm>
#include <cassert>
#include <cmath>
#include <cstdint>
#include <span>
#include <utility>
#include <vector>
#include <iterator>
#include <iostream>
#include <concepts>
#include <limits>
#include <type_traits>
#line 2 "src/fft/engines/split.hpp"
#line 10 "src/fft/engines/split.hpp"
#line 2 "src/fft/core.hpp"
#line 8 "src/fft/core.hpp"
#line 2 "src/fft/common.hpp"
#line 8 "src/fft/common.hpp"
/**
* Author: Andrew He
* Source: http://neerc.ifmo.ru/trains/toulouse/2017/fft2.pdf
* Papers about accuracy: http://www.daemonology.net/papers/fft.pdf, http://www.cs.berkeley.edu/~fateman/papers/fftvsothers.pdf
* For integers rounding works if $(|a| + |b|)\max(a, b) < \mathtt{\sim} 10^9$, or in theory maybe $10^6$.
*
* Abstraction layers:
* fft_core<num> FFT itself and other ops on rings with 2^k-th roots of unity. We use bit-reversed indexing in the frequency domain.
*
* engines Engines for packing/unpacking arbitrary rings for convolution (the `engine` concept).
* Still expose (opaque) transform-domain objects for caching/fusion.
*
* multiply layer Wrappers for convolving bounded sequences: track length/truncation.
*
* value types series::vec<E, exact> - R[[x]]
* series::exact<E> - exact (finite-support) power series
* series::trunc<E> - truncated prefix of an (infinite) power series
*
* polynomials - R[x]. Under x -> 1/x a polynomial becomes a Laurent polynomial in 1/x;
* shifting by x^{deg P} (reversal) lands it in R[[x]], and we store that exact series.
* poly::vec<E> - polynomial type, supporting natural indexing
* poly::form<E> - finite-support linear forms, via the pairing <P, S> = [x^0] P(1/x) S(x)
* a linear form is one side of this pairing, applied to the other
*
* online_multiplier<E> - online (relaxed) multiplication of 2 sequences in n log^2 n time
* ap_sampled_poly<E> - a polynomial stored as its evaluations on an arithmetic progression
*/
namespace ecnerwala {
template<class T> int sz(T&& arg) { using std::size; return int(size(std::forward<T>(arg))); }
inline int nextPow2(int s) { return 1 << (s > 1 ? 32 - __builtin_clz(s-1) : 0); }
namespace fft {
using std::swap;
using std::vector;
using std::min;
using std::max;
// Reusable scratch buffers. Not thread-safe by default: this is deliberately plain
// static storage so single-threaded programs pay no TLS indirection; define
// ECNERWALA_FFT_POOL_STORAGE to `thread_local` for multithreaded use.
#ifndef ECNERWALA_FFT_POOL_STORAGE
#define ECNERWALA_FFT_POOL_STORAGE
#endif
template <typename T> struct buffer_pool {
static inline ECNERWALA_FFT_POOL_STORAGE std::vector<std::vector<T>> free_list;
struct handle {
std::vector<T> v;
explicit handle(int n) {
if (!free_list.empty()) {
v = std::move(free_list.back());
free_list.pop_back();
}
v.assign(n, T());
}
handle(const handle&) = delete;
handle& operator=(const handle&) = delete;
handle(handle&& o) noexcept : v(std::move(o.v)) {}
~handle() { if (v.capacity()) free_list.push_back(std::move(v)); }
T& operator[](int i) { return v[i]; }
operator std::span<T>() { return std::span<T>(v); }
std::span<T> span() { return std::span<T>(v); }
};
static handle get(int n) { return handle(n); }
};
/* namespace fft */ }
/* namespace ecnerwala */ }
#line 2 "src/modnum.hpp"
#line 9 "src/modnum.hpp"
template <typename T> T mod_inv_in_range(T a, T m) {
// assert(0 <= a && a < m);
T x = a, y = m;
// abs coeff of a in x and y (they're always opposite sign)
T vx = 1, vy = 0;
bool swap = false;
while (x) {
T k = y / x;
y %= x;
vy += k * vx;
std::swap(x, y);
std::swap(vx, vy);
swap ^= 1;
}
assert(y == 1);
return swap ? vy : m - vy;
}
template <typename T> struct extended_gcd_result {
T gcd;
T coeff_a, coeff_b;
};
template <typename T> extended_gcd_result<T> extended_gcd(T a, T b) {
T x = a, y = b;
// coeff of a and b in x and y
T ax = 1, ay = 0;
T bx = 0, by = 1;
while (x) {
T k = y / x;
y %= x;
ay -= k * ax;
by -= k * bx;
std::swap(x, y);
std::swap(ax, ay);
std::swap(bx, by);
}
return {y, ay, by};
}
template <typename T> T mod_inv(T a, T m) {
a %= m;
a = a < 0 ? a + m : a;
return mod_inv_in_range(a, m);
}
// Derives the boilerplate operator surface of a number type from its compound
// ops, ==, neg(), and inv().
// Bodies are only instantiated on use, so a type may omit some of the
// underlying pieces if the corresponding derived ops are never called.
template <typename Self>
struct num_ops {
Self operator+ () const { return static_cast<const Self&>(*this); }
Self operator- () const { return static_cast<const Self&>(*this).neg(); }
friend Self operator ++ (Self& a, int) { Self r = a; ++a; return r; }
friend Self operator -- (Self& a, int) { Self r = a; --a; return r; }
friend Self operator + (const Self& a, const Self& b) { return Self(a) += b; }
friend Self operator - (const Self& a, const Self& b) { return Self(a) -= b; }
friend Self operator * (const Self& a, const Self& b) { return Self(a) *= b; }
friend Self operator / (const Self& a, const Self& b) { return Self(a) /= b; }
friend bool operator != (const Self& a, const Self& b) { return !(a == b); }
friend Self neg(const Self& a) { return a.neg(); }
friend Self inv(const Self& a) { return a.inv(); }
};
// Storage and arithmetic for numbers mod Self::MOD, as a reduced
// representative v in [0, MOD) of unsigned type V.
// The type provides static MOD (of type V), reduce (value -> representative),
// and *=;
// everything else is derived here, valid for any MOD up to V's full range
// (sums and differences are tracked mod 2^bits, so no headroom is needed).
// Hooks may be overridden in the type's own body (e.g. a faster += / -=).
template <typename Self, typename V>
struct mod_ops : num_ops<Self> {
static_assert(std::unsigned_integral<V>);
V v;
struct is_reduced_tag {};
mod_ops() : v(0) {}
mod_ops(V v_, is_reduced_tag) : v(v_) { assert(v < Self::MOD); }
template <std::integral I> mod_ops(I x) : v(Self::reduce(x)) {}
static Self from_reduced(V v) { return Self(v, is_reduced_tag{}); }
// A negative value reduces via its nonnegative complement: x = -1 - ~x.
static V reduce(std::signed_integral auto x) {
using U = std::make_unsigned_t<decltype(x)>;
return x < 0 ? V(Self::MOD - 1 - Self::reduce(U(~x))) : Self::reduce(U(x));
}
explicit operator V() const { return v; }
std::make_signed_t<V> balanced() const {
return std::make_signed_t<V>(Self::MOD-v > v ? v : v - Self::MOD);
}
friend bool operator == (const Self& a, const Self& b) { return a.v == b.v; }
friend std::ostream& operator << (std::ostream& out, const Self& n) { return out << n.v; }
friend std::istream& operator >> (std::istream& in, Self& n) { int64_t v_; in >> v_; n = Self(v_); return in; }
Self& operator ++ () {
++v;
if (v == Self::MOD) v = 0;
return self();
}
Self& operator -- () {
if (v == 0) v = Self::MOD;
--v;
return self();
}
Self& operator += (const Self& o) { v = Self::sub_mod_raw(v, Self::MOD - o.v); return self(); }
Self& operator -= (const Self& o) { v = Self::sub_mod_raw(v, o.v); return self(); }
Self& operator /= (const Self& o) { return self() *= o.inv(); }
// Returns a - b mod MOD, for b in [0, MOD]; wraparound detects the underflow.
static V sub_mod_raw(V a, V b) { return a < b ? a - b + Self::MOD : a - b; }
Self neg() const { return from_reduced(v ? Self::MOD - v : 0); }
Self inv() const { return from_reduced(mod_inv_in_range(v, Self::MOD)); }
private:
Self& self() { return static_cast<Self&>(*this); }
};
template <auto MOD_> struct modnum : mod_ops<modnum<MOD_>, std::make_unsigned_t<decltype(MOD_)>> {
using Self = modnum;
static_assert(MOD_ > 0, "MOD must be positive");
using V = std::make_unsigned_t<decltype(MOD_)>;
static constexpr V MOD = V(MOD_);
using base = mod_ops<modnum, V>;
using base::base;
using base::v;
using base::reduce;
static V reduce(std::unsigned_integral auto x) { return V(x % MOD); }
explicit operator std::make_signed_t<V>() const
requires (MOD <= V(std::numeric_limits<std::make_signed_t<V>>::max()))
{
return std::make_signed_t<V>(v);
}
Self& operator *= (const Self& o) {
if constexpr (sizeof(V) <= 4) v = V(uint64_t(v) * o.v % MOD);
else v = V(__uint128_t(v) * o.v % MOD);
return *this;
}
};
struct mod_goldilocks : mod_ops<mod_goldilocks, uint64_t> {
using Self = mod_goldilocks;
static constexpr uint64_t MOD = 0xffffffff00000001ull;
static constexpr uint64_t EPS = -MOD;
// We have 2^32 is a primitive 6th root of unity.
// Note that omega_8 + omega_8^7 == 2^24 - 2^72 == sqrt(2)
// We'll pick the root so that 2^24 - 2^72 is our primitive 384th root of unity.
static constexpr uint64_t PRIMITIVE_ROOT = 2717;
using base = mod_ops<mod_goldilocks, uint64_t>;
using base::base;
using base::reduce;
mod_goldilocks() = default;
mod_goldilocks(__int128_t a) : base(a < 0 ? uint64_t(MOD - 1 - __uint128_t(~a) % MOD) : uint64_t(__uint128_t(a) % MOD), is_reduced_tag{}) {}
mod_goldilocks(__uint128_t a) : base(uint64_t(a % MOD), is_reduced_tag{}) {}
// Avoids the division: any uint64_t is within MOD of reduced.
static uint64_t reduce(std::unsigned_integral auto x) {
static_assert(sizeof(x) <= 8);
uint64_t a = x;
return a >= MOD ? a - MOD : a;
}
// returns a-b, assuming -MOD <= a-b, e.g. b <= MOD
static uint64_t sub_mod_raw(uint64_t a, uint64_t b) {
#if defined(__x86_64__)
// TODO: We could try to write this using intrinsics, but GCC sometimes produces the wrong code.
uint64_t res_wrapped = a;
uint64_t adjustment = b;
asm (
// AT&T syntax: SRC DST
"sub %[y], %[x]\n\t"
// Trick from plonky2 implementation:
// After sub, flag CF is set iff we underflowed. We want to correct by EPS == 2^32 - 1 iff C is set.
// sbb (subtract with borrow) computes DST <- DST - SRC - CF
// Thus, we can use the 32-bit form of sbb on a dummy register to load CF ? EPS : 0.
// Here, we'll just reuse the original register holding b.
"sbb %k[y], %k[y]\n\t"
: [x] "+r"(res_wrapped),
[y] "+r"(adjustment)
:
: "cc"
);
#else
uint64_t res_wrapped = a - b;
uint64_t adjustment = (res_wrapped > a) ? EPS : 0;
#endif
return res_wrapped - adjustment;
}
// Reduce lo + 2^64 * mi + 2^96 * hi, where hi <= MOD
static uint64_t reduce_u160_raw(uint64_t lo, uint32_t mi, uint64_t hi) {
// result = lo - hi + EPS * mi
// 0 <= lo <= 2^64 - 1 = MOD + EPS - 1
// 0 <= EPS * mi <= (2^32 - 1) * EPS = MOD - 1 - EPS
// 0 <= hi <= MOD
// -MOD <= lo - hi + EPS * mi <= 2*MOD-2
// so we do have some leeway
return sub_mod_raw(sub_mod_raw(lo, hi), MOD-(uint64_t(mi)<<32)+mi);
}
static uint64_t reduce_u128_raw(__uint128_t v) {
uint64_t hi = uint64_t(v >> 64);
uint64_t lo = uint64_t(v);
uint32_t hi_hi = uint32_t(hi >> 32);
uint32_t hi_lo = uint32_t(hi);
return reduce_u160_raw(lo, hi_lo, hi_hi);
}
Self& operator *= (Self o) {
v = reduce_u128_raw(__uint128_t(v) * __uint128_t(o.v));
return *this;
}
};
template <typename T> T power(T a, long long b) {
assert(b >= 0);
T r = 1; while (b) { if (b & 1) r *= a; b >>= 1; a *= a; } return r;
}
template <typename U, typename V> struct pairnum : num_ops<pairnum<U, V>> {
using Self = pairnum;
U u;
V v;
pairnum() : u(0), v(0) {}
pairnum(long long val) : u(val), v(val) {}
pairnum(const U& u_, const V& v_) : u(u_), v(v_) {}
friend std::ostream& operator << (std::ostream& out, const Self& n) { return out << '(' << n.u << ',' << ' ' << n.v << ')'; }
friend std::istream& operator >> (std::istream& in, Self& n) { long long val; in >> val; n = Self(val); return in; }
friend bool operator == (const Self& a, const Self& b) { return a.u == b.u && a.v == b.v; }
Self inv() const {
return Self(u.inv(), v.inv());
}
Self neg() const {
return Self(u.neg(), v.neg());
}
Self& operator ++ () {
++u, ++v;
return *this;
}
Self& operator -- () {
--u, --v;
return *this;
}
Self& operator += (const Self& o) {
u += o.u;
v += o.v;
return *this;
}
Self& operator -= (const Self& o) {
u -= o.u;
v -= o.v;
return *this;
}
Self& operator *= (const Self& o) {
u *= o.u;
v *= o.v;
return *this;
}
Self& operator /= (const Self& o) {
u /= o.u;
v /= o.v;
return *this;
}
};
template <typename tag> struct dynamic_modnum : mod_ops<dynamic_modnum<tag>, uint32_t> {
using Self = dynamic_modnum;
private:
inline static uint32_t MOD_ = 0;
inline static uint64_t BARRETT_M = 0;
public:
// Make only the const-reference public, to force the use of set_mod
static constexpr uint32_t const& MOD = MOD_;
using base = mod_ops<dynamic_modnum, uint32_t>;
using base::base;
using base::v;
using base::reduce;
// Barret reduction taken from KACTL:
/**
* Author: Simon Lindholm
* Date: 2020-05-30
* License: CC0
* Source: https://en.wikipedia.org/wiki/Barrett_reduction
* Description: Compute $a \% b$ about 5 times faster than usual, where $b$ is constant but not known at compile time.
* Returns a value congruent to $a \pmod b$ in the range $[0, 2b)$.
* Status: proven correct, stress-tested
* Measured as having 4 times lower latency, and 8 times higher throughput, see stress-test.
* Details:
* More precisely, it can be proven that the result equals 0 only if $a = 0$,
* and otherwise lies in $[1, (1 + a/2^64) * b)$.
*/
static void set_mod(int mod) {
assert(mod > 0);
MOD_ = uint32_t(mod);
BARRETT_M = (uint64_t(-1) / MOD);
}
static uint32_t barrett_reduce_partial(uint64_t a) {
return uint32_t(a - uint64_t((__uint128_t(BARRETT_M) * a) >> 64) * MOD);
}
static uint32_t barrett_reduce(uint64_t a) {
int32_t res = int32_t(barrett_reduce_partial(a) - MOD);
return uint32_t((res < 0) ? res + int32_t(MOD) : res);
}
struct mod_reader {
friend std::istream& operator >> (std::istream& i, mod_reader) {
int mod; i >> mod;
Self::set_mod(mod);
return i;
}
};
static mod_reader MOD_READER() {
return mod_reader();
}
static uint32_t reduce(std::unsigned_integral auto x) {
static_assert(sizeof(x) <= 8);
return barrett_reduce(x);
}
explicit operator int() const { return int(v); }
Self& operator *= (const Self& o) {
v = barrett_reduce(uint64_t(v) * o.v);
return *this;
}
};
template <typename T> struct mod_constraint {
T v, mod;
friend mod_constraint operator & (mod_constraint a, mod_constraint b) {
if (a.mod < b.mod) std::swap(a, b);
if (b.mod == 1) return a;
extended_gcd_result<T> egcd = extended_gcd<T>(a.mod, b.mod);
assert(a.v % egcd.gcd == b.v % egcd.gcd);
T extra = b.v - a.v % b.mod;
extra /= egcd.gcd;
extra *= egcd.coeff_a;
extra %= b.mod / egcd.gcd;
extra += (extra < 0) ? b.mod / egcd.gcd : 0;
return mod_constraint{
a.v + extra * a.mod,
a.mod * (b.mod / egcd.gcd)
};
}
};
#line 11 "src/fft/core.hpp"
namespace ecnerwala::fft {
// ==== core: roots, buffers, raw transforms ====
// Complex
template <typename dbl> struct cplx { /// start-hash
dbl x, y;
cplx(dbl x_ = 0, dbl y_ = 0) : x(x_), y(y_) { }
friend cplx operator+(cplx a, cplx b) { return cplx(a.x + b.x, a.y + b.y); }
friend cplx operator-(cplx a, cplx b) { return cplx(a.x - b.x, a.y - b.y); }
friend cplx operator*(cplx a, cplx b) { return cplx(a.x * b.x - a.y * b.y, a.x * b.y + a.y * b.x); }
friend cplx conj(cplx a) { return cplx(a.x, -a.y); }
friend cplx inv(cplx a) { dbl n = (a.x*a.x+a.y*a.y); return cplx(a.x/n,-a.y/n); }
};
// getRoot implementations
template <typename num> struct getRoot {
static num f(int k) = delete;
};
template <typename dbl> struct getRoot<cplx<dbl>> {
static cplx<dbl> f(int k) {
#ifndef M_PI
#define M_PI 3.14159265358979323846
#endif
dbl a=2*M_PI/k;
return cplx<dbl>(cos(a),sin(a));
}
};
template <int MOD> struct primitive_root {
static const int value;
};
// 998244353 = (119 << 23) + 1 = 2^30 - 2^26 - 2^23 + 1
template <> struct primitive_root<998244353> {
static const int value = 3;
};
// babybear prime
template <> struct primitive_root<(15 << 27) + 1> {
static const int value = 31;
};
// koalabear prime
template <> struct primitive_root<(127 << 24) + 1> {
static const int value = 3;
};
template <> struct primitive_root<(7 << 26) + 1> {
static const int value = 3;
};
template <> struct primitive_root<(5 << 25) + 1> {
static const int value = 3;
};
template <int MOD> struct getRoot<modnum<MOD>> {
static modnum<MOD> f(int k) {
assert((MOD-1)%k == 0);
return power(modnum<MOD>(primitive_root<MOD>::value), (MOD-1)/k);
}
};
template <> struct getRoot<mod_goldilocks> {
static mod_goldilocks f(int k) {
assert((mod_goldilocks::MOD-1)%k == 0);
return power(mod_goldilocks(mod_goldilocks::PRIMITIVE_ROOT), (mod_goldilocks::MOD-1)/k);
}
};
// We take the bit-reverse convention: the coefficient of a[i] -> b[j] is omega^{i * bit_reverse(j)}.
// This means that the size 2^{k-1} transform is the prefix of the size 2^k transform (wrapping the input).
//
// We mostly work with spans here:
// spans of transforms are expected to have length exactly 2^k
// spans of inputs/outputs are expected to have length [0, 2^{k+1})
//
// Inputs/outputs are treated mod x^{2^k} - 1.
// Their length is allowed to be bigger than 2^k mostly to perform ops on sequences of size n+1 with only transforms of size n.
// The upper bound of 2^{k+1} is arbitrary: we could tighten it to 2^k + 1 or loosen it to infinity, this is just a "defensive" choice.
template <typename num> struct fft_core {
static inline vector<int> rev;
// rt[2^k + i] = 1^{i / 2^(k+1)}
// TODO: can we get rid of inv_rt; alternatively, should we store inv_rt in bit-reverse order?
static inline vector<num> rt, inv_rt;
static void init(int n) {
if (n <= sz(rt)) return;
rev.resize(n);
for (int i = 0; i < n; i++) {
rev[i] = (rev[i>>1] | ((i&1)*n)) >> 1;
}
rt.reserve(n); inv_rt.reserve(n);
while (sz(rt) < 2 && sz(rt) < n) { rt.push_back(num(1)); inv_rt.push_back(num(1)); }
for (int k = sz(rt); k < n; k *= 2) {
rt.resize(2*k); inv_rt.resize(2*k);
num z = getRoot<num>::f(2*k);
num iz = inv(z);
for (int i = k/2; i < k; i++) {
rt[2*i] = rt[i], rt[2*i+1] = rt[i]*z;
inv_rt[2*i] = inv_rt[i], inv_rt[2*i+1] = inv_rt[i]*iz;
}
}
}
// bit-reversal of i as a log2(n)-bit number
static int brev(int i, int n) {
int s = __builtin_ctz(unsigned(sz(rev)/n));
return rev[i] >> s;
}
// index of the conjugate evaluation point, in the returned bit-reversed order
static int conj_index(int j) {
return j == 0 ? 0 : j ^ ((1 << (31 - __builtin_clz(unsigned(j)))) - 1);
}
static void forward(std::span<num> a) {
int n = sz(a);
if (n <= 1) return;
init(n);
for (int k = n/2; k >= 1; k /= 2) {
for (int i = 0; i < n; i += 2*k) {
for (int j = 0; j < k; j++) {
num u = a[i+j], v = a[i+j+k];
a[i+j] = u + v;
a[i+j+k] = (u - v) * rt[j+k];
}
}
}
}
static void inverse(std::span<num> a) {
int n = sz(a);
if (n <= 1) return;
init(n);
for (int k = 1; k < n; k *= 2) {
for (int i = 0; i < n; i += 2*k) {
for (int j = 0; j < k; j++) {
num t = inv_rt[j+k] * a[i+j+k];
a[i+j+k] = a[i+j] - t;
a[i+j] = a[i+j] + t;
}
}
}
}
// Extend a size 2^{k-1} transform to size 2^k; we need the coefficients.
// t must have size 2^k, and coeffs must have size at most 2^{k+1}.
static void extend(std::span<num> t, std::span<const num> coeffs) {
int n = sz(t) / 2;
assert(sz(coeffs) <= 2 * n);
init(sz(t));
auto b = t.subspan(n, n);
int lo = min(sz(coeffs), n);
for (int i = 0; i < lo; i++) {
// rt[n + i] = w_{2n}^i
b[i] = coeffs[i] * rt[n + i];
}
std::fill(b.begin() + lo, b.end(), num(0));
for (int i = n; i < sz(coeffs); i++) {
b[i - n] = b[i - n] - coeffs[i] * rt[i];
}
forward(b);
}
// Consider t = transform(P) and P(x) = E(x^2) + O(x^2) * x
// `even_half` and `odd_half` extract a size 2^{k-1} transform of E/O, respectively.
static void even_half(std::span<const num> t, std::span<num> out) {
int n = sz(out);
assert(sz(t) >= 2*n);
num half = inv(num(2));
for (int j = 0; j < n; j++) out[j] = (t[2*j] + t[2*j+1]) * half;
}
static void odd_half(std::span<const num> t, std::span<num> out) {
int n = sz(out);
assert(sz(t) >= 2*n);
init(2*n);
num half = inv(num(2));
for (int j = 0; j < n; j++) {
// entry j of the size-2n transform pairs (w, -w) with w = w_{2n}^{brev(j, n)}
out[j] = (t[2*j] - t[2*j+1]) * half * inv_rt[n + brev(j, n)];
}
}
};
/* namespace ecnerwala::fft */ }
#line 2 "src/fft/engine.hpp"
#line 7 "src/fft/engine.hpp"
#line 9 "src/fft/engine.hpp"
namespace ecnerwala::fft {
// ==== engine concept ====
// Output operations for the finish step to express arbitrary fusion into the output buffer.
struct assign_op { template <typename T> void operator()(T& d, T v) const { d = v; } };
struct add_op { template <typename T> void operator()(T& d, T v) const { d += v; } };
struct sub_op { template <typename T> void operator()(T& d, T v) const { d -= v; } };
struct add_twice_op { template <typename T> void operator()(T& d, T v) const { d += v + v; } };
// `engine` contract
// engine represents a way of packing/unpacking sequences over an arbitrary ring into FFT-style transforms.
// We expect transforms/products of transforms to be linear but potentially lossy/imprecise, so we'll track precision
// at compile-time as a template parameter.
//
// E::value_type The ring we operate over
// E::unit_scale 0 or 1 depending on whether there's error that can accumulate
// E::commutative A marker for whether the ring is commutative
//
// transformed_t<A> The transform of a sequence. This object owns its data buffer.
// product_t<A> The product of 2 transforms. May equal transformed_t, particularly when unit_scale = 0.
//
// transformed alias for transformed_t<unit_scale>
// product alias for product_t<unit_scale>
//
// The basic multiplication API is
// transform(span<const value_type> in, int n) -> transformed_t<unit_scale>
// mul(transformed_t<A>, transformed_t<B>, int n) -> product_t<A*B>
// mul2(a1, b1, a2, b2, int n) -> product_t<A1*B1 + A2*B2>, computing a1*b1 + a2*b2 in one pass
// finish(product_t<A>, span<value_type>& out, Op) -> void
//
// Input span can be length up to 2n.
// Output spans can be length up to n; only the prefix that exists is filled.
// finish applies Op exactly once per out element, in index order, to
// value_type targets (so ops may be stateful).
// Transforms can be longer than necessary, and only the relevant prefix is used.
//
// For non-exact engines, there's some subtlety in whether we wrap before or after packing.
// We will choose to wrap *after* packing, which hurts error bounds but makes the prefix condition more uniform.
//
// Additionally, we have APIs to take advantage of linearity in both transformed and product space:
// add(transformed_t<A>, transformed_t<B>) -> transformed_t<A+B>
// add(product_t<K1>, product_t<K2>) -> product_t<K1+K2>
//
// Finally, we expose some additional fast-transform optimization paths.
// extend_to only operates on transformed_t<unit_scale>; the others are scale-generic.
// downsample is also defined on product_t (halving before finish saves inverse-transform work).
// extend_to build (if empty) or grow a transform to size m by repeated doubling; feed the same coefficients (sz <= 2m) every time
// (each doubling step reads only the coefficients that fit, so zero-padded buffers are fine:
// coefficients past twice the existing transform's size must be zero, or it couldn't be a prefix)
// downsample compute the half-sized transform/product of just the even (odd = false) or odd terms of the input
// negate_arg size n transform of A(-x)
template <typename E>
concept engine = requires(
std::span<const typename E::value_type> in,
std::span<typename E::value_type> out,
typename E::transformed& t,
const typename E::transformed& ct,
typename E::product& p,
const typename E::product& cp,
int n
) {
typename E::value_type;
{ E::transform(in, n) } -> std::same_as<typename E::transformed>;
{ ct.size() } -> std::same_as<int>;
E::extend_to(t, n, in);
{ E::downsample(ct, n, false) } -> std::same_as<typename E::transformed>;
{ E::downsample(cp, n, false) } -> std::same_as<typename E::product>;
{ E::negate_arg(ct, n) } -> std::same_as<typename E::transformed>;
{ E::mul(ct, ct, n) } -> std::same_as<typename E::product>;
{ E::sq(ct, n) } -> std::same_as<typename E::product>;
{ E::mul2(ct, ct, ct, ct, n) } -> std::same_as<typename E::template product_t<2 * E::unit_scale>>;
E::finish(std::move(p), out);
E::finish(std::move(p), out, add_op{});
E::finish(E::add(std::move(p), std::move(p)), out);
{ E::add(E::transform(in, n), ct) } -> std::same_as<typename E::template transformed_t<2 * E::unit_scale>>;
{ E::add(std::move(p), std::move(p)) } -> std::same_as<typename E::template product_t<2 * E::unit_scale>>;
requires std::same_as<std::remove_cvref_t<decltype(E::commutative)>, bool>;
requires std::same_as<std::remove_cvref_t<decltype(E::unit_scale)>, int>;
};
// Constrains two engine-parameterized value types to share the same engine.
template <typename A, typename B>
concept same_engine = std::same_as<typename A::engine_t, typename B::engine_t>;
// short spelling for E::transformed at use sites
template <engine E> using transformed = typename E::transformed;
/* namespace ecnerwala::fft */ }
#line 13 "src/fft/engines/split.hpp"
namespace ecnerwala::fft::engines {
// Multiplies mod `mnum` by splitting values into balanced 15-bit halves (each limb in
// [-2^14, 2^14], from the balanced representative |v| <= MOD/2) packed into one complex
// transform per operand.
template <typename mnum> struct split {
static_assert(sizeof(decltype(mnum::MOD)) <= 4, "limbs must fit 15 bits");
using value_type = mnum;
static constexpr bool commutative = true;
static constexpr int unit_scale = 1;
using cnum = cplx<double>;
using core = fft_core<cnum>;
template <int A = 1> struct transformed_t {
vector<cnum> v;
int size() const { return sz(v); }
transformed_t() = default;
explicit transformed_t(vector<cnum>&& v_) : v(std::move(v_)) {}
template <int A2> requires (A2 != A) explicit(A2 > A) transformed_t(transformed_t<A2>&& o)
: v(std::move(o.v)) {}
};
using transformed = transformed_t<1>;
template <int K> struct product_t {
// After finish's inverse transforms: lo = (lo*lo, hi*lo), hi = (lo*hi, hi*hi).
vector<cnum> lo, hi;
int size() const { return sz(lo); }
product_t() = default;
product_t(vector<cnum>&& lo_, vector<cnum>&& hi_) : lo(std::move(lo_)), hi(std::move(hi_)) {}
template <int K2> requires (K2 != K) explicit(K2 > K) product_t(product_t<K2>&& o)
: lo(std::move(o.lo)), hi(std::move(o.hi)) {}
};
using product = product_t<1>;
static cnum pack(mnum x) {
int64_t v = x.balanced();
int64_t hi = (v + (1 << 14)) >> 15;
return cnum(double(v - (hi << 15)), double(hi));
}
static transformed transform(std::span<const mnum> a, int n) {
assert(sz(a) <= 2 * n);
transformed r;
r.v.assign(n, cnum(0));
for (int i = 0; i < sz(a); i++) {
int j = i < n ? i : i - n;
r.v[j] = r.v[j] + pack(a[i]);
}
core::forward(std::span<cnum>(r.v));
return r;
}
static void extend_to(transformed& t, int m, std::span<const mnum> coeffs) {
assert(!(m & (m-1)) && sz(coeffs) <= 2 * m);
if (t.size() >= m) return;
if (t.size() == 0) { t = transform(coeffs, m); return; }
auto buf = buffer_pool<cnum>::get(sz(coeffs));
for (int i = 0; i < sz(coeffs); i++) buf[i] = pack(coeffs[i]);
while (t.size() < m) {
int s = t.size();
t.v.resize(2 * s);
// coeffs past 2s are zero: they didn't fit in the transform we're a prefix of
core::extend(
std::span<cnum>(t.v),
std::span<const cnum>(buf.span()).first(size_t(min(sz(coeffs), 2 * s)))
);
}
}
static void downsample_core(std::span<const cnum> in, std::span<cnum> out, bool odd) {
if (odd) core::odd_half(in, out);
else core::even_half(in, out);
}
template <int A> static transformed_t<A> downsample(const transformed_t<A>& t, int n, bool odd) {
transformed_t<A> r; r.v.resize(n);
downsample_core(std::span<const cnum>(t.v), std::span<cnum>(r.v), odd);
return r;
}
template <int K> static product_t<K> downsample(const product_t<K>& p, int n, bool odd) {
product_t<K> r; r.lo.resize(n); r.hi.resize(n);
downsample_core(std::span<const cnum>(p.lo), std::span<cnum>(r.lo), odd);
downsample_core(std::span<const cnum>(p.hi), std::span<cnum>(r.hi), odd);
return r;
}
template <int A> static transformed_t<A> negate_arg(const transformed_t<A>& t, int n) {
assert(n >= 2 && t.size() >= n);
transformed_t<A> r; r.v.resize(n);
for (int j = 0; j < n; j++) r.v[j] = t.v[j ^ 1];
return r;
}
template <int A, int B> static transformed_t<A + B> add(transformed_t<A>&& a, const transformed_t<B>& b) {
transformed_t<A + B> r{std::move(a.v)};
add_into(r.v, b.v);
return r;
}
// Unpacks b's transform into transforms of its low/high halves via conjugate
// symmetry, then multiplies both against a's (still packed) transform. The scale
// parameter only affects the bookkeeping, so the body is a shared untyped impl.
static void mul_impl(const vector<cnum>& a, const vector<cnum>& b, vector<cnum>& lo, vector<cnum>& hi, int n, bool acc = false) {
core::init(n);
lo.resize(n); hi.resize(n);
for (int i = 0; i < n; i++) {
int ci = core::conj_index(i);
cnum g0 = (b[i] + conj(b[ci])) * cnum(0.5);
cnum t = (b[i] - conj(b[ci])) * cnum(0.5);
cnum g1 = cnum(t.y, -t.x);
if (acc) {
lo[i] = lo[i] + a[i] * g0;
hi[i] = hi[i] + a[i] * g1;
} else {
lo[i] = a[i] * g0;
hi[i] = a[i] * g1;
}
}
}
template <int A, int B> static product_t<A * B> mul(const transformed_t<A>& a, const transformed_t<B>& b, int n) {
assert(a.size() >= n && b.size() >= n);
product_t<A * B> p;
mul_impl(a.v, b.v, p.lo, p.hi, n);
return p;
}
template <int A> static product_t<A * A> sq(const transformed_t<A>& a, int n) { return mul(a, a, n); }
template <int A1, int B1, int A2, int B2>
static product_t<A1 * B1 + A2 * B2> mul2(
const transformed_t<A1>& a1, const transformed_t<B1>& b1,
const transformed_t<A2>& a2, const transformed_t<B2>& b2,
int n
) {
assert(a1.size() >= n && b1.size() >= n && a2.size() >= n && b2.size() >= n);
product_t<A1 * B1 + A2 * B2> p;
mul_impl(a1.v, b1.v, p.lo, p.hi, n);
mul_impl(a2.v, b2.v, p.lo, p.hi, n, true);
return p;
}
static void add_into(vector<cnum>& a, const vector<cnum>& b) {
assert(sz(a) == sz(b));
for (int i = 0; i < sz(a); i++) a[i] = a[i] + b[i];
}
template <int K1, int K2> static product_t<K1 + K2> add(product_t<K1>&& a, product_t<K2>&& b) {
product_t<K1 + K2> r{std::move(a.lo), std::move(a.hi)};
add_into(r.lo, b.lo);
add_into(r.hi, b.hi);
return r;
}
template <int K = 1, typename Op = assign_op> static void finish(product_t<K>&& p, std::span<mnum> out, Op op = {}) {
// The fp error budget is divided by the accumulated scale; K <= 2 is very
// conservative (balanced limbs already left ~2x headroom at max lengths).
static_assert(K <= 2, "split: accumulated scale too large");
int n = p.size();
assert(sz(out) <= n);
core::inverse(std::span<cnum>(p.lo));
core::inverse(std::span<cnum>(p.hi));
const int64_t m = mnum::MOD;
double d = 1.0 / double(n);
// llround + a final wrap so negative half-products (e.g. from negate_arg'd
// transforms) reconstruct correctly.
for (int i = 0; i < sz(out); i++) {
int64_t v = (llround(p.lo[i].x * d)
+ (llround(p.lo[i].y * d) % m << 15)
+ (llround(p.hi[i].x * d) % m << 15)
+ (llround(p.hi[i].y * d) % m << 30)) % m;
if (v < 0) v += m;
op(out[i], mnum(v));
}
}
};
/* namespace ecnerwala::fft::engines */ }
// clang-format off
// @formatter:off
#pragma GCC diagnostic push
#pragma GCC diagnostic ignored "-Wpragmas"
#pragma GCC diagnostic ignored "-Wunknown-warning-option"
#pragma GCC diagnostic ignored "-Wmisleading-indentation"
#pragma GCC diagnostic ignored "-Wmultistatement-macros"
#include <bits/stdc++.h>
#include <cassert>
// src/fft/common.hpp
namespace ecnerwala{
template<class T>int sz(T&&arg){using std::size;return int(size(std::forward<T>(arg)));}
inline int nextPow2(int s){return 1<<(s>1?32-__builtin_clz(s-1):0);}
namespace fft{
using std::swap;
using std::vector;
using std::min;
using std::max;
#ifndef ECNERWALA_FFT_POOL_STORAGE
#define ECNERWALA_FFT_POOL_STORAGE
#endif
template<typename T>struct buffer_pool{
static inline ECNERWALA_FFT_POOL_STORAGE std::vector<std::vector<T>>free_list;
struct handle{
std::vector<T>v;
explicit handle(int n){
if(!free_list.empty()){
v=std::move(free_list.back());
free_list.pop_back();
}
v.assign(n,T());
}
handle(const handle&)=delete;
handle&operator=(const handle&)=delete;
handle(handle&&o)noexcept:v(std::move(o.v)){}
~handle(){if(v.capacity())free_list.push_back(std::move(v));}
T&operator[](int i){return v[i];}
operator std::span<T>(){return std::span<T>(v);}
std::span<T>span(){return std::span<T>(v);}
};
static handle get(int n){return handle(n);}
};
}
}
// src/modnum.hpp
template<typename T>T mod_inv_in_range(T a,T m){
T x=a,y=m;
T vx=1,vy=0;
bool swap=false;
while(x){
T k=y/x;
y%=x;
vy+=k*vx;
std::swap(x,y);
std::swap(vx,vy);
swap^=1;
}
assert(y==1);
return swap?vy:m-vy;
}
template<typename T>struct extended_gcd_result{
T gcd;
T coeff_a,coeff_b;
};
template<typename T>extended_gcd_result<T>extended_gcd(T a,T b){
T x=a,y=b;
T ax=1,ay=0;
T bx=0,by=1;
while(x){
T k=y/x;
y%=x;
ay-=k*ax;
by-=k*bx;
std::swap(x,y);
std::swap(ax,ay);
std::swap(bx,by);
}
return{y,ay,by};
}
template<typename T>T mod_inv(T a,T m){
a%=m;
a=a<0?a+m:a;
return mod_inv_in_range(a,m);
}
template<typename Self>
struct num_ops{
Self operator+()const{return static_cast<const Self&>(*this);}
Self operator-()const{return static_cast<const Self&>(*this).neg();}
friend Self operator++(Self&a,int){Self r=a;++a;return r;}
friend Self operator--(Self&a,int){Self r=a;--a;return r;}
friend Self operator+(const Self&a,const Self&b){return Self(a)+=b;}
friend Self operator-(const Self&a,const Self&b){return Self(a)-=b;}
friend Self operator*(const Self&a,const Self&b){return Self(a)*=b;}
friend Self operator/(const Self&a,const Self&b){return Self(a)/=b;}
friend bool operator!=(const Self&a,const Self&b){return!(a==b);}
friend Self neg(const Self&a){return a.neg();}
friend Self inv(const Self&a){return a.inv();}
};
template<typename Self,typename V>
struct mod_ops:num_ops<Self>{
static_assert(std::unsigned_integral<V>);
V v;
struct is_reduced_tag{};
mod_ops():v(0){}
mod_ops(V v_,is_reduced_tag):v(v_){assert(v<Self::MOD);}
template<std::integral I>mod_ops(I x):v(Self::reduce(x)){}
static Self from_reduced(V v){return Self(v,is_reduced_tag{});}
static V reduce(std::signed_integral auto x){
using U=std::make_unsigned_t<decltype(x)>;
return x<0?V(Self::MOD-1-Self::reduce(U(~x))):Self::reduce(U(x));
}
explicit operator V()const{return v;}
std::make_signed_t<V>balanced()const{
return std::make_signed_t<V>(Self::MOD-v>v?v:v-Self::MOD);
}
friend bool operator==(const Self&a,const Self&b){return a.v==b.v;}
friend std::ostream&operator<<(std::ostream&out,const Self&n){return out<<n.v;}
friend std::istream&operator>>(std::istream&in,Self&n){int64_t v_;in>>v_;n=Self(v_);return in;}
Self&operator++(){
++v;
if(v==Self::MOD)v=0;
return self();
}
Self&operator--(){
if(v==0)v=Self::MOD;
--v;
return self();
}
Self&operator+=(const Self&o){v=Self::sub_mod_raw(v,Self::MOD-o.v);return self();}
Self&operator-=(const Self&o){v=Self::sub_mod_raw(v,o.v);return self();}
Self&operator/=(const Self&o){return self()*=o.inv();}
static V sub_mod_raw(V a,V b){return a<b?a-b+Self::MOD:a-b;}
Self neg()const{return from_reduced(v?Self::MOD-v:0);}
Self inv()const{return from_reduced(mod_inv_in_range(v,Self::MOD));}
private:
Self&self(){return static_cast<Self&>(*this);}
};
template<auto MOD_>struct modnum:mod_ops<modnum<MOD_>,std::make_unsigned_t<decltype(MOD_)>>{
using Self=modnum;
static_assert(MOD_>0,"MOD must be positive");
using V=std::make_unsigned_t<decltype(MOD_)>;
static constexpr V MOD=V(MOD_);
using base=mod_ops<modnum,V>;
using base::base;
using base::v;
using base::reduce;
static V reduce(std::unsigned_integral auto x){return V(x%MOD);}
explicit operator std::make_signed_t<V>()const
requires(MOD<=V(std::numeric_limits<std::make_signed_t<V>>::max()))
{
return std::make_signed_t<V>(v);
}
Self&operator*=(const Self&o){
if constexpr(sizeof(V)<=4)v=V(uint64_t(v)*o.v%MOD);
else v=V(__uint128_t(v)*o.v%MOD);
return*this;
}
};
struct mod_goldilocks:mod_ops<mod_goldilocks,uint64_t>{
using Self=mod_goldilocks;
static constexpr uint64_t MOD=0xffffffff00000001ull;
static constexpr uint64_t EPS=-MOD;
static constexpr uint64_t PRIMITIVE_ROOT=2717;
using base=mod_ops<mod_goldilocks,uint64_t>;
using base::base;
using base::reduce;
mod_goldilocks()=default;
mod_goldilocks(__int128_t a):base(a<0?uint64_t(MOD-1-__uint128_t(~a)%MOD):uint64_t(__uint128_t(a)%MOD),is_reduced_tag{}){}
mod_goldilocks(__uint128_t a):base(uint64_t(a%MOD),is_reduced_tag{}){}
static uint64_t reduce(std::unsigned_integral auto x){
static_assert(sizeof(x)<=8);
uint64_t a=x;
return a>=MOD?a-MOD:a;
}
static uint64_t sub_mod_raw(uint64_t a,uint64_t b){
#if defined(__x86_64__)
uint64_t res_wrapped=a;
uint64_t adjustment=b;
asm(
"sub %[y], %[x]\n\t"
"sbb %k[y], %k[y]\n\t"
:[x]"+r"(res_wrapped),
[y]"+r"(adjustment)
:
:"cc"
);
#else
uint64_t res_wrapped=a-b;
uint64_t adjustment=(res_wrapped>a)?EPS:0;
#endif
return res_wrapped-adjustment;
}
static uint64_t reduce_u160_raw(uint64_t lo,uint32_t mi,uint64_t hi){
return sub_mod_raw(sub_mod_raw(lo,hi),MOD-(uint64_t(mi)<<32)+mi);
}
static uint64_t reduce_u128_raw(__uint128_t v){
uint64_t hi=uint64_t(v>>64);
uint64_t lo=uint64_t(v);
uint32_t hi_hi=uint32_t(hi>>32);
uint32_t hi_lo=uint32_t(hi);
return reduce_u160_raw(lo,hi_lo,hi_hi);
}
Self&operator*=(Self o){
v=reduce_u128_raw(__uint128_t(v)*__uint128_t(o.v));
return*this;
}
};
template<typename T>T power(T a,long long b){
assert(b>=0);
T r=1;while(b){if(b&1)r*=a;b>>=1;a*=a;}return r;
}
template<typename U,typename V>struct pairnum:num_ops<pairnum<U,V>>{
using Self=pairnum;
U u;
V v;
pairnum():u(0),v(0){}
pairnum(long long val):u(val),v(val){}
pairnum(const U&u_,const V&v_):u(u_),v(v_){}
friend std::ostream&operator<<(std::ostream&out,const Self&n){return out<<'('<<n.u<<','<<' '<<n.v<<')';}
friend std::istream&operator>>(std::istream&in,Self&n){long long val;in>>val;n=Self(val);return in;}
friend bool operator==(const Self&a,const Self&b){return a.u==b.u&&a.v==b.v;}
Self inv()const{
return Self(u.inv(),v.inv());
}
Self neg()const{
return Self(u.neg(),v.neg());
}
Self&operator++(){
++u,++v;
return*this;
}
Self&operator--(){
--u,--v;
return*this;
}
Self&operator+=(const Self&o){
u+=o.u;
v+=o.v;
return*this;
}
Self&operator-=(const Self&o){
u-=o.u;
v-=o.v;
return*this;
}
Self&operator*=(const Self&o){
u*=o.u;
v*=o.v;
return*this;
}
Self&operator/=(const Self&o){
u/=o.u;
v/=o.v;
return*this;
}
};
template<typename tag>struct dynamic_modnum:mod_ops<dynamic_modnum<tag>,uint32_t>{
using Self=dynamic_modnum;
private:
inline static uint32_t MOD_=0;
inline static uint64_t BARRETT_M=0;
public:
static constexpr uint32_t const&MOD=MOD_;
using base=mod_ops<dynamic_modnum,uint32_t>;
using base::base;
using base::v;
using base::reduce;
static void set_mod(int mod){
assert(mod>0);
MOD_=uint32_t(mod);
BARRETT_M=(uint64_t(-1)/MOD);
}
static uint32_t barrett_reduce_partial(uint64_t a){
return uint32_t(a-uint64_t((__uint128_t(BARRETT_M)*a)>>64)*MOD);
}
static uint32_t barrett_reduce(uint64_t a){
int32_t res=int32_t(barrett_reduce_partial(a)-MOD);
return uint32_t((res<0)?res+int32_t(MOD):res);
}
struct mod_reader{
friend std::istream&operator>>(std::istream&i,mod_reader){
int mod;i>>mod;
Self::set_mod(mod);
return i;
}
};
static mod_reader MOD_READER(){
return mod_reader();
}
static uint32_t reduce(std::unsigned_integral auto x){
static_assert(sizeof(x)<=8);
return barrett_reduce(x);
}
explicit operator int()const{return int(v);}
Self&operator*=(const Self&o){
v=barrett_reduce(uint64_t(v)*o.v);
return*this;
}
};
template<typename T>struct mod_constraint{
T v,mod;
friend mod_constraint operator&(mod_constraint a,mod_constraint b){
if(a.mod<b.mod)std::swap(a,b);
if(b.mod==1)return a;
extended_gcd_result<T>egcd=extended_gcd<T>(a.mod,b.mod);
assert(a.v%egcd.gcd==b.v%egcd.gcd);
T extra=b.v-a.v%b.mod;
extra/=egcd.gcd;
extra*=egcd.coeff_a;
extra%=b.mod/egcd.gcd;
extra+=(extra<0)?b.mod/egcd.gcd:0;
return mod_constraint{
a.v+extra*a.mod,
a.mod*(b.mod/egcd.gcd)
};
}
};
// src/fft/core.hpp
namespace ecnerwala::fft{
template<typename dbl>struct cplx{
dbl x,y;
cplx(dbl x_=0,dbl y_=0):x(x_),y(y_){}
friend cplx operator+(cplx a,cplx b){return cplx(a.x+b.x,a.y+b.y);}
friend cplx operator-(cplx a,cplx b){return cplx(a.x-b.x,a.y-b.y);}
friend cplx operator*(cplx a,cplx b){return cplx(a.x*b.x-a.y*b.y,a.x*b.y+a.y*b.x);}
friend cplx conj(cplx a){return cplx(a.x,-a.y);}
friend cplx inv(cplx a){dbl n=(a.x*a.x+a.y*a.y);return cplx(a.x/n,-a.y/n);}
};
template<typename num>struct getRoot{
static num f(int k)=delete;
};
template<typename dbl>struct getRoot<cplx<dbl>>{
static cplx<dbl>f(int k){
#ifndef M_PI
#define M_PI 3.14159265358979323846
#endif
dbl a=2*M_PI/k;
return cplx<dbl>(cos(a),sin(a));
}
};
template<int MOD>struct primitive_root{
static const int value;
};
template<>struct primitive_root<998244353>{
static const int value=3;
};
template<>struct primitive_root<(15<<27)+1>{
static const int value=31;
};
template<>struct primitive_root<(127<<24)+1>{
static const int value=3;
};
template<>struct primitive_root<(7<<26)+1>{
static const int value=3;
};
template<>struct primitive_root<(5<<25)+1>{
static const int value=3;
};
template<int MOD>struct getRoot<modnum<MOD>>{
static modnum<MOD>f(int k){
assert((MOD-1)%k==0);
return power(modnum<MOD>(primitive_root<MOD>::value),(MOD-1)/k);
}
};
template<>struct getRoot<mod_goldilocks>{
static mod_goldilocks f(int k){
assert((mod_goldilocks::MOD-1)%k==0);
return power(mod_goldilocks(mod_goldilocks::PRIMITIVE_ROOT),(mod_goldilocks::MOD-1)/k);
}
};
template<typename num>struct fft_core{
static inline vector<int>rev;
static inline vector<num>rt,inv_rt;
static void init(int n){
if(n<=sz(rt))return;
rev.resize(n);
for(int i=0;i<n;i++){
rev[i]=(rev[i>>1]|((i&1)*n))>>1;
}
rt.reserve(n);inv_rt.reserve(n);
while(sz(rt)<2&&sz(rt)<n){rt.push_back(num(1));inv_rt.push_back(num(1));}
for(int k=sz(rt);k<n;k*=2){
rt.resize(2*k);inv_rt.resize(2*k);
num z=getRoot<num>::f(2*k);
num iz=inv(z);
for(int i=k/2;i<k;i++){
rt[2*i]=rt[i],rt[2*i+1]=rt[i]*z;
inv_rt[2*i]=inv_rt[i],inv_rt[2*i+1]=inv_rt[i]*iz;
}
}
}
static int brev(int i,int n){
int s=__builtin_ctz(unsigned(sz(rev)/n));
return rev[i]>>s;
}
static int conj_index(int j){
return j==0?0:j^((1<<(31-__builtin_clz(unsigned(j))))-1);
}
static void forward(std::span<num>a){
int n=sz(a);
if(n<=1)return;
init(n);
for(int k=n/2;k>=1;k/=2){
for(int i=0;i<n;i+=2*k){
for(int j=0;j<k;j++){
num u=a[i+j],v=a[i+j+k];
a[i+j]=u+v;
a[i+j+k]=(u-v)*rt[j+k];
}
}
}
}
static void inverse(std::span<num>a){
int n=sz(a);
if(n<=1)return;
init(n);
for(int k=1;k<n;k*=2){
for(int i=0;i<n;i+=2*k){
for(int j=0;j<k;j++){
num t=inv_rt[j+k]*a[i+j+k];
a[i+j+k]=a[i+j]-t;
a[i+j]=a[i+j]+t;
}
}
}
}
static void extend(std::span<num>t,std::span<const num>coeffs){
int n=sz(t)/2;
assert(sz(coeffs)<=2*n);
init(sz(t));
auto b=t.subspan(n,n);
int lo=min(sz(coeffs),n);
for(int i=0;i<lo;i++){
b[i]=coeffs[i]*rt[n+i];
}
std::fill(b.begin()+lo,b.end(),num(0));
for(int i=n;i<sz(coeffs);i++){
b[i-n]=b[i-n]-coeffs[i]*rt[i];
}
forward(b);
}
static void even_half(std::span<const num>t,std::span<num>out){
int n=sz(out);
assert(sz(t)>=2*n);
num half=inv(num(2));
for(int j=0;j<n;j++)out[j]=(t[2*j]+t[2*j+1])*half;
}
static void odd_half(std::span<const num>t,std::span<num>out){
int n=sz(out);
assert(sz(t)>=2*n);
init(2*n);
num half=inv(num(2));
for(int j=0;j<n;j++){
out[j]=(t[2*j]-t[2*j+1])*half*inv_rt[n+brev(j,n)];
}
}
};
}
// src/fft/engine.hpp
namespace ecnerwala::fft{
struct assign_op{template<typename T>void operator()(T&d,T v)const{d=v;}};
struct add_op{template<typename T>void operator()(T&d,T v)const{d+=v;}};
struct sub_op{template<typename T>void operator()(T&d,T v)const{d-=v;}};
struct add_twice_op{template<typename T>void operator()(T&d,T v)const{d+=v+v;}};
template<typename E>
concept engine=requires(
std::span<const typename E::value_type>in,
std::span<typename E::value_type>out,
typename E::transformed&t,
const typename E::transformed&ct,
typename E::product&p,
const typename E::product&cp,
int n
){
typename E::value_type;
{E::transform(in,n)}->std::same_as<typename E::transformed>;
{ct.size()}->std::same_as<int>;
E::extend_to(t,n,in);
{E::downsample(ct,n,false)}->std::same_as<typename E::transformed>;
{E::downsample(cp,n,false)}->std::same_as<typename E::product>;
{E::negate_arg(ct,n)}->std::same_as<typename E::transformed>;
{E::mul(ct,ct,n)}->std::same_as<typename E::product>;
{E::sq(ct,n)}->std::same_as<typename E::product>;
{E::mul2(ct,ct,ct,ct,n)}->std::same_as<typename E::template product_t<2*E::unit_scale>>;
E::finish(std::move(p),out);
E::finish(std::move(p),out,add_op{});
E::finish(E::add(std::move(p),std::move(p)),out);
{E::add(E::transform(in,n),ct)}->std::same_as<typename E::template transformed_t<2*E::unit_scale>>;
{E::add(std::move(p),std::move(p))}->std::same_as<typename E::template product_t<2*E::unit_scale>>;
requires std::same_as<std::remove_cvref_t<decltype(E::commutative)>,bool>;
requires std::same_as<std::remove_cvref_t<decltype(E::unit_scale)>,int>;
};
template<typename A,typename B>
concept same_engine=std::same_as<typename A::engine_t,typename B::engine_t>;
template<engine E>using transformed=typename E::transformed;
}
// src/fft/engines/split.hpp
namespace ecnerwala::fft::engines{
template<typename mnum>struct split{
static_assert(sizeof(decltype(mnum::MOD))<=4,"limbs must fit 15 bits");
using value_type=mnum;
static constexpr bool commutative=true;
static constexpr int unit_scale=1;
using cnum=cplx<double>;
using core=fft_core<cnum>;
template<int A=1>struct transformed_t{
vector<cnum>v;
int size()const{return sz(v);}
transformed_t()=default;
explicit transformed_t(vector<cnum>&&v_):v(std::move(v_)){}
template<int A2>requires(A2!=A)explicit(A2>A)transformed_t(transformed_t<A2>&&o)
:v(std::move(o.v)){}
};
using transformed=transformed_t<1>;
template<int K>struct product_t{
vector<cnum>lo,hi;
int size()const{return sz(lo);}
product_t()=default;
product_t(vector<cnum>&&lo_,vector<cnum>&&hi_):lo(std::move(lo_)),hi(std::move(hi_)){}
template<int K2>requires(K2!=K)explicit(K2>K)product_t(product_t<K2>&&o)
:lo(std::move(o.lo)),hi(std::move(o.hi)){}
};
using product=product_t<1>;
static cnum pack(mnum x){
int64_t v=x.balanced();
int64_t hi=(v+(1<<14))>>15;
return cnum(double(v-(hi<<15)),double(hi));
}
static transformed transform(std::span<const mnum>a,int n){
assert(sz(a)<=2*n);
transformed r;
r.v.assign(n,cnum(0));
for(int i=0;i<sz(a);i++){
int j=i<n?i:i-n;
r.v[j]=r.v[j]+pack(a[i]);
}
core::forward(std::span<cnum>(r.v));
return r;
}
static void extend_to(transformed&t,int m,std::span<const mnum>coeffs){
assert(!(m&(m-1))&&sz(coeffs)<=2*m);
if(t.size()>=m)return;
if(t.size()==0){t=transform(coeffs,m);return;}
auto buf=buffer_pool<cnum>::get(sz(coeffs));
for(int i=0;i<sz(coeffs);i++)buf[i]=pack(coeffs[i]);
while(t.size()<m){
int s=t.size();
t.v.resize(2*s);
core::extend(
std::span<cnum>(t.v),
std::span<const cnum>(buf.span()).first(size_t(min(sz(coeffs),2*s)))
);
}
}
static void downsample_core(std::span<const cnum>in,std::span<cnum>out,bool odd){
if(odd)core::odd_half(in,out);
else core::even_half(in,out);
}
template<int A>static transformed_t<A>downsample(const transformed_t<A>&t,int n,bool odd){
transformed_t<A>r;r.v.resize(n);
downsample_core(std::span<const cnum>(t.v),std::span<cnum>(r.v),odd);
return r;
}
template<int K>static product_t<K>downsample(const product_t<K>&p,int n,bool odd){
product_t<K>r;r.lo.resize(n);r.hi.resize(n);
downsample_core(std::span<const cnum>(p.lo),std::span<cnum>(r.lo),odd);
downsample_core(std::span<const cnum>(p.hi),std::span<cnum>(r.hi),odd);
return r;
}
template<int A>static transformed_t<A>negate_arg(const transformed_t<A>&t,int n){
assert(n>=2&&t.size()>=n);
transformed_t<A>r;r.v.resize(n);
for(int j=0;j<n;j++)r.v[j]=t.v[j^1];
return r;
}
template<int A,int B>static transformed_t<A+B>add(transformed_t<A>&&a,const transformed_t<B>&b){
transformed_t<A+B>r{std::move(a.v)};
add_into(r.v,b.v);
return r;
}
static void mul_impl(const vector<cnum>&a,const vector<cnum>&b,vector<cnum>&lo,vector<cnum>&hi,int n,bool acc=false){
core::init(n);
lo.resize(n);hi.resize(n);
for(int i=0;i<n;i++){
int ci=core::conj_index(i);
cnum g0=(b[i]+conj(b[ci]))*cnum(0.5);
cnum t=(b[i]-conj(b[ci]))*cnum(0.5);
cnum g1=cnum(t.y,-t.x);
if(acc){
lo[i]=lo[i]+a[i]*g0;
hi[i]=hi[i]+a[i]*g1;
}else{
lo[i]=a[i]*g0;
hi[i]=a[i]*g1;
}
}
}
template<int A,int B>static product_t<A*B>mul(const transformed_t<A>&a,const transformed_t<B>&b,int n){
assert(a.size()>=n&&b.size()>=n);
product_t<A*B>p;
mul_impl(a.v,b.v,p.lo,p.hi,n);
return p;
}
template<int A>static product_t<A*A>sq(const transformed_t<A>&a,int n){return mul(a,a,n);}
template<int A1,int B1,int A2,int B2>
static product_t<A1*B1+A2*B2>mul2(
const transformed_t<A1>&a1,const transformed_t<B1>&b1,
const transformed_t<A2>&a2,const transformed_t<B2>&b2,
int n
){
assert(a1.size()>=n&&b1.size()>=n&&a2.size()>=n&&b2.size()>=n);
product_t<A1*B1+A2*B2>p;
mul_impl(a1.v,b1.v,p.lo,p.hi,n);
mul_impl(a2.v,b2.v,p.lo,p.hi,n,true);
return p;
}
static void add_into(vector<cnum>&a,const vector<cnum>&b){
assert(sz(a)==sz(b));
for(int i=0;i<sz(a);i++)a[i]=a[i]+b[i];
}
template<int K1,int K2>static product_t<K1+K2>add(product_t<K1>&&a,product_t<K2>&&b){
product_t<K1+K2>r{std::move(a.lo),std::move(a.hi)};
add_into(r.lo,b.lo);
add_into(r.hi,b.hi);
return r;
}
template<int K=1,typename Op=assign_op>static void finish(product_t<K>&&p,std::span<mnum>out,Op op={}){
static_assert(K<=2,"split: accumulated scale too large");
int n=p.size();
assert(sz(out)<=n);
core::inverse(std::span<cnum>(p.lo));
core::inverse(std::span<cnum>(p.hi));
const int64_t m=mnum::MOD;
double d=1.0/double(n);
for(int i=0;i<sz(out);i++){
int64_t v=(llround(p.lo[i].x*d)
+(llround(p.lo[i].y*d)%m<<15)
+(llround(p.hi[i].x*d)%m<<15)
+(llround(p.hi[i].y*d)%m<<30))%m;
if(v<0)v+=m;
op(out[i],mnum(v));
}
}
};
}
#pragma GCC diagnostic pop
// clang-format on
// @formatter:on
#pragma once
#include <algorithm>
#include <cassert>
#include <cmath>
#include <cstdint>
#include <span>
#include <utility>
#include <vector>
#include "fft/core.hpp"
#include "fft/engine.hpp"
namespace ecnerwala::fft::engines {
// Multiplies mod `mnum` by splitting values into balanced 15-bit halves (each limb in
// [-2^14, 2^14], from the balanced representative |v| <= MOD/2) packed into one complex
// transform per operand.
template <typename mnum> struct split {
static_assert(sizeof(decltype(mnum::MOD)) <= 4, "limbs must fit 15 bits");
using value_type = mnum;
static constexpr bool commutative = true;
static constexpr int unit_scale = 1;
using cnum = cplx<double>;
using core = fft_core<cnum>;
template <int A = 1> struct transformed_t {
vector<cnum> v;
int size() const { return sz(v); }
transformed_t() = default;
explicit transformed_t(vector<cnum>&& v_) : v(std::move(v_)) {}
template <int A2> requires (A2 != A) explicit(A2 > A) transformed_t(transformed_t<A2>&& o)
: v(std::move(o.v)) {}
};
using transformed = transformed_t<1>;
template <int K> struct product_t {
// After finish's inverse transforms: lo = (lo*lo, hi*lo), hi = (lo*hi, hi*hi).
vector<cnum> lo, hi;
int size() const { return sz(lo); }
product_t() = default;
product_t(vector<cnum>&& lo_, vector<cnum>&& hi_) : lo(std::move(lo_)), hi(std::move(hi_)) {}
template <int K2> requires (K2 != K) explicit(K2 > K) product_t(product_t<K2>&& o)
: lo(std::move(o.lo)), hi(std::move(o.hi)) {}
};
using product = product_t<1>;
static cnum pack(mnum x) {
int64_t v = x.balanced();
int64_t hi = (v + (1 << 14)) >> 15;
return cnum(double(v - (hi << 15)), double(hi));
}
static transformed transform(std::span<const mnum> a, int n) {
assert(sz(a) <= 2 * n);
transformed r;
r.v.assign(n, cnum(0));
for (int i = 0; i < sz(a); i++) {
int j = i < n ? i : i - n;
r.v[j] = r.v[j] + pack(a[i]);
}
core::forward(std::span<cnum>(r.v));
return r;
}
static void extend_to(transformed& t, int m, std::span<const mnum> coeffs) {
assert(!(m & (m-1)) && sz(coeffs) <= 2 * m);
if (t.size() >= m) return;
if (t.size() == 0) { t = transform(coeffs, m); return; }
auto buf = buffer_pool<cnum>::get(sz(coeffs));
for (int i = 0; i < sz(coeffs); i++) buf[i] = pack(coeffs[i]);
while (t.size() < m) {
int s = t.size();
t.v.resize(2 * s);
// coeffs past 2s are zero: they didn't fit in the transform we're a prefix of
core::extend(
std::span<cnum>(t.v),
std::span<const cnum>(buf.span()).first(size_t(min(sz(coeffs), 2 * s)))
);
}
}
static void downsample_core(std::span<const cnum> in, std::span<cnum> out, bool odd) {
if (odd) core::odd_half(in, out);
else core::even_half(in, out);
}
template <int A> static transformed_t<A> downsample(const transformed_t<A>& t, int n, bool odd) {
transformed_t<A> r; r.v.resize(n);
downsample_core(std::span<const cnum>(t.v), std::span<cnum>(r.v), odd);
return r;
}
template <int K> static product_t<K> downsample(const product_t<K>& p, int n, bool odd) {
product_t<K> r; r.lo.resize(n); r.hi.resize(n);
downsample_core(std::span<const cnum>(p.lo), std::span<cnum>(r.lo), odd);
downsample_core(std::span<const cnum>(p.hi), std::span<cnum>(r.hi), odd);
return r;
}
template <int A> static transformed_t<A> negate_arg(const transformed_t<A>& t, int n) {
assert(n >= 2 && t.size() >= n);
transformed_t<A> r; r.v.resize(n);
for (int j = 0; j < n; j++) r.v[j] = t.v[j ^ 1];
return r;
}
template <int A, int B> static transformed_t<A + B> add(transformed_t<A>&& a, const transformed_t<B>& b) {
transformed_t<A + B> r{std::move(a.v)};
add_into(r.v, b.v);
return r;
}
// Unpacks b's transform into transforms of its low/high halves via conjugate
// symmetry, then multiplies both against a's (still packed) transform. The scale
// parameter only affects the bookkeeping, so the body is a shared untyped impl.
static void mul_impl(const vector<cnum>& a, const vector<cnum>& b, vector<cnum>& lo, vector<cnum>& hi, int n, bool acc = false) {
core::init(n);
lo.resize(n); hi.resize(n);
for (int i = 0; i < n; i++) {
int ci = core::conj_index(i);
cnum g0 = (b[i] + conj(b[ci])) * cnum(0.5);
cnum t = (b[i] - conj(b[ci])) * cnum(0.5);
cnum g1 = cnum(t.y, -t.x);
if (acc) {
lo[i] = lo[i] + a[i] * g0;
hi[i] = hi[i] + a[i] * g1;
} else {
lo[i] = a[i] * g0;
hi[i] = a[i] * g1;
}
}
}
template <int A, int B> static product_t<A * B> mul(const transformed_t<A>& a, const transformed_t<B>& b, int n) {
assert(a.size() >= n && b.size() >= n);
product_t<A * B> p;
mul_impl(a.v, b.v, p.lo, p.hi, n);
return p;
}
template <int A> static product_t<A * A> sq(const transformed_t<A>& a, int n) { return mul(a, a, n); }
template <int A1, int B1, int A2, int B2>
static product_t<A1 * B1 + A2 * B2> mul2(
const transformed_t<A1>& a1, const transformed_t<B1>& b1,
const transformed_t<A2>& a2, const transformed_t<B2>& b2,
int n
) {
assert(a1.size() >= n && b1.size() >= n && a2.size() >= n && b2.size() >= n);
product_t<A1 * B1 + A2 * B2> p;
mul_impl(a1.v, b1.v, p.lo, p.hi, n);
mul_impl(a2.v, b2.v, p.lo, p.hi, n, true);
return p;
}
static void add_into(vector<cnum>& a, const vector<cnum>& b) {
assert(sz(a) == sz(b));
for (int i = 0; i < sz(a); i++) a[i] = a[i] + b[i];
}
template <int K1, int K2> static product_t<K1 + K2> add(product_t<K1>&& a, product_t<K2>&& b) {
product_t<K1 + K2> r{std::move(a.lo), std::move(a.hi)};
add_into(r.lo, b.lo);
add_into(r.hi, b.hi);
return r;
}
template <int K = 1, typename Op = assign_op> static void finish(product_t<K>&& p, std::span<mnum> out, Op op = {}) {
// The fp error budget is divided by the accumulated scale; K <= 2 is very
// conservative (balanced limbs already left ~2x headroom at max lengths).
static_assert(K <= 2, "split: accumulated scale too large");
int n = p.size();
assert(sz(out) <= n);
core::inverse(std::span<cnum>(p.lo));
core::inverse(std::span<cnum>(p.hi));
const int64_t m = mnum::MOD;
double d = 1.0 / double(n);
// llround + a final wrap so negative half-products (e.g. from negate_arg'd
// transforms) reconstruct correctly.
for (int i = 0; i < sz(out); i++) {
int64_t v = (llround(p.lo[i].x * d)
+ (llround(p.lo[i].y * d) % m << 15)
+ (llround(p.hi[i].x * d) % m << 15)
+ (llround(p.hi[i].y * d) % m << 30)) % m;
if (v < 0) v += m;
op(out[i], mnum(v));
}
}
};
/* namespace ecnerwala::fft::engines */ }