Multidimensional Convolution
(math/multivariate_convolution.hpp)
- View this file on GitHub
- Last update: 2026-08-10 17:30:05+09:00
- Include:
#include "math/multivariate_convolution.hpp"
Overview
Fast convolution of multidimensional arrays. Inputs may use a flat vector plus
explicit dimensions, or nested std::vectors whose dimensions are inferred.
In the flat representation, the first dimension is contiguous: for dimensions
n, coordinates (i[0], ..., i[k - 1]) are stored at
i[0] + i[1] * n[0] + ... + i[k - 1] * n[0] * ... * n[k - 2].
Two products are available:
- truncated convolution discards every term whose exponent reaches the bound in any variable, corresponding to reduction modulo $(x_0^{n_0}, \ldots, x_{k-1}^{n_{k-1}})$;
- cyclic convolution wraps every exponent in each variable, corresponding to reduction modulo $(1-x_0^{n_0}, \ldots, 1-x_{k-1}^{n_{k-1}})$.
API
| Function | Description | Complexity |
|---|---|---|
template <class Mint> std::vector<Mint> multivariate_convolution_truncated(const std::vector<int>& dimensions, const std::vector<Mint>& first, const std::vector<Mint>& second) |
Returns the first $n_i$ coefficients in every dimension of the ordinary multidimensional convolution, in flat form. | $O(kN\log N+k^2N)$ time and $O(kN)$ memory. |
template <class Nested> Nested multivariate_convolution_truncated(const Nested& first, const Nested& second) |
Returns the same truncated convolution as a nested vector with the same shape as the inputs. | $O(kN\log N+k^2N)$ time and $O(kN)$ memory. |
template <class Mint> std::vector<Mint> multivariate_convolution_cyclic(const std::vector<int>& dimensions, const std::vector<Mint>& first, const std::vector<Mint>& second) |
Returns the multidimensional convolution with indices wrapping around each dimension, in flat form. | $O(N\sum_i \log n_i + P(M))$ time and $O(N+\max_i n_i)$ memory when every $n_i$ divides $M-1$; otherwise $O(L\log L)$ time and $O(L)$ memory. |
template <class Nested> Nested multivariate_convolution_cyclic(const Nested& first, const Nested& second) |
Returns the same cyclic convolution as a nested vector with the same shape as the inputs. | $O(N\sum_i \log n_i + P(M))$ time and $O(N+\max_i n_i)$ memory when every $n_i$ divides $M-1$; otherwise $O(L\log L)$ time and $O(L)$ memory. |
Here, k = dimensions.size(), $n_i$ is the size of dimension i, and
$N=\prod_i n_i$, $M$ is the modulus, $L=\prod_{i:n_i>1}(2n_i-1)$, and $P(M)$
is the cost of one primitive_root(M) call. Both input arrays and the returned
array have length $N$.
Returned Values
Let first[p] mean the coefficient at multidimensional index
$p=(p_0,\ldots,p_{k-1})$, and define second[q] similarly. Every coordinate is
in the range $0\leq p_i,q_i<n_i$.
For multivariate_convolution_truncated, the returned coefficient at index $t$
is
Thus, terms with $p_i+q_i\geq n_i$ in any dimension do not appear in the returned array. The result has the original dimensions, rather than the full convolution dimensions $(2n_0-1,\ldots,2n_{k-1}-1)$.
For multivariate_convolution_cyclic, the returned coefficient at index $t$ is
The modulo is applied independently in every dimension, so overflowing indices wrap around instead of being discarded.
The flat overloads return a vector of length $N$ using the index order described
in the overview. The nested overloads return the same coefficients in a nested
vector with the same type and shape as the inputs. In particular,
result[i2][i1][i0] is the coefficient at index $(i_0,i_1,i_2)$.
Requirements and Behavior
Every dimension must be positive. An empty dimension vector represents a zero-variable polynomial, so both arrays must contain one scalar.
multivariate_convolution_truncated requires a static-modulus Mint. If S
is the smallest power of two with $S \geq 2N-1$, then S must divide
Mint::mod() - 1. In particular, m1une::math::modint998244353 supports all
sizes up to the usual $2^{22}$ coefficient limit for this routine.
multivariate_convolution_cyclic accepts either ModInt<mod> or
DynamicModInt<id>. When every $n_i$ divides Mint::mod() - 1, it uses a
multidimensional DFT and the modulus must admit a primitive root. With a static
modulus, power-of-two axes use the native NTT; other axis lengths use geometric
evaluation. Otherwise it works for arbitrary positive dimensions by embedding
each nontrivial dimension in mixed radix $2n_i-1$, taking one ordinary
convolution, and folding the result modulo the original dimensions. The
embedded input arrays end at index
$(L-1)/2$, so the ordinary convolution has exactly $L$ output coefficients and
uses the smallest power-of-two transform length that can hold them. The
fallback’s supported transform length and coefficient bound are those of
fps::convolution. For a dynamic modint, call set_mod before constructing or
reading coefficients.
No overload modifies its arguments.
For a nested input, Nested must be one or more levels of std::vector with a
modint scalar type. Both inputs must be nonempty, rectangular, and have the same
shape. The innermost vector is the first, contiguous dimension: for example,
values[i2][i1][i0] represents coordinate (i0, i1, i2). Flattening and
rebuilding the nested vectors take an additional $O(N)$ time and memory.
Example
#include "math/modint.hpp"
#include "math/multivariate_convolution.hpp"
#include <vector>
using mint = m1une::math::modint998244353;
int main() {
// Shape 2 by 2. Indices are (0,0), (1,0), (0,1), (1,1).
std::vector<int> dimensions = {2, 2};
std::vector<mint> first = {1, 2, 3, 4};
std::vector<mint> second = {5, 6, 7, 8};
std::vector<mint> truncated =
m1une::math::multivariate_convolution_truncated(
dimensions, first, second
);
std::vector<mint> cyclic =
m1une::math::multivariate_convolution_cyclic(
dimensions, first, second
);
// truncated is 5, 16, 22, 60.
// cyclic is 70, 68, 62, 60.
std::vector<std::vector<mint>> nested_first(2, std::vector<mint>(2));
nested_first[0][0] = 1;
nested_first[0][1] = 2;
nested_first[1][0] = 3;
nested_first[1][1] = 4;
std::vector<std::vector<mint>> nested_second(2, std::vector<mint>(2));
nested_second[0][0] = 5;
nested_second[0][1] = 6;
nested_second[1][0] = 7;
nested_second[1][1] = 8;
auto nested_truncated = m1une::math::multivariate_convolution_truncated(
nested_first, nested_second
);
auto nested_cyclic = m1une::math::multivariate_convolution_cyclic(
nested_first, nested_second
);
// nested_truncated has rows (5, 16) and (22, 60).
// nested_cyclic has rows (70, 68) and (62, 60).
}
Depends on
Convolution
(math/fps/convolution.hpp)
math/fps/internal/ntt998_faster.hpp
ModInt
(math/modint.hpp)
64-bit Prime Factorization
(math/prime_factorization.hpp)
Primitive Root
(math/primitive_root.hpp)
Required by
Verified with
verify/math/math_algorithms.test.cpp
verify/math/multivariate_convolution_cyclic.test.cpp
verify/math/multivariate_convolution_truncated.test.cpp
Code
#ifndef M1UNE_MATH_MULTIVARIATE_CONVOLUTION_HPP
#define M1UNE_MATH_MULTIVARIATE_CONVOLUTION_HPP 1
#include <algorithm>
#include <cassert>
#include <cstdint>
#include <limits>
#include <type_traits>
#include <utility>
#include <vector>
#include "fps/convolution.hpp"
#include "primitive_root.hpp"
namespace m1une {
namespace math {
namespace internal {
template <class T>
struct nested_vector_traits {
using scalar_type = T;
static constexpr int depth = 0;
};
template <class T, class Allocator>
struct nested_vector_traits<std::vector<T, Allocator>> {
using scalar_type = typename nested_vector_traits<T>::scalar_type;
static constexpr int depth = nested_vector_traits<T>::depth + 1;
};
template <class Nested>
void nested_vector_shape(const Nested& values, std::vector<int>& shape) {
if constexpr (nested_vector_traits<Nested>::depth > 0) {
assert(!values.empty());
assert(values.size() <= std::size_t(std::numeric_limits<int>::max()));
shape.push_back(int(values.size()));
nested_vector_shape(values.front(), shape);
}
}
template <class Nested, class Mint>
void flatten_nested_vector(
const Nested& values,
const std::vector<int>& shape,
int level,
std::vector<Mint>& flattened
) {
if constexpr (nested_vector_traits<Nested>::depth == 0) {
flattened.push_back(values);
} else {
assert(level < int(shape.size()));
assert(int(values.size()) == shape[level]);
for (const auto& child : values) {
flatten_nested_vector(child, shape, level + 1, flattened);
}
}
}
template <class Nested, class Mint>
void rebuild_nested_vector(
Nested& values,
const std::vector<int>& shape,
int level,
const std::vector<Mint>& flattened,
int& position
) {
if constexpr (nested_vector_traits<Nested>::depth == 0) {
assert(position < int(flattened.size()));
values = flattened[position++];
} else {
assert(level < int(shape.size()));
values.resize(shape[level]);
for (auto& child : values) {
rebuild_nested_vector(child, shape, level + 1, flattened, position);
}
}
}
template <class Nested>
std::vector<int> flatten_multivariate_inputs(
const Nested& first,
const Nested& second,
std::vector<typename nested_vector_traits<Nested>::scalar_type>& flattened_first,
std::vector<typename nested_vector_traits<Nested>::scalar_type>& flattened_second
) {
std::vector<int> shape;
nested_vector_shape(first, shape);
assert(int(shape.size()) == nested_vector_traits<Nested>::depth);
std::vector<int> second_shape;
nested_vector_shape(second, second_shape);
assert(second_shape == shape);
flatten_nested_vector(first, shape, 0, flattened_first);
flatten_nested_vector(second, shape, 0, flattened_second);
std::reverse(shape.begin(), shape.end());
return shape;
}
template <class Nested>
Nested rebuild_multivariate_result(
std::vector<int> dimensions,
const std::vector<typename nested_vector_traits<Nested>::scalar_type>& flattened
) {
std::reverse(dimensions.begin(), dimensions.end());
Nested result;
int position = 0;
rebuild_nested_vector(result, dimensions, 0, flattened, position);
assert(position == int(flattened.size()));
return result;
}
inline int multivariate_coefficient_count(const std::vector<int>& dimensions) {
int64_t count = 1;
for (int dimension : dimensions) {
assert(dimension > 0);
count *= dimension;
assert(count <= std::numeric_limits<int>::max());
}
return int(count);
}
inline std::vector<int> multivariate_colors(const std::vector<int>& dimensions) {
const int variable_count = int(dimensions.size());
const int coefficient_count = multivariate_coefficient_count(dimensions);
std::vector<int> color(coefficient_count);
if (variable_count == 0) return color;
for (int index = 0; index < coefficient_count; index++) {
int sum = 0;
int stride = 1;
for (int variable = 0; variable + 1 < variable_count; variable++) {
stride *= dimensions[variable];
sum += index / stride;
}
color[index] = sum % variable_count;
}
return color;
}
template <class Mint>
std::vector<Mint> geometric_evaluation(
const std::vector<Mint>& polynomial, Mint ratio
) {
const int size = int(polynomial.size());
if (size <= 64) {
std::vector<Mint> result(size);
Mint point = 1;
for (int i = 0; i < size; i++) {
Mint power = 1;
for (const Mint& coefficient : polynomial) {
result[i] += coefficient * power;
power *= point;
}
point *= ratio;
}
return result;
}
auto triangular_powers = [](Mint base, int length) {
std::vector<Mint> result(length);
if (length == 0) return result;
result[0] = 1;
Mint power = 1;
for (int i = 0; i + 1 < length; i++) {
result[i + 1] = result[i] * power;
power *= base;
}
return result;
};
std::vector<Mint> positive = triangular_powers(ratio, 2 * size - 1);
std::vector<Mint> negative = triangular_powers(ratio.inv(), size);
std::vector<Mint> scaled(polynomial);
for (int i = 0; i < size; i++) scaled[i] *= negative[i];
std::reverse(scaled.begin(), scaled.end());
std::vector<Mint> product = fps::convolution(scaled, positive);
std::vector<Mint> result(size);
for (int i = 0; i < size; i++) result[i] = product[size - 1 + i] * negative[i];
return result;
}
template <class Mint>
std::vector<Mint> cyclic_fourier_transform(
std::vector<Mint> values, Mint ratio, bool inverse
) {
if constexpr (fps::internal::has_static_modulus<Mint>::value) {
const int size = int(values.size());
if ((size & (size - 1)) == 0) {
// Keep normalization outside the per-axis transforms, matching
// the arbitrary-length DFT path below.
fps::internal::ntt(values, inverse, false);
return values;
}
}
return geometric_evaluation(values, ratio);
}
} // namespace internal
template <class Mint>
std::vector<Mint> multivariate_convolution_truncated(
const std::vector<int>& dimensions,
const std::vector<Mint>& first,
const std::vector<Mint>& second
) {
static_assert(
fps::internal::has_static_modulus<Mint>::value,
"truncated multivariate convolution requires a static-modulus type"
);
const int variable_count = int(dimensions.size());
const int coefficient_count = internal::multivariate_coefficient_count(dimensions);
assert(int(first.size()) == coefficient_count);
assert(int(second.size()) == coefficient_count);
if (variable_count == 0) return {first[0] * second[0]};
int64_t transform_size_64 = 1;
while (transform_size_64 < 2LL * coefficient_count - 1) transform_size_64 <<= 1;
assert(transform_size_64 <= std::numeric_limits<int>::max());
const int transform_size = int(transform_size_64);
assert((Mint::mod() - 1) % uint32_t(transform_size) == 0);
const std::vector<int> color = internal::multivariate_colors(dimensions);
std::vector<std::vector<Mint>> transformed_first(
variable_count, std::vector<Mint>(transform_size)
);
std::vector<std::vector<Mint>> transformed_second(
variable_count, std::vector<Mint>(transform_size)
);
for (int i = 0; i < coefficient_count; i++) {
transformed_first[color[i]][i] = first[i];
transformed_second[color[i]][i] = second[i];
}
for (int group = 0; group < variable_count; group++) {
fps::internal::ntt(transformed_first[group], false);
fps::internal::ntt(transformed_second[group], false);
}
std::vector<std::vector<Mint>> transformed_result(
variable_count, std::vector<Mint>(transform_size)
);
for (int left = 0; left < variable_count; left++) {
for (int right = 0; right < variable_count; right++) {
std::vector<Mint>& destination =
transformed_result[(left + right) % variable_count];
const std::vector<Mint>& left_values = transformed_first[left];
const std::vector<Mint>& right_values = transformed_second[right];
for (int i = 0; i < transform_size; i++) {
destination[i] += left_values[i] * right_values[i];
}
}
}
for (int group = 0; group < variable_count; group++) {
fps::internal::ntt(transformed_result[group], true);
}
std::vector<Mint> result(coefficient_count);
for (int i = 0; i < coefficient_count; i++) {
result[i] = transformed_result[color[i]][i];
}
return result;
}
template <
class Nested,
std::enable_if_t<(internal::nested_vector_traits<Nested>::depth > 0), int> = 0
>
Nested multivariate_convolution_truncated(
const Nested& first,
const Nested& second
) {
using Mint = typename internal::nested_vector_traits<Nested>::scalar_type;
std::vector<Mint> flattened_first, flattened_second;
std::vector<int> dimensions = internal::flatten_multivariate_inputs(
first, second, flattened_first, flattened_second
);
std::vector<Mint> flattened_result = multivariate_convolution_truncated(
dimensions, flattened_first, flattened_second
);
return internal::rebuild_multivariate_result<Nested>(
std::move(dimensions), flattened_result
);
}
template <class Mint>
std::vector<Mint> multivariate_convolution_cyclic(
const std::vector<int>& dimensions,
const std::vector<Mint>& first,
const std::vector<Mint>& second
) {
const int coefficient_count = internal::multivariate_coefficient_count(dimensions);
assert(int(first.size()) == coefficient_count);
assert(int(second.size()) == coefficient_count);
if (dimensions.empty()) return {first[0] * second[0]};
const uint32_t modulus = Mint::mod();
bool has_all_roots = true;
for (int dimension : dimensions) {
if ((modulus - 1) % uint32_t(dimension) != 0) has_all_roots = false;
}
if (!has_all_roots) {
std::vector<int> reduced_dimensions;
for (int dimension : dimensions) {
if (dimension != 1) reduced_dimensions.push_back(dimension);
}
if (reduced_dimensions.empty()) return {first[0] * second[0]};
std::vector<int> widened_dimensions(reduced_dimensions.size());
for (int i = 0; i < int(reduced_dimensions.size()); i++) {
const int64_t widened = 2LL * reduced_dimensions[i] - 1;
assert(widened <= std::numeric_limits<int>::max());
widened_dimensions[i] = int(widened);
}
const int widened_count =
internal::multivariate_coefficient_count(widened_dimensions);
// The largest embedded input index uses coordinate dimension - 1 on
// every axis. Its double is widened_count - 1, so convolving arrays
// ending at this index produces exactly the widened mixed-radix box.
// In particular, fps::convolution chooses the smallest transform that
// contains widened_count coefficients, instead of one that contains
// 2 * widened_count - 1 coefficients due to trailing zeroes.
int64_t maximum_embedded_index = 0;
int64_t widened_stride = 1;
for (int variable = 0; variable < int(reduced_dimensions.size()); variable++) {
maximum_embedded_index +=
int64_t(reduced_dimensions[variable] - 1) * widened_stride;
widened_stride *= widened_dimensions[variable];
}
assert(widened_stride == widened_count);
assert(2 * maximum_embedded_index + 1 == widened_count);
assert(maximum_embedded_index < std::numeric_limits<int>::max());
const int embedded_input_count = int(maximum_embedded_index) + 1;
std::vector<Mint> widened_first(embedded_input_count);
std::vector<Mint> widened_second(embedded_input_count);
for (int index = 0; index < coefficient_count; index++) {
int remaining = index;
int widened_index = 0;
int embedding_stride = 1;
for (int variable = 0; variable < int(reduced_dimensions.size()); variable++) {
const int coordinate = remaining % reduced_dimensions[variable];
remaining /= reduced_dimensions[variable];
widened_index += coordinate * embedding_stride;
embedding_stride *= widened_dimensions[variable];
}
widened_first[widened_index] = first[index];
widened_second[widened_index] = second[index];
}
std::vector<Mint> widened_product =
fps::convolution(widened_first, widened_second);
assert(int(widened_product.size()) == widened_count);
std::vector<Mint> result(coefficient_count);
for (int widened_index = 0; widened_index < widened_count; widened_index++) {
int remaining = widened_index;
int index = 0;
int stride = 1;
for (int variable = 0; variable < int(reduced_dimensions.size()); variable++) {
const int coordinate = remaining % widened_dimensions[variable];
remaining /= widened_dimensions[variable];
index += (coordinate % reduced_dimensions[variable]) * stride;
stride *= reduced_dimensions[variable];
}
result[index] += widened_product[widened_index];
}
return result;
}
const uint64_t generator = primitive_root(modulus);
assert(generator != 0);
std::vector<Mint> transformed_first(first);
std::vector<Mint> transformed_second(second);
int stride = 1;
for (int dimension : dimensions) {
assert((modulus - 1) % uint32_t(dimension) == 0);
const Mint root = Mint(generator).pow((modulus - 1) / dimension);
for (int block = 0; block < coefficient_count; block += stride * dimension) {
for (int offset = 0; offset < stride; offset++) {
std::vector<Mint> first_line(dimension);
std::vector<Mint> second_line(dimension);
for (int i = 0; i < dimension; i++) {
first_line[i] = transformed_first[block + offset + stride * i];
second_line[i] = transformed_second[block + offset + stride * i];
}
first_line = internal::cyclic_fourier_transform(
std::move(first_line), root, false
);
second_line = internal::cyclic_fourier_transform(
std::move(second_line), root, false
);
for (int i = 0; i < dimension; i++) {
transformed_first[block + offset + stride * i] = first_line[i];
transformed_second[block + offset + stride * i] = second_line[i];
}
}
}
stride *= dimension;
}
for (int i = 0; i < coefficient_count; i++) {
transformed_first[i] *= transformed_second[i];
}
stride = 1;
for (int dimension : dimensions) {
const Mint inverse_root =
Mint(generator).pow((modulus - 1) / dimension).inv();
for (int block = 0; block < coefficient_count; block += stride * dimension) {
for (int offset = 0; offset < stride; offset++) {
std::vector<Mint> line(dimension);
for (int i = 0; i < dimension; i++) {
line[i] = transformed_first[block + offset + stride * i];
}
line = internal::cyclic_fourier_transform(
std::move(line), inverse_root, true
);
for (int i = 0; i < dimension; i++) {
transformed_first[block + offset + stride * i] = line[i];
}
}
}
stride *= dimension;
}
const Mint inverse_size = Mint(coefficient_count).inv();
for (Mint& value : transformed_first) value *= inverse_size;
return transformed_first;
}
template <
class Nested,
std::enable_if_t<(internal::nested_vector_traits<Nested>::depth > 0), int> = 0
>
Nested multivariate_convolution_cyclic(
const Nested& first,
const Nested& second
) {
using Mint = typename internal::nested_vector_traits<Nested>::scalar_type;
std::vector<Mint> flattened_first, flattened_second;
std::vector<int> dimensions = internal::flatten_multivariate_inputs(
first, second, flattened_first, flattened_second
);
std::vector<Mint> flattened_result = multivariate_convolution_cyclic(
dimensions, flattened_first, flattened_second
);
return internal::rebuild_multivariate_result<Nested>(
std::move(dimensions), flattened_result
);
}
} // namespace math
} // namespace m1une
#endif // M1UNE_MATH_MULTIVARIATE_CONVOLUTION_HPP#line 1 "math/multivariate_convolution.hpp"
#include <algorithm>
#include <cassert>
#include <cstdint>
#include <limits>
#include <type_traits>
#include <utility>
#include <vector>
#line 1 "math/fps/convolution.hpp"
#line 5 "math/fps/convolution.hpp"
#include <array>
#line 8 "math/fps/convolution.hpp"
#include <cstring>
#include <new>
#line 13 "math/fps/convolution.hpp"
#if defined(__GNUC__) && !defined(__clang__) && \
(defined(__x86_64__) || defined(__i386__)) && \
!defined(M1UNE_FPS_DISABLE_X86_SIMD)
#include <immintrin.h>
#define M1UNE_FPS_HAS_X86_SIMD 1
#pragma GCC push_options
#pragma GCC target("avx2,bmi")
#endif
#line 1 "math/fps/internal/ntt998_faster.hpp"
#ifdef M1UNE_FPS_HAS_X86_SIMD
#line 9 "math/fps/internal/ntt998_faster.hpp"
#include <immintrin.h>
namespace m1une {
namespace fps {
namespace internal {
namespace fast998_v2 {
// Fixed-modulus AVX2 transform with an in-register degree-8 residue product.
using u32=unsigned;
using u64=unsigned long long;
using idt=std::size_t;
using I256=__m256i;
inline void store256(void*p,I256 x){
_mm256_store_si256((I256*)p,x);
}
inline I256 load256(const void*p){
return _mm256_load_si256((const I256*)p);
}
constexpr u32 shrk(u32 x,u32 M){
return std::min(x,x-M);
}
constexpr u32 dilt(u32 x,u32 M){
return std::min(x,x+M);
}
constexpr u32 reduce(u64 x,u32 niv,u32 M){
return (x+u64(u32(x)*niv)*M)>>32;
}
constexpr u32 mul(u32 x,u32 y,u32 niv,u32 M){
return reduce(u64(x)*y,niv,M);
}
constexpr u32 mul_s(u32 x,u32 y,u32 niv,u32 M){
return shrk(reduce(u64(x)*y,niv,M),M);
}
constexpr u32 qpw(u32 a,u32 b,u32 niv,u32 M,u32 r){
for(;b;b>>=1,a=mul(a,a,niv,M)){
if(b&1){
r=mul(r,a,niv,M);
}
}
return r;
}
constexpr u32 qpw_s(u32 a,u32 b,u32 niv,u32 M,u32 r){
return shrk(qpw(a,b,niv,M,r),M);
}
inline I256 shrk32(I256 x,I256 M){
return _mm256_min_epu32(x,_mm256_sub_epi32(x,M));
}
inline I256 dilt32(I256 x,I256 M){
return _mm256_min_epu32(x,_mm256_add_epi32(x,M));
}
inline I256 Ladd32(I256 x,I256 y,I256){
return _mm256_add_epi32(x,y);
}
inline I256 Lsub32(I256 x,I256 y,I256 M){
return _mm256_add_epi32(_mm256_sub_epi32(x,y),M);
}
inline I256 add32(I256 x,I256 y,I256 M){
return shrk32(_mm256_add_epi32(x,y),M);
}
inline I256 sub32(I256 x,I256 y,I256 M){
return dilt32(_mm256_sub_epi32(x,y),M);
}
template<int msk>inline I256 neg32_m(I256 x,I256 M){
return _mm256_blend_epi32(x,_mm256_sub_epi32(M,x),msk);
}
inline I256 reduce(I256 a,I256 b,I256 niv,I256 M){
I256 c=_mm256_mul_epu32(a,niv),d=_mm256_mul_epu32(b,niv);
c=_mm256_mul_epu32(c,M),d=_mm256_mul_epu32(d,M);
return _mm256_blend_epi32(_mm256_srli_epi64(_mm256_add_epi64(a,c),32),_mm256_add_epi64(b,d),0xaa);
}
inline I256 mul(I256 a,I256 b,I256 niv,I256 M){
return reduce(_mm256_mul_epu32(a,b),_mm256_mul_epu32(_mm256_srli_epi64(a,32),_mm256_srli_epi64(b,32)),niv,M);
}
inline I256 mul_s(I256 a,I256 b,I256 niv,I256 M){
return shrk32(mul(a,b,niv,M),M);
}
inline I256 mul_bsm(I256 a,I256 b,I256 niv,I256 M){
return reduce(_mm256_mul_epu32(a,b),_mm256_mul_epu32(_mm256_srli_epi64(a,32),b),niv,M);
}
inline I256 mul_bsmfxd(I256 a,I256 b,I256 bniv,I256 M){
I256 cc=_mm256_mul_epu32(a,bniv),dd=_mm256_mul_epu32(_mm256_srli_epi64(a,32),bniv);
I256 c=_mm256_mul_epu32(a,b),d=_mm256_mul_epu32(_mm256_srli_epi64(a,32),b);
cc=_mm256_mul_epu32(cc,M),dd=_mm256_mul_epu32(dd,M);
return _mm256_blend_epi32(_mm256_srli_epi64(_mm256_add_epi64(c,cc),32),_mm256_add_epi64(d,dd),0xaa);
}
inline I256 mul_bfxd(I256 a,I256 b,I256 bniv,I256 M){
I256 cc=_mm256_mul_epu32(a,bniv),dd=_mm256_mul_epu32(_mm256_srli_epi64(a,32),_mm256_srli_epi64(bniv,32));
I256 c=_mm256_mul_epu32(a,b),d=_mm256_mul_epu32(_mm256_srli_epi64(a,32),_mm256_srli_epi64(b,32));
cc=_mm256_mul_epu32(cc,M),dd=_mm256_mul_epu32(dd,M);
return _mm256_blend_epi32(_mm256_srli_epi64(_mm256_add_epi64(c,cc),32),_mm256_add_epi64(d,dd),0xaa);
}
inline I256 mul_upd_rt(I256 a,I256 bu,I256 M){
I256 cc=_mm256_mul_epu32(a,bu),c=_mm256_mul_epu32(a,_mm256_srli_epi64(bu,32));
cc=_mm256_mul_epu32(cc,M);
return shrk32(_mm256_srli_epi64(_mm256_add_epi64(c,cc),32),M);
}
constexpr auto _mxlg=26,_lg_itth=6;
constexpr auto _itth=idt(1)<<_lg_itth;
static_assert(_lg_itth%2==0);
struct FNTT32_info{
u32 mod,mod2,niv,one,r2,r3,img,imgniv,RT1[_mxlg];
alignas(32) std::array<u32,8> rt3[_mxlg-2],rt3i[_mxlg-2],bwbr,bwb,bwbi,rt4[_mxlg-3],rt4niv[_mxlg-3],rt4i[_mxlg-3],rt4iniv[_mxlg-3],pr2,pr4,pr2niv,pr4niv,pr2i,pr2iniv,pr4i,pr4iniv;
constexpr FNTT32_info(const u32 m):mod(m),mod2(m*2),niv([&]{u32 n=2+m;for(int i=0;i<4;++i){n*=2+m*n;}return n;}()),one((-m)%m),r2((-u64(m))%m),r3(mul_s(r2,r2,niv,m)),img{},imgniv{},RT1{},rt3{},rt3i{},bwbr{},bwb{},bwbi{},rt4{},rt4niv{},rt4i{},rt4iniv{},pr2{},pr4{},pr2niv{},pr4niv{},pr2i{},pr2iniv{},pr4i{},pr4iniv{}{
const int k=__builtin_ctz(m-1);
u32 _g=mul(3,r2,niv,mod);
for(;;++_g){
if(qpw_s(_g,mod>>1,niv,mod,one)!=one){
break;
}
}
_g=qpw(_g,mod>>k,niv,mod,one);
u32 rt1[_mxlg-1],rt1i[_mxlg-1];
rt1[k-2]=_g,rt1i[k-2]=qpw(_g,mod-2,niv,mod,one);
for(int i=k-2;i>0;--i){
rt1[i-1]=mul(rt1[i],rt1[i],niv,mod);
rt1i[i-1]=mul(rt1i[i],rt1i[i],niv,mod);
}
RT1[k-1]=qpw_s(_g,3,niv,mod,one);
for(int i=k-1;i>0;--i){
RT1[i-1]=mul_s(RT1[i],RT1[i],niv,mod);
}
img=rt1[0],imgniv=img*niv;
bwbr={one,0,one,0,one};
bwb={rt1[1],0,rt1[0],0,mod-mul_s(rt1[0],rt1[1],niv,mod)};
bwbi={rt1i[1],0,rt1i[0],0,mul_s(rt1i[0],rt1i[1],niv,mod)};
u32 pr=one,pri=one;
for(int i=0;i<k-2;++i){
const u32 r=mul_s(pr,rt1[i+1],niv,mod),ri=mul_s(pri,rt1i[i+1],niv,mod);
const u32 r2=mul_s(r,r,niv,mod),r2i=mul_s(ri,ri,niv,mod);
const u32 r3=mul_s(r,r2,niv,mod),r3i=mul_s(ri,r2i,niv,mod);
rt3[i]={r*niv,r,r2*niv,r2,r3*niv,r3};
rt3i[i]={ri*niv,ri,r2i*niv,r2i,r3i*niv,r3i};
pr=mul(pr,rt1i[i+1],niv,mod),pri=mul(pri,rt1[i+1],niv,mod);
}
pr=one,pri=one;
for(int i=0;i<k-3;++i){
const u32 r=mul_s(pr,rt1[i+2],niv,mod),ri=mul_s(pri,rt1i[i+2],niv,mod);
rt4[i][0]=rt4i[i][0]=one;
for(int j=1;j<8;++j){
rt4[i][j]=mul_s(rt4[i][j-1],r,niv,mod);
rt4i[i][j]=mul_s(rt4i[i][j-1],ri,niv,mod);
}
for(int j=0;j<8;++j){
rt4niv[i][j]=rt4[i][j]*niv;
rt4iniv[i][j]=rt4i[i][j]*niv;
}
pr=mul(pr,rt1i[i+2],niv,mod),pri=mul(pri,rt1[i+2],niv,mod);
}
pr2={one,one,one,img,one,one,one,img};
pr4={one,one,one,one,one,rt1[1],img,mul_s(img,rt1[1],niv,mod)};
const u32 nr2=mod-r2,imgr2=mul_s(img,r2,niv,mod);
pr2i={nr2,nr2,nr2,imgr2,nr2,nr2,nr2,imgr2};
pr4i={one,one,one,one,one,rt1i[1],rt1i[0],mul_s(rt1i[0],rt1i[1],niv,mod)};
for(int j=0;j<8;++j){
pr2niv[j]=pr2[j]*niv,pr4niv[j]=pr4[j]*niv;
pr2iniv[j]=pr2i[j]*niv,pr4iniv[j]=pr4i[j]*niv;
}
}
};
inline void vector_dif(I256*const f,const idt n,const FNTT32_info*info){
alignas(32) std::array<u32,8> st_1[_mxlg>>1];
const I256 Mod=_mm256_set1_epi32(info->mod),Mod2=_mm256_set1_epi32(info->mod2),Niv=_mm256_set1_epi32(info->niv);
const I256 Img=_mm256_set1_epi32(info->img),ImgNiv=_mm256_set1_epi32(info->imgniv),id=_mm256_setr_epi32(0,2,0,4,0,2,0,4);
const int lgn=__builtin_ctzll(n);
std::fill(st_1,st_1+(lgn>>1),info->bwb);
const idt nn=n>>(lgn&1),m=std::min(n,_itth),mm=std::min(nn,_itth);
// I256 rr=_mm256_set1_epi32(info->one);
if(nn!=n){
for(idt i=0;i<nn;++i){
auto const p0=f+i,p1=f+nn+i;
const auto f0=load256(p0),f1=load256(p1);
const auto g0=add32(f0,f1,Mod2),g1=Lsub32(f0,f1,Mod2);
store256(p0,g0),store256(p1,g1);
}
}
for(idt L=nn>>2;L>0;L>>=2){
for(idt i=0;i<L;++i){
auto const p0=f+i,p1=p0+L,p2=p1+L,p3=p2+L;
const auto f1=load256(p1),f3=load256(p3),f2=load256(p2),f0=load256(p0);
const auto g3=mul_bsmfxd(Lsub32(f1,f3,Mod2),Img,ImgNiv,Mod),g1=add32(f1,f3,Mod2);
const auto g0=add32(f0,f2,Mod2),g2=sub32(f0,f2,Mod2);
const auto h0=add32(g0,g1,Mod2),h1=Lsub32(g0,g1,Mod2);
const auto h2=Ladd32(g2,g3,Mod2),h3=Lsub32(g2,g3,Mod2);
store256(p0,h0),store256(p1,h1),store256(p2,h2),store256(p3,h3);
}
}
for(idt j=0;j<n;j+=m){
int t=((j==0)?std::min(_lg_itth,lgn):__builtin_ctzll(j))&-2,p=(t-2)>>1;
for(idt L=(idt(1)<<t)>>2;L>=_itth;L>>=2,t-=2,--p){
auto rt=load256(st_1+p);
const auto r1=_mm256_permutevar8x32_epi32(rt,id);
const auto r1Niv=_mm256_permutevar8x32_epi32(_mm256_mul_epu32(rt,Niv),id);
rt=mul_upd_rt(rt,load256(info->rt3+__builtin_ctzll(~j>>t)),Mod);
const auto r2=_mm256_shuffle_epi32(r1,_MM_PERM_BBBB),nr3=_mm256_shuffle_epi32(r1,_MM_PERM_DDDD);
const auto r2Niv=_mm256_shuffle_epi32(r1Niv,_MM_PERM_BBBB),nr3Niv=_mm256_shuffle_epi32(r1Niv,_MM_PERM_DDDD);
store256(st_1+p,rt);
for(idt i=0;i<L;++i){
auto const p0=f+i+j,p1=p0+L,p2=p1+L,p3=p2+L;
const auto f1=load256(p1),f3=load256(p3),f2=load256(p2),f0=load256(p0);
const auto g1=mul_bsmfxd(f1,r1,r1Niv,Mod),ng3=mul_bsmfxd(f3,nr3,nr3Niv,Mod);
const auto g2=mul_bsmfxd(f2,r2,r2Niv,Mod),g0=shrk32(f0,Mod2);
const auto h3=mul_bsmfxd(Ladd32(g1,ng3,Mod2),Img,ImgNiv,Mod),h1=sub32(g1,ng3,Mod2);
const auto h0=add32(g0,g2,Mod2),h2=sub32(g0,g2,Mod2);
const auto u0=Ladd32(h0,h1,Mod2),u1=Lsub32(h0,h1,Mod2);
const auto u2=Ladd32(h2,h3,Mod2),u3=Lsub32(h2,h3,Mod2);
store256(p0,u0),store256(p1,u1),store256(p2,u2),store256(p3,u3);
}
}
I256*const g=f+j;
for(idt l=mm,L=mm>>2;L;l=L,L>>=2,t-=2,--p){
auto rt=load256(st_1+p);
for(idt i=(j==0?l:0),k=(j+i)>>t;i<m;i+=l,++k){
const auto r1=_mm256_permutevar8x32_epi32(rt,id);
const auto r2=_mm256_shuffle_epi32(r1,_MM_PERM_BBBB);
const auto nr3=_mm256_shuffle_epi32(r1,_MM_PERM_DDDD);
for(idt j=0;j<L;++j){
auto const p0=g+i+j,p1=p0+L,p2=p1+L,p3=p2+L;
const auto f1=load256(p1),f3=load256(p3),f2=load256(p2),f0=load256(p0);
const auto g1=mul_bsm(f1,r1,Niv,Mod),ng3=mul_bsm(f3,nr3,Niv,Mod);
const auto g2=mul_bsm(f2,r2,Niv,Mod),g0=shrk32(f0,Mod2);
const auto h3=mul_bsmfxd(Ladd32(g1,ng3,Mod2),Img,ImgNiv,Mod),h1=sub32(g1,ng3,Mod2);
const auto h0=add32(g0,g2,Mod2),h2=sub32(g0,g2,Mod2);
const auto u0=Ladd32(h0,h1,Mod2),u1=Lsub32(h0,h1,Mod2);
const auto u2=Ladd32(h2,h3,Mod2),u3=Lsub32(h2,h3,Mod2);
store256(p0,u0),store256(p1,u1),store256(p2,u2),store256(p3,u3);
}
rt=mul_upd_rt(rt,load256(info->rt3+__builtin_ctzll(~k)),Mod);
}
store256(st_1+p,rt);
}
// const auto pr2=load256(&info->pr2),pr4=load256(&info->pr4);
// const auto pr2Niv=load256(&info->pr2niv),pr4Niv=load256(&info->pr4niv);
// for(idt i=j;i<j+m;++i){
// auto fi=load256(f+i);
// fi=mul(fi,rr,Niv,Mod);
// rr=shrk32(mul_bfxd(rr,load256(info->rt4+__builtin_ctzll(~i)),load256(info->rt4niv+__builtin_ctzll(~i)),Mod),Mod);
// fi=mul_bfxd(Ladd32(neg32_m<0xf0>(fi,Mod2),_mm256_permute2x128_si256(fi,fi,1),Mod2),pr4,pr4Niv,Mod);
// fi=mul_bfxd(Ladd32(neg32_m<0xcc>(fi,Mod2),_mm256_shuffle_epi32(fi,0x4e),Mod2),pr2,pr2Niv,Mod);
// fi=sub32(_mm256_shuffle_epi32(fi,0xb1),neg32_m<0x55>(fi,Mod2),Mod2);
// store256(f+i,fi);
// }
}
}
template<bool shrk=false>inline void vector_dit(I256*const f,idt n,const FNTT32_info*const info){
alignas(32) std::array<u32,8> st_1[_mxlg>>1];
const I256 Mod=_mm256_set1_epi32(info->mod),Mod2=_mm256_set1_epi32(info->mod2),Niv=_mm256_set1_epi32(info->niv);
const I256 Img=_mm256_set1_epi32(info->img),ImgNiv=_mm256_set1_epi32(info->imgniv),id=_mm256_setr_epi32(0,2,0,4,0,2,0,4);
const int lgn=__builtin_ctzll(n);
std::fill(st_1,st_1+(_lg_itth>>1),info->bwbr);
std::fill(st_1+(_lg_itth>>1),st_1+(_mxlg>>1),info->bwbi);
const idt nn=n>>(lgn&1),mm=std::min(nn,_itth);
// I256 rr=_mm256_set1_epi32((info->mod-1)>>(lgn+3));
for(idt j=0;j<n;j+=mm){
// const auto pr2=load256(&info->pr2i),pr4=load256(&info->pr4i);
// const auto pr2Niv=load256(&info->pr2iniv),pr4Niv=load256(&info->pr4iniv);
// for(idt i=j;i<j+mm;++i){
// auto fi=load256(f+i);
// const auto rt=rr;
// rr=shrk32(mul_bfxd(rr,load256(info->rt4i+__builtin_ctzll(~i)),load256(info->rt4iniv+__builtin_ctzll(~i)),Mod),Mod);
// fi=mul_bfxd(Ladd32(neg32_m<0xaa>(fi,Mod2),_mm256_shuffle_epi32(fi,0xb1),Mod2),pr2,pr2Niv,Mod);
// fi=mul_bfxd(Ladd32(neg32_m<0xcc>(fi,Mod2),_mm256_shuffle_epi32(fi,0x4e),Mod2),pr4,pr4Niv,Mod);
// fi=mul(Ladd32(neg32_m<0xf0>(fi,Mod2),_mm256_permute2x128_si256(fi,fi,1),Mod2),rt,Niv,Mod);
// store256(f+i,fi);
// }
I256*const g=f+j;
int t=2,p=0;
for(idt l=4,L=1;l<=mm;L=l,l<<=2,t+=2,++p){
auto rt=load256(st_1+p);
for(idt i=0,k=j>>t;i<mm;i+=l,++k){
const auto r1=_mm256_permutevar8x32_epi32(rt,id);
const auto r2=_mm256_shuffle_epi32(r1,_MM_PERM_BBBB);
const auto r3=_mm256_shuffle_epi32(r1,_MM_PERM_DDDD);
for(idt j=0;j<L;++j){
auto const p0=g+i+j,p1=p0+L,p2=p1+L,p3=p2+L;
const auto f0=load256(p0),f1=load256(p1),f2=load256(p2),f3=load256(p3);
const auto g0=add32(f0,f1,Mod2),g1=sub32(f0,f1,Mod2);
const auto g2=add32(f2,f3,Mod2),g3=mul_bsmfxd(Lsub32(f3,f2,Mod2),Img,ImgNiv,Mod);
const auto h0=Ladd32(g0,g2,Mod2),h1=Ladd32(g1,g3,Mod2);
const auto h2=Lsub32(g0,g2,Mod2),h3=Lsub32(g1,g3,Mod2);
const auto u0=shrk32(h0,Mod2),u1=mul_bsm(h1,r1,Niv,Mod);
const auto u2=mul_bsm(h2,r2,Niv,Mod),u3=mul_bsm(h3,r3,Niv,Mod);
store256(p0,u0),store256(p1,u1),store256(p2,u2),store256(p3,u3);
}
rt=mul_upd_rt(rt,load256(info->rt3i+__builtin_ctzll(~k)),Mod);
}
store256(st_1+p,rt);
}
int tt=std::min(__builtin_ctzll(~(j>>_lg_itth))+_lg_itth,lgn);
for(idt L=_itth,l=L<<2;t<=tt;L=l,l<<=2,t+=2,++p){
if((j+_itth)==l){
if(shrk && l==n){
for(idt i=0;i<L;++i){
auto const p0=f+i,p1=p0+L,p2=p1+L,p3=p2+L;
const auto f2=load256(p2),f3=load256(p3),f0=load256(p0),f1=load256(p1);
const auto g3=mul_bsmfxd(Lsub32(f3,f2,Mod2),Img,ImgNiv,Mod),g2=add32(f2,f3,Mod2);
const auto g0=add32(f0,f1,Mod2),g1=sub32(f0,f1,Mod2);
const auto h0=add32(g0,g2,Mod2),h1=add32(g1,g3,Mod2);
const auto h2=sub32(g0,g2,Mod2),h3=sub32(g1,g3,Mod2);
const auto u0=shrk32(h0,Mod),u1=shrk32(h1,Mod);
const auto u2=shrk32(h2,Mod),u3=shrk32(h3,Mod);
store256(p0,u0),store256(p1,u1),store256(p2,u2),store256(p3,u3);
}
}
else{
for(idt i=0;i<L;++i){
auto const p0=f+i,p1=p0+L,p2=p1+L,p3=p2+L;
const auto f2=load256(p2),f3=load256(p3),f0=load256(p0),f1=load256(p1);
const auto g3=mul_bsmfxd(Lsub32(f3,f2,Mod2),Img,ImgNiv,Mod),g2=add32(f2,f3,Mod2);
const auto g0=add32(f0,f1,Mod2),g1=sub32(f0,f1,Mod2);
const auto h0=add32(g0,g2,Mod2),h1=add32(g1,g3,Mod2);
const auto h2=sub32(g0,g2,Mod2),h3=sub32(g1,g3,Mod2);
store256(p0,h0),store256(p1,h1),store256(p2,h2),store256(p3,h3);
}
}
}
else{
auto rt=load256(st_1+p);
const auto r1=_mm256_permutevar8x32_epi32(rt,id);
const auto r1Niv=_mm256_permutevar8x32_epi32(_mm256_mul_epu32(rt,Niv),id);
rt=mul_upd_rt(rt,load256(info->rt3i+__builtin_ctzll(~j>>t)),Mod);
const auto r2=_mm256_shuffle_epi32(r1,_MM_PERM_BBBB),r3=_mm256_shuffle_epi32(r1,_MM_PERM_DDDD);
const auto r2Niv=_mm256_shuffle_epi32(r1Niv,_MM_PERM_BBBB),r3Niv=_mm256_shuffle_epi32(r1Niv,_MM_PERM_DDDD);
store256(st_1+p,rt);
for(idt i=0;i<L;++i){
auto const p0=f+j+_itth-l+i,p1=p0+L,p2=p1+L,p3=p2+L;
const auto f0=load256(p0),f1=load256(p1),f2=load256(p2),f3=load256(p3);
const auto g0=add32(f0,f1,Mod2),g1=sub32(f0,f1,Mod2);
const auto g2=add32(f2,f3,Mod2),g3=mul_bsmfxd(Lsub32(f3,f2,Mod2),Img,ImgNiv,Mod);
const auto h0=Ladd32(g0,g2,Mod2),h1=Ladd32(g1,g3,Mod2);
const auto h2=Lsub32(g0,g2,Mod2),h3=Lsub32(g1,g3,Mod2);
const auto u0=shrk32(h0,Mod2),u1=mul_bsmfxd(h1,r1,r1Niv,Mod);
const auto u2=mul_bsmfxd(h2,r2,r2Niv,Mod),u3=mul_bsmfxd(h3,r3,r3Niv,Mod);
store256(p0,u0),store256(p1,u1),store256(p2,u2),store256(p3,u3);
}
}
}
}
if(shrk && nn==n && n<=_itth){
for(idt i=0;i<n;++i){
const auto f0=load256(f+i);
store256(f+i,shrk32(f0,Mod));
}
}
if(nn!=n){
for(idt i=0;i<nn;++i){
auto const p0=f+i,p1=f+nn+i;
const auto f0=load256(p0),f1=load256(p1);
const auto g0=add32(f0,f1,Mod2),g1=sub32(f0,f1,Mod2);
if constexpr(shrk){
const auto h0=shrk32(g0,Mod),h1=shrk32(g1,Mod);
store256(p0,h0),store256(p1,h1);
}
else{
store256(p0,g0),store256(p1,g1);
}
}
}
}
// Returns fx * f[0,8) * g[0,8) (mod x^8 - ww).
[[gnu::always_inline]] inline I256 convolve8(const I256*f,const I256*g,I256 ww,I256 fx,I256 Niv,I256 Mod,I256 Mod2){
const auto raa=load256(f),rbb=load256(g);
const auto taa=shrk32(raa,Mod2),bb=shrk32(mul_bsm(rbb,fx,Niv,Mod),Mod);
const auto aw=shrk32(mul_bsm(taa,ww,Niv,Mod),Mod);
const auto aa=shrk32(taa,Mod);
const auto awa=_mm256_permute2x128_si256(aa,aw,3);
const auto b0=_mm256_permute4x64_epi64(bb,0x00),b1=_mm256_shuffle_epi32(b0,_MM_PERM_CDAB);
const auto a0=aa,a1=_mm256_srli_epi64(a0,32);
const auto aw7=_mm256_alignr_epi8(aa,awa,12);
auto res00=_mm256_mul_epu32(a0,b0);
auto res01=_mm256_mul_epu32(a1,b0);
auto res10=_mm256_mul_epu32(aw7,b1);
auto res11=_mm256_mul_epu32(a0,b1);
const auto b2=_mm256_permute4x64_epi64(bb,0x55),b3=_mm256_shuffle_epi32(b2,_MM_PERM_CDAB);
const auto aw6=_mm256_alignr_epi8(aa,awa,8);
const auto aw5=_mm256_alignr_epi8(aa,awa,4);
res00=_mm256_add_epi64(res00,_mm256_mul_epu32(aw6,b2));
res01=_mm256_add_epi64(res01,_mm256_mul_epu32(aw7,b2));
res10=_mm256_add_epi64(res10,_mm256_mul_epu32(aw5,b3));
res11=_mm256_add_epi64(res11,_mm256_mul_epu32(aw6,b3));
const auto b4=_mm256_permute4x64_epi64(bb,0xaa),b5=_mm256_shuffle_epi32(b4,_MM_PERM_CDAB);
const auto aw3=_mm256_alignr_epi8(awa,aw,12);
res00=_mm256_add_epi64(res00,_mm256_mul_epu32(awa,b4));
res01=_mm256_add_epi64(res01,_mm256_mul_epu32(aw5,b4));
res10=_mm256_add_epi64(res10,_mm256_mul_epu32(aw3,b5));
res11=_mm256_add_epi64(res11,_mm256_mul_epu32(awa,b5));
const auto b6=_mm256_permute4x64_epi64(bb,0xff),b7=_mm256_shuffle_epi32(b6,_MM_PERM_CDAB);
const auto aw2=_mm256_alignr_epi8(awa,aw,8);
const auto aw1=_mm256_alignr_epi8(awa,aw,4);
res00=_mm256_add_epi64(res00,_mm256_mul_epu32(aw2,b6));
res01=_mm256_add_epi64(res01,_mm256_mul_epu32(aw3,b6));
res10=_mm256_add_epi64(res10,_mm256_mul_epu32(aw1,b7));
res11=_mm256_add_epi64(res11,_mm256_mul_epu32(aw2,b7));
res00=_mm256_add_epi64(res00,res10);
res01=_mm256_add_epi64(res01,res11);
return shrk32(reduce(res00,res01,Niv,Mod),Mod2);
}
inline void vector_convolution_direct(I256*f,const I256*g,idt lm,const FNTT32_info*const info){
u32 RR=info->one;
const auto mod=info->mod,niv=info->niv;
const auto Fx=_mm256_set1_epi32(mul_s((mod-((mod-1)>>(__builtin_ctzll(lm)))),info->r3,niv,mod));
const auto Niv=_mm256_set1_epi32(niv),Mod=_mm256_set1_epi32(mod),Mod2=_mm256_set1_epi32(info->mod2);
for(idt i=0;i<lm;++i){
store256(f+i,convolve8(f+i,g+i,_mm256_set1_epi32(RR),Fx,Niv,Mod,Mod2));
RR=mul(RR,info->RT1[__builtin_ctzll(~i)],niv,mod);
}
}
inline void vector_convolution_accumulate(I256*const result,const I256*const f,
const I256*const g,idt lm,
const FNTT32_info*const info){
u32 RR=info->one;
const auto mod=info->mod,niv=info->niv;
const auto Fx=_mm256_set1_epi32(mul_s((mod-((mod-1)>>(__builtin_ctzll(lm)))),info->r3,niv,mod));
const auto Niv=_mm256_set1_epi32(niv),Mod=_mm256_set1_epi32(mod),Mod2=_mm256_set1_epi32(info->mod2);
for(idt i=0;i<lm;++i){
const auto product=convolve8(f+i,g+i,_mm256_set1_epi32(RR),Fx,Niv,Mod,Mod2);
store256(result+i,add32(load256(result+i),product,Mod2));
RR=mul(RR,info->RT1[__builtin_ctzll(~i)],niv,mod);
}
}
} // namespace fast998_v2
} // namespace internal
} // namespace fps
} // namespace m1une
#endif // M1UNE_FPS_HAS_X86_SIMD
#line 24 "math/fps/convolution.hpp"
#ifdef M1UNE_FPS_HAS_X86_SIMD
#pragma GCC pop_options
#endif
#line 1 "math/modint.hpp"
#line 6 "math/modint.hpp"
#include <iostream>
#line 9 "math/modint.hpp"
namespace m1une {
namespace math {
template <uint32_t Modulus>
struct ModInt {
static_assert(0 < Modulus, "Modulus must be positive");
private:
uint32_t _v;
public:
static constexpr uint32_t mod() {
return Modulus;
}
static constexpr ModInt raw(uint32_t v) noexcept {
ModInt x;
x._v = v;
return x;
}
constexpr ModInt() noexcept : _v(0) {}
template <class Integer, std::enable_if_t<std::is_integral_v<Integer>, int> = 0>
constexpr ModInt(Integer v) noexcept {
if constexpr (std::is_signed_v<Integer>) {
int64_t x = static_cast<int64_t>(v) % static_cast<int64_t>(Modulus);
if (x < 0) x += Modulus;
_v = static_cast<uint32_t>(x);
} else {
_v = static_cast<uint32_t>(static_cast<uint64_t>(v) % Modulus);
}
}
constexpr uint32_t val() const noexcept {
return _v;
}
constexpr ModInt& operator++() noexcept {
_v++;
if (_v == Modulus) _v = 0;
return *this;
}
constexpr ModInt& operator--() noexcept {
if (_v == 0) _v = Modulus;
_v--;
return *this;
}
constexpr ModInt operator++(int) noexcept {
ModInt res = *this;
++*this;
return res;
}
constexpr ModInt operator--(int) noexcept {
ModInt res = *this;
--*this;
return res;
}
constexpr ModInt& operator+=(const ModInt& rhs) noexcept {
_v += rhs._v;
if (_v >= Modulus) _v -= Modulus;
return *this;
}
constexpr ModInt& operator-=(const ModInt& rhs) noexcept {
_v -= rhs._v;
if (_v >= Modulus) _v += Modulus;
return *this;
}
constexpr ModInt& operator*=(const ModInt& rhs) noexcept {
uint64_t z = _v;
z *= rhs._v;
_v = static_cast<uint32_t>(z % Modulus);
return *this;
}
constexpr ModInt& operator/=(const ModInt& rhs) noexcept {
return *this *= rhs.inv();
}
constexpr ModInt operator+(const ModInt& rhs) const noexcept {
return ModInt(*this) += rhs;
}
constexpr ModInt operator-(const ModInt& rhs) const noexcept {
return ModInt(*this) -= rhs;
}
constexpr ModInt operator*(const ModInt& rhs) const noexcept {
return ModInt(*this) *= rhs;
}
constexpr ModInt operator/(const ModInt& rhs) const noexcept {
return ModInt(*this) /= rhs;
}
constexpr bool operator==(const ModInt& rhs) const noexcept {
return _v == rhs._v;
}
constexpr bool operator!=(const ModInt& rhs) const noexcept {
return _v != rhs._v;
}
constexpr ModInt pow(long long n) const noexcept {
ModInt res = raw(1 % Modulus);
ModInt x = n < 0 ? inv() : *this;
uint64_t exponent = n < 0 ? uint64_t(-(n + 1)) + 1 : uint64_t(n);
while (exponent > 0) {
if (exponent & 1) res *= x;
x *= x;
exponent >>= 1;
}
return res;
}
constexpr ModInt inv() const noexcept {
int64_t a = _v, b = Modulus, u = 1, v = 0;
while (b) {
int64_t t = a / b;
a -= t * b;
std::swap(a, b);
u -= t * v;
std::swap(u, v);
}
assert(a == 1);
u %= Modulus;
if (u < 0) u += Modulus;
return raw(static_cast<uint32_t>(u));
}
friend std::ostream& operator<<(std::ostream& os, const ModInt& rhs) {
return os << rhs._v;
}
friend std::istream& operator>>(std::istream& is, ModInt& rhs) {
long long v;
is >> v;
rhs = ModInt(v);
return is;
}
};
using modint998244353 = ModInt<998244353>;
using modint1000000007 = ModInt<1000000007>;
template <int Id = 0>
struct DynamicModInt {
private:
uint32_t _v;
inline static uint32_t _mod = 1;
public:
static uint32_t mod() noexcept {
return _mod;
}
static void set_mod(uint32_t modulus) noexcept {
assert(modulus > 0);
assert(modulus <= uint32_t(1) << 31);
_mod = modulus;
}
static DynamicModInt raw(uint32_t v) noexcept {
assert(v < _mod);
DynamicModInt x;
x._v = v;
return x;
}
DynamicModInt() noexcept : _v(0) {}
template <class Integer, std::enable_if_t<std::is_integral_v<Integer>, int> = 0>
DynamicModInt(Integer v) noexcept {
if constexpr (std::is_signed_v<Integer>) {
int64_t x = static_cast<int64_t>(v) % static_cast<int64_t>(_mod);
if (x < 0) x += _mod;
_v = static_cast<uint32_t>(x);
} else {
_v = static_cast<uint32_t>(static_cast<uint64_t>(v) % _mod);
}
}
uint32_t val() const noexcept {
return _v;
}
DynamicModInt& operator++() noexcept {
_v++;
if (_v == _mod) _v = 0;
return *this;
}
DynamicModInt& operator--() noexcept {
if (_v == 0) _v = _mod;
_v--;
return *this;
}
DynamicModInt operator++(int) noexcept {
DynamicModInt result = *this;
++*this;
return result;
}
DynamicModInt operator--(int) noexcept {
DynamicModInt result = *this;
--*this;
return result;
}
DynamicModInt& operator+=(const DynamicModInt& rhs) noexcept {
_v += rhs._v;
if (_v >= _mod) _v -= _mod;
return *this;
}
DynamicModInt& operator-=(const DynamicModInt& rhs) noexcept {
_v -= rhs._v;
if (_v >= _mod) _v += _mod;
return *this;
}
DynamicModInt& operator*=(const DynamicModInt& rhs) noexcept {
_v = static_cast<uint32_t>(uint64_t(_v) * rhs._v % _mod);
return *this;
}
DynamicModInt& operator/=(const DynamicModInt& rhs) noexcept {
return *this *= rhs.inv();
}
DynamicModInt operator+(const DynamicModInt& rhs) const noexcept {
return DynamicModInt(*this) += rhs;
}
DynamicModInt operator-(const DynamicModInt& rhs) const noexcept {
return DynamicModInt(*this) -= rhs;
}
DynamicModInt operator*(const DynamicModInt& rhs) const noexcept {
return DynamicModInt(*this) *= rhs;
}
DynamicModInt operator/(const DynamicModInt& rhs) const noexcept {
return DynamicModInt(*this) /= rhs;
}
bool operator==(const DynamicModInt& rhs) const noexcept {
return _v == rhs._v;
}
bool operator!=(const DynamicModInt& rhs) const noexcept {
return _v != rhs._v;
}
DynamicModInt pow(long long exponent) const noexcept {
DynamicModInt result = raw(1 % _mod);
DynamicModInt base = exponent < 0 ? inv() : *this;
uint64_t magnitude =
exponent < 0 ? uint64_t(-(exponent + 1)) + 1 : uint64_t(exponent);
while (magnitude > 0) {
if (magnitude & 1) result *= base;
base *= base;
magnitude >>= 1;
}
return result;
}
DynamicModInt inv() const noexcept {
int64_t a = _v, b = _mod, u = 1, v = 0;
while (b) {
int64_t quotient = a / b;
a -= quotient * b;
std::swap(a, b);
u -= quotient * v;
std::swap(u, v);
}
assert(a == 1);
u %= _mod;
if (u < 0) u += _mod;
return raw(static_cast<uint32_t>(u));
}
friend std::ostream& operator<<(std::ostream& os, const DynamicModInt& rhs) {
return os << rhs._v;
}
friend std::istream& operator>>(std::istream& is, DynamicModInt& rhs) {
long long value;
is >> value;
rhs = DynamicModInt(value);
return is;
}
};
} // namespace math
} // namespace m1une
#line 29 "math/fps/convolution.hpp"
namespace m1une {
namespace fps {
namespace internal {
template <class Mint, class = void>
struct has_static_modulus : std::false_type {};
template <class Mint>
struct has_static_modulus<
Mint, std::void_t<decltype(std::integral_constant<uint32_t, Mint::mod()>{})>>
: std::true_type {};
constexpr uint32_t primitive_root_constexpr(uint32_t mod) {
if (mod == 2) return 1;
if (mod == 167772161) return 3;
if (mod == 469762049) return 3;
if (mod == 754974721) return 11;
if (mod == 998244353) return 3;
if (mod == 1224736769) return 3;
uint32_t divisors[32] = {};
int count = 0;
uint32_t x = mod - 1;
for (uint32_t p = 2; uint64_t(p) * p <= x; p++) {
if (x % p != 0) continue;
divisors[count++] = p;
while (x % p == 0) x /= p;
}
if (x > 1) divisors[count++] = x;
for (uint32_t g = 2;; g++) {
bool ok = true;
for (int i = 0; i < count; i++) {
uint64_t value = 1;
uint64_t base = g;
uint32_t exponent = (mod - 1) / divisors[i];
while (exponent > 0) {
if (exponent & 1) value = value * base % mod;
base = base * base % mod;
exponent >>= 1;
}
if (value == 1) {
ok = false;
break;
}
}
if (ok) return g;
}
}
constexpr int two_adic_order(uint32_t x) {
int result = 0;
while ((x & 1) == 0) {
x >>= 1;
result++;
}
return result;
}
template <class Mint>
struct NttRoots {
static constexpr int max_base = two_adic_order(Mint::mod() - 1);
std::array<Mint, max_base + 1> root;
std::array<Mint, max_base + 1> inverse_root;
std::array<Mint, max_base> rate;
std::array<Mint, max_base> inverse_rate;
std::array<Mint, max_base> rate_radix4;
std::array<Mint, max_base> inverse_rate_radix4;
NttRoots() {
constexpr uint32_t primitive_root = primitive_root_constexpr(Mint::mod());
for (int level = 1; level <= max_base; level++) {
root[level] = Mint(primitive_root).pow((Mint::mod() - 1) >> level);
inverse_root[level] = root[level].inv();
}
Mint product = 1;
Mint inverse_product = 1;
for (int i = 0; i + 1 < max_base; i++) {
rate[i] = root[i + 2] * product;
inverse_rate[i] = inverse_root[i + 2] * inverse_product;
product *= inverse_root[i + 2];
inverse_product *= root[i + 2];
}
product = 1;
inverse_product = 1;
for (int i = 0; i + 2 < max_base; i++) {
rate_radix4[i] = root[i + 3] * product;
inverse_rate_radix4[i] = inverse_root[i + 3] * inverse_product;
product *= inverse_root[i + 3];
inverse_product *= root[i + 3];
}
}
};
template <class Mint>
const NttRoots<Mint>& ntt_roots() {
static const NttRoots<Mint> roots;
return roots;
}
template <class Mint>
void ntt(std::vector<Mint>& a, bool inverse, bool normalize = true) {
const int n = int(a.size());
assert(n > 0 && (n & (n - 1)) == 0);
assert((Mint::mod() - 1) % uint32_t(n) == 0);
const auto& roots = ntt_roots<Mint>();
const int height = two_adic_order(uint32_t(n));
if (!inverse) {
int phase = 0;
while (phase < height) {
if (height - phase == 1) {
const int width = 1 << (height - phase - 1);
Mint twiddle = 1;
for (int block = 0; block < (1 << phase); block++) {
const int offset = block << (height - phase);
for (int i = 0; i < width; i++) {
const Mint left = a[offset + i];
const Mint right = a[offset + i + width] * twiddle;
a[offset + i] = left + right;
a[offset + i + width] = left - right;
}
if (block + 1 != (1 << phase))
twiddle *= roots.rate[__builtin_ctz(~uint32_t(block))];
}
phase++;
continue;
}
const int width = 1 << (height - phase - 2);
Mint twiddle = 1;
const Mint imaginary = roots.root[2];
for (int block = 0; block < (1 << phase); block++) {
const Mint twiddle2 = twiddle * twiddle;
const Mint twiddle3 = twiddle2 * twiddle;
const int offset = block << (height - phase);
for (int i = 0; i < width; i++) {
const uint64_t mod2 = uint64_t(Mint::mod()) * Mint::mod();
const uint64_t a0 = a[offset + i].val();
const uint64_t a1 = uint64_t(a[offset + i + width].val()) * twiddle.val();
const uint64_t a2 =
uint64_t(a[offset + i + 2 * width].val()) * twiddle2.val();
const uint64_t a3 =
uint64_t(a[offset + i + 3 * width].val()) * twiddle3.val();
const uint64_t a1na3i =
uint64_t(Mint(a1 + mod2 - a3).val()) * imaginary.val();
const uint64_t negative_a2 = mod2 - a2;
a[offset + i] = Mint(a0 + a2 + a1 + a3);
a[offset + i + width] = Mint(a0 + a2 + 2 * mod2 - a1 - a3);
a[offset + i + 2 * width] = Mint(a0 + negative_a2 + a1na3i);
a[offset + i + 3 * width] = Mint(a0 + negative_a2 + mod2 - a1na3i);
}
if (block + 1 != (1 << phase))
twiddle *= roots.rate_radix4[__builtin_ctz(~uint32_t(block))];
}
phase += 2;
}
} else {
int phase = height;
while (phase > 0) {
if (phase == 1) {
const int width = 1 << (height - phase);
Mint twiddle = 1;
for (int block = 0; block < (1 << (phase - 1)); block++) {
const int offset = block << (height - phase + 1);
for (int i = 0; i < width; i++) {
const Mint left = a[offset + i];
const Mint right = a[offset + i + width];
a[offset + i] = left + right;
a[offset + i + width] = (left - right) * twiddle;
}
if (block + 1 != (1 << (phase - 1)))
twiddle *= roots.inverse_rate[__builtin_ctz(~uint32_t(block))];
}
phase--;
continue;
}
const int width = 1 << (height - phase);
Mint twiddle = 1;
const Mint inverse_imaginary = roots.inverse_root[2];
for (int block = 0; block < (1 << (phase - 2)); block++) {
const Mint twiddle2 = twiddle * twiddle;
const Mint twiddle3 = twiddle2 * twiddle;
const int offset = block << (height - phase + 2);
for (int i = 0; i < width; i++) {
const uint64_t a0 = a[offset + i].val();
const uint64_t a1 = a[offset + i + width].val();
const uint64_t a2 = a[offset + i + 2 * width].val();
const uint64_t a3 = a[offset + i + 3 * width].val();
const uint64_t a2na3i =
uint64_t(Mint((Mint::mod() + a2 - a3) * inverse_imaginary.val()).val());
a[offset + i] = Mint(a0 + a1 + a2 + a3);
a[offset + i + width] =
Mint((a0 + Mint::mod() - a1 + a2na3i) * twiddle.val());
a[offset + i + 2 * width] = Mint(
(a0 + a1 + 2ULL * Mint::mod() - a2 - a3) * twiddle2.val());
a[offset + i + 3 * width] = Mint(
(a0 + Mint::mod() - a1 + Mint::mod() - a2na3i) * twiddle3.val());
}
if (block + 1 != (1 << (phase - 2)))
twiddle *= roots.inverse_rate_radix4[__builtin_ctz(~uint32_t(block))];
}
phase -= 2;
}
if (normalize) {
const Mint inverse_n = Mint(n).inv();
for (Mint& value : a) value *= inverse_n;
}
}
}
#ifdef M1UNE_FPS_HAS_X86_SIMD
#pragma GCC push_options
#pragma GCC target("avx2,bmi")
template <class Mint>
__attribute__((target("avx2,bmi"), hot))
std::vector<Mint> convolution_998244353_simd(const std::vector<Mint>& a,
const std::vector<Mint>& b) {
const int result_size = int(a.size() + b.size() - 1);
int n = 1;
while (n < result_size) n <<= 1;
const bool squaring = &a == &b;
auto* transformed_a = static_cast<uint32_t*>(
::operator new[](sizeof(uint32_t) * n, std::align_val_t(32)));
auto* transformed_b = squaring
? transformed_a
: static_cast<uint32_t*>(::operator new[](
sizeof(uint32_t) * n, std::align_val_t(32)));
if constexpr (std::is_same_v<Mint, math::ModInt<998244353>>) {
static_assert(sizeof(Mint) == sizeof(uint32_t) && std::is_trivially_copyable_v<Mint>);
std::memcpy(transformed_a, a.data(), sizeof(uint32_t) * a.size());
if (!squaring)
std::memcpy(transformed_b, b.data(), sizeof(uint32_t) * b.size());
} else {
for (int i = 0; i < int(a.size()); i++) transformed_a[i] = a[i].val();
if (!squaring)
for (int i = 0; i < int(b.size()); i++) transformed_b[i] = b[i].val();
}
std::memset(transformed_a + a.size(), 0, sizeof(uint32_t) * (n - a.size()));
if (!squaring)
std::memset(transformed_b + b.size(), 0, sizeof(uint32_t) * (n - b.size()));
static constexpr fast998_v2::FNTT32_info transform(998244353);
const std::size_t vector_size = std::size_t(n) >> 3;
fast998_v2::vector_dif(reinterpret_cast<__m256i*>(transformed_a), vector_size, &transform);
if (!squaring)
fast998_v2::vector_dif(reinterpret_cast<__m256i*>(transformed_b), vector_size,
&transform);
fast998_v2::vector_convolution_direct(
reinterpret_cast<__m256i*>(transformed_a),
reinterpret_cast<const __m256i*>(transformed_b), vector_size, &transform);
fast998_v2::vector_dit<true>(reinterpret_cast<__m256i*>(transformed_a), vector_size,
&transform);
std::vector<Mint> result(result_size);
for (int j = 0; j < result_size; j++) result[j] = Mint::raw(transformed_a[j]);
::operator delete[](transformed_a, std::align_val_t(32));
if (!squaring) ::operator delete[](transformed_b, std::align_val_t(32));
return result;
}
#pragma GCC pop_options
#endif
} // namespace internal
template <class Mint>
std::vector<Mint> convolution_naive(const std::vector<Mint>& a, const std::vector<Mint>& b) {
if (a.empty() || b.empty()) return {};
std::vector<Mint> result(a.size() + b.size() - 1);
if (a.size() < b.size()) {
for (int i = 0; i < int(a.size()); i++) {
for (int j = 0; j < int(b.size()); j++) result[i + j] += a[i] * b[j];
}
} else {
for (int j = 0; j < int(b.size()); j++) {
for (int i = 0; i < int(a.size()); i++) result[i + j] += a[i] * b[j];
}
}
return result;
}
template <class Mint>
std::vector<Mint> convolution_ntt(const std::vector<Mint>& a, const std::vector<Mint>& b) {
const int result_size = int(a.size() + b.size() - 1);
int n = 1;
while (n < result_size) n <<= 1;
assert((Mint::mod() - 1) % uint32_t(n) == 0);
#ifdef M1UNE_FPS_HAS_X86_SIMD
if constexpr (Mint::mod() == 998244353) {
if (n >= 64 && __builtin_cpu_supports("avx2"))
return internal::convolution_998244353_simd(a, b);
}
#endif
// Allocate the padded buffers directly. Constructing from the inputs and
// then resizing used to allocate and copy both large operands twice.
const bool squaring = &a == &b;
std::vector<Mint> fa(n);
std::copy(a.begin(), a.end(), fa.begin());
internal::ntt(fa, false);
const Mint inverse_n = Mint(n).inv();
if (squaring) {
for (int i = 0; i < n; i++) fa[i] *= fa[i] * inverse_n;
} else {
std::vector<Mint> fb(n);
std::copy(b.begin(), b.end(), fb.begin());
internal::ntt(fb, false);
for (int i = 0; i < n; i++) fa[i] *= fb[i] * inverse_n;
}
internal::ntt(fa, true, false);
fa.resize(result_size);
return fa;
}
namespace internal {
template <class Mint>
std::vector<Mint> convolution_998244353_blocked_scalar(const std::vector<Mint>& a,
const std::vector<Mint>& b,
int transform_size) {
assert(Mint::mod() == 998244353);
assert(transform_size >= 2 && (transform_size & (transform_size - 1)) == 0);
assert((Mint::mod() - 1) % uint32_t(transform_size) == 0);
const int block_size = transform_size / 2;
const int a_blocks = int((a.size() + block_size - 1) / block_size);
const int b_blocks = int((b.size() + block_size - 1) / block_size);
auto transform_blocks = [&](const std::vector<Mint>& values, int block_count) {
std::vector<std::vector<Mint>> blocks;
blocks.reserve(block_count);
for (int block = 0; block < block_count; block++) {
const int begin = block * block_size;
const int count = std::min(block_size, int(values.size()) - begin);
std::vector<Mint> transformed(transform_size);
std::copy_n(values.begin() + begin, count, transformed.begin());
ntt(transformed, false);
blocks.emplace_back(std::move(transformed));
}
return blocks;
};
std::vector<std::vector<Mint>> transformed_a = transform_blocks(a, a_blocks);
std::vector<std::vector<Mint>> transformed_b = transform_blocks(b, b_blocks);
const int result_size = int(a.size() + b.size() - 1);
std::vector<Mint> result(result_size);
std::vector<Mint> transformed_result(transform_size);
for (int diagonal = 0; diagonal < a_blocks + b_blocks - 1; diagonal++) {
std::fill(transformed_result.begin(), transformed_result.end(), Mint(0));
const int first_a = std::max(0, diagonal - (b_blocks - 1));
const int last_a = std::min(a_blocks - 1, diagonal);
for (int a_block = first_a; a_block <= last_a; a_block++) {
const int b_block = diagonal - a_block;
for (int i = 0; i < transform_size; i++)
transformed_result[i] +=
transformed_a[a_block][i] * transformed_b[b_block][i];
}
ntt(transformed_result, true);
const int output_offset = diagonal * block_size;
const int output_count = std::min(transform_size, result_size - output_offset);
for (int i = 0; i < output_count; i++)
result[output_offset + i] += transformed_result[i];
}
return result;
}
#ifdef M1UNE_FPS_HAS_X86_SIMD
class AlignedUint32Buffer {
private:
uint32_t* data_;
public:
explicit AlignedUint32Buffer(std::size_t size)
: data_(static_cast<uint32_t*>(
::operator new[](sizeof(uint32_t) * size, std::align_val_t(32)))) {}
AlignedUint32Buffer(const AlignedUint32Buffer&) = delete;
AlignedUint32Buffer& operator=(const AlignedUint32Buffer&) = delete;
AlignedUint32Buffer(AlignedUint32Buffer&& other) noexcept : data_(other.data_) {
other.data_ = nullptr;
}
AlignedUint32Buffer& operator=(AlignedUint32Buffer&& other) noexcept {
if (this == &other) return *this;
::operator delete[](data_, std::align_val_t(32));
data_ = other.data_;
other.data_ = nullptr;
return *this;
}
~AlignedUint32Buffer() {
::operator delete[](data_, std::align_val_t(32));
}
uint32_t* data() {
return data_;
}
const uint32_t* data() const {
return data_;
}
};
template <class Mint>
__attribute__((target("avx2,bmi"), hot))
std::vector<Mint> convolution_998244353_blocked_simd(const std::vector<Mint>& a,
const std::vector<Mint>& b,
int transform_size) {
assert(Mint::mod() == 998244353);
assert(transform_size >= 64 && (transform_size & (transform_size - 1)) == 0);
assert((Mint::mod() - 1) % uint32_t(transform_size) == 0);
const int block_size = transform_size / 2;
const int a_blocks = int((a.size() + block_size - 1) / block_size);
const int b_blocks = int((b.size() + block_size - 1) / block_size);
static constexpr fast998_v2::FNTT32_info transform(998244353);
const std::size_t vector_size = std::size_t(transform_size) / 8;
auto transform_blocks = [&](const std::vector<Mint>& values, int block_count) {
std::vector<AlignedUint32Buffer> blocks;
blocks.reserve(block_count);
for (int block = 0; block < block_count; block++) {
const int begin = block * block_size;
const int count = std::min(block_size, int(values.size()) - begin);
AlignedUint32Buffer transformed(transform_size);
if constexpr (std::is_same_v<Mint, math::ModInt<998244353>>) {
static_assert(sizeof(Mint) == sizeof(uint32_t) &&
std::is_trivially_copyable_v<Mint>);
std::memcpy(transformed.data(), values.data() + begin,
sizeof(uint32_t) * count);
} else {
for (int i = 0; i < count; i++)
transformed.data()[i] = values[begin + i].val();
}
std::memset(transformed.data() + count, 0,
sizeof(uint32_t) * (transform_size - count));
fast998_v2::vector_dif(reinterpret_cast<__m256i*>(transformed.data()),
vector_size, &transform);
blocks.emplace_back(std::move(transformed));
}
return blocks;
};
std::vector<AlignedUint32Buffer> transformed_a = transform_blocks(a, a_blocks);
std::vector<AlignedUint32Buffer> transformed_b = transform_blocks(b, b_blocks);
const int result_size = int(a.size() + b.size() - 1);
std::vector<Mint> result(result_size);
AlignedUint32Buffer transformed_result(transform_size);
for (int diagonal = 0; diagonal < a_blocks + b_blocks - 1; diagonal++) {
std::memset(transformed_result.data(), 0, sizeof(uint32_t) * transform_size);
const int first_a = std::max(0, diagonal - (b_blocks - 1));
const int last_a = std::min(a_blocks - 1, diagonal);
for (int a_block = first_a; a_block <= last_a; a_block++) {
const int b_block = diagonal - a_block;
fast998_v2::vector_convolution_accumulate(
reinterpret_cast<__m256i*>(transformed_result.data()),
reinterpret_cast<const __m256i*>(transformed_a[a_block].data()),
reinterpret_cast<const __m256i*>(transformed_b[b_block].data()),
vector_size, &transform);
}
fast998_v2::vector_dit<true>(
reinterpret_cast<__m256i*>(transformed_result.data()), vector_size,
&transform);
const int output_offset = diagonal * block_size;
const int output_count = std::min(transform_size, result_size - output_offset);
for (int i = 0; i < output_count; i++) {
uint32_t value = result[output_offset + i].val() + transformed_result.data()[i];
if (value >= Mint::mod()) value -= Mint::mod();
result[output_offset + i] = Mint::raw(value);
}
}
return result;
}
#endif
template <class Mint>
std::vector<Mint> convolution_998244353_blocked(const std::vector<Mint>& a,
const std::vector<Mint>& b,
int transform_size = 1 << 23) {
#ifdef M1UNE_FPS_HAS_X86_SIMD
if (transform_size >= 64 && __builtin_cpu_supports("avx2"))
return convolution_998244353_blocked_simd(a, b, transform_size);
#endif
return convolution_998244353_blocked_scalar(a, b, transform_size);
}
} // namespace internal
template <class Mint>
std::vector<Mint> convolution(const std::vector<Mint>& a, const std::vector<Mint>& b) {
if (a.empty() || b.empty()) return {};
if (std::min(a.size(), b.size()) <= 32) return convolution_naive(a, b);
const int result_size = int(a.size() + b.size() - 1);
int n = 1;
while (n < result_size) n <<= 1;
if constexpr (internal::has_static_modulus<Mint>::value) {
if constexpr (Mint::mod() == 998244353) {
if (n > (1 << 23))
return internal::convolution_998244353_blocked(a, b);
}
if ((Mint::mod() - 1) % uint32_t(n) == 0) return convolution_ntt(a, b);
}
using Mint1 = math::ModInt<167772161>;
using Mint2 = math::ModInt<469762049>;
using Mint3 = math::ModInt<754974721>;
assert(n <= (1 << 24));
[[maybe_unused]] const unsigned __int128 coefficient_bound =
static_cast<unsigned __int128>(std::min(a.size(), b.size())) * (Mint::mod() - 1) *
(Mint::mod() - 1);
[[maybe_unused]] const unsigned __int128 crt_modulus =
static_cast<unsigned __int128>(Mint1::mod()) * Mint2::mod() * Mint3::mod();
assert(coefficient_bound < crt_modulus);
auto converted_convolution = [&]<class OtherMint>() {
std::vector<OtherMint> converted_a(a.size());
std::vector<OtherMint> converted_b(b.size());
for (int i = 0; i < int(a.size()); i++) converted_a[i] = OtherMint(a[i].val());
for (int i = 0; i < int(b.size()); i++) converted_b[i] = OtherMint(b[i].val());
return convolution_ntt(converted_a, converted_b);
};
std::vector<Mint1> c1 = converted_convolution.template operator()<Mint1>();
std::vector<Mint2> c2 = converted_convolution.template operator()<Mint2>();
std::vector<Mint3> c3 = converted_convolution.template operator()<Mint3>();
static const uint64_t inverse_mod1_mod2 = Mint2(Mint1::mod()).inv().val();
static const uint64_t mod1_mod3 = Mint1::mod() % Mint3::mod();
static const uint64_t mod1_mod2_mod3 =
mod1_mod3 * (Mint2::mod() % Mint3::mod()) % Mint3::mod();
static const uint64_t inverse_mod1_mod2_mod3 = Mint3(uint32_t(mod1_mod2_mod3)).inv().val();
const uint64_t target_mod = Mint::mod();
const uint64_t mod1_target = Mint1::mod() % target_mod;
const uint64_t mod1_mod2_target = mod1_target * (Mint2::mod() % target_mod) % target_mod;
std::vector<Mint> result(result_size);
for (int i = 0; i < result_size; i++) {
const uint64_t r1 = c1[i].val();
const uint64_t r2 = c2[i].val();
const uint64_t r3 = c3[i].val();
const uint64_t first =
(r2 + Mint2::mod() - r1 % Mint2::mod()) % Mint2::mod() * inverse_mod1_mod2 %
Mint2::mod();
const uint64_t combined_mod3 =
(r1 % Mint3::mod() + mod1_mod3 * (first % Mint3::mod())) % Mint3::mod();
const uint64_t second =
(r3 + Mint3::mod() - combined_mod3) % Mint3::mod() * inverse_mod1_mod2_mod3 %
Mint3::mod();
uint64_t value = r1 % target_mod;
value = (value + mod1_target * (first % target_mod)) % target_mod;
value = (value + mod1_mod2_target * (second % target_mod)) % target_mod;
result[i] = Mint::raw(uint32_t(value));
}
return result;
}
} // namespace fps
} // namespace m1une
#ifdef M1UNE_FPS_HAS_X86_SIMD
#undef M1UNE_FPS_HAS_X86_SIMD
#endif
#line 1 "math/primitive_root.hpp"
#line 6 "math/primitive_root.hpp"
#include <numeric>
#line 9 "math/primitive_root.hpp"
#line 1 "math/prime_factorization.hpp"
#line 10 "math/prime_factorization.hpp"
namespace m1une {
namespace math {
namespace internal {
inline uint64_t multiply_mod(uint64_t a, uint64_t b, uint64_t mod) {
return static_cast<uint64_t>(static_cast<unsigned __int128>(a) * b % mod);
}
inline uint64_t power_mod(uint64_t base, uint64_t exponent, uint64_t mod) {
uint64_t result = 1;
while (exponent > 0) {
if (exponent & 1) result = multiply_mod(result, base, mod);
base = multiply_mod(base, base, mod);
exponent >>= 1;
}
return result;
}
inline uint64_t pollard_random() {
static uint64_t state = 0x123456789abcdef0ULL;
state += 0x9e3779b97f4a7c15ULL;
uint64_t value = state;
value = (value ^ (value >> 30)) * 0xbf58476d1ce4e5b9ULL;
value = (value ^ (value >> 27)) * 0x94d049bb133111ebULL;
return value ^ (value >> 31);
}
} // namespace internal
inline bool is_prime(uint64_t value) {
if (value < 2) return false;
for (uint64_t prime : {2ULL, 3ULL, 5ULL, 7ULL, 11ULL, 13ULL, 17ULL, 19ULL, 23ULL, 29ULL, 31ULL, 37ULL}) {
if (value % prime == 0) return value == prime;
}
uint64_t odd_part = value - 1;
int power_of_two = 0;
while ((odd_part & 1) == 0) {
odd_part >>= 1;
power_of_two++;
}
for (uint64_t base : {2ULL, 325ULL, 9375ULL, 28178ULL, 450775ULL, 9780504ULL, 1795265022ULL}) {
if (base % value == 0) continue;
uint64_t x = internal::power_mod(base % value, odd_part, value);
if (x == 1 || x == value - 1) continue;
bool composite = true;
for (int i = 1; i < power_of_two; i++) {
x = internal::multiply_mod(x, x, value);
if (x == value - 1) {
composite = false;
break;
}
}
if (composite) return false;
}
return true;
}
namespace internal {
inline uint64_t pollard_rho(uint64_t value) {
for (uint64_t prime : {2ULL, 3ULL, 5ULL, 7ULL, 11ULL, 13ULL, 17ULL, 19ULL, 23ULL, 29ULL, 31ULL, 37ULL}) {
if (value % prime == 0) return prime;
}
while (true) {
const uint64_t constant = pollard_random() % (value - 1) + 1;
uint64_t y = pollard_random() % (value - 1) + 1;
uint64_t x = 0;
uint64_t saved_y = 0;
uint64_t gcd = 1;
uint64_t segment_length = 1;
auto advance = [&](uint64_t current) {
return static_cast<uint64_t>(
(static_cast<unsigned __int128>(multiply_mod(current, current, value)) + constant) % value);
};
while (gcd == 1) {
x = y;
for (uint64_t i = 0; i < segment_length; i++) y = advance(y);
for (uint64_t offset = 0; offset < segment_length && gcd == 1; offset += 128) {
saved_y = y;
uint64_t product = 1;
const uint64_t block = std::min<uint64_t>(128, segment_length - offset);
for (uint64_t i = 0; i < block; i++) {
y = advance(y);
const uint64_t difference = x > y ? x - y : y - x;
product = multiply_mod(product, difference, value);
}
gcd = std::gcd(product, value);
}
segment_length <<= 1;
}
if (gcd == value) {
do {
saved_y = advance(saved_y);
const uint64_t difference = x > saved_y ? x - saved_y : saved_y - x;
gcd = std::gcd(difference, value);
} while (gcd == 1);
}
if (gcd != value) return gcd;
}
}
inline void factor_recursively(uint64_t value, std::vector<uint64_t>& factors) {
if (value == 1) return;
if (is_prime(value)) {
factors.push_back(value);
return;
}
const uint64_t divisor = pollard_rho(value);
factor_recursively(divisor, factors);
factor_recursively(value / divisor, factors);
}
} // namespace internal
inline std::vector<uint64_t> prime_factors(uint64_t value) {
assert(value >= 1);
std::vector<uint64_t> result;
internal::factor_recursively(value, result);
std::sort(result.begin(), result.end());
return result;
}
inline std::vector<std::pair<uint64_t, int>> prime_factorize(uint64_t value) {
std::vector<uint64_t> factors = prime_factors(value);
std::vector<std::pair<uint64_t, int>> result;
for (uint64_t prime : factors) {
if (result.empty() || result.back().first != prime) {
result.emplace_back(prime, 1);
} else {
result.back().second++;
}
}
return result;
}
inline std::vector<uint64_t> divisors(uint64_t value) {
std::vector<uint64_t> result = {1};
for (const auto& factor : prime_factorize(value)) {
const int current_size = int(result.size());
uint64_t power = 1;
for (int exponent = 1; exponent <= factor.second; exponent++) {
power *= factor.first;
for (int i = 0; i < current_size; i++) {
result.push_back(result[i] * power);
}
}
}
std::sort(result.begin(), result.end());
return result;
}
inline uint64_t euler_phi(uint64_t value) {
assert(value >= 1);
uint64_t result = value;
for (const auto& factor : prime_factorize(value)) {
result = result / factor.first * (factor.first - 1);
}
return result;
}
inline int mobius(uint64_t value) {
assert(value >= 1);
int result = 1;
for (const auto& factor : prime_factorize(value)) {
if (factor.second >= 2) return 0;
result = -result;
}
return result;
}
} // namespace math
} // namespace m1une
#line 11 "math/primitive_root.hpp"
namespace m1une {
namespace math {
inline bool has_primitive_root(uint64_t mod) {
if (mod == 2 || mod == 4) return true;
if (mod < 2) return false;
uint64_t odd_part = mod;
if ((odd_part & 1) == 0) {
odd_part >>= 1;
if ((odd_part & 1) == 0) return false;
}
return prime_factorize(odd_part).size() == 1;
}
// Returns the smallest positive primitive root modulo mod.
// Returns 0 when no primitive root exists.
inline uint64_t primitive_root(uint64_t mod) {
assert(mod >= 2);
if (mod == 2) return 1;
if (!has_primitive_root(mod)) return 0;
const uint64_t phi = euler_phi(mod);
const std::vector<std::pair<uint64_t, int>> factors = prime_factorize(phi);
for (uint64_t candidate = 2; candidate < mod; candidate++) {
if (std::gcd(candidate, mod) != 1) continue;
bool generator = true;
for (const auto& factor : factors) {
if (internal::power_mod(candidate, phi / factor.first, mod) == 1) {
generator = false;
break;
}
}
if (generator) return candidate;
}
return 0;
}
} // namespace math
} // namespace m1une
#line 14 "math/multivariate_convolution.hpp"
namespace m1une {
namespace math {
namespace internal {
template <class T>
struct nested_vector_traits {
using scalar_type = T;
static constexpr int depth = 0;
};
template <class T, class Allocator>
struct nested_vector_traits<std::vector<T, Allocator>> {
using scalar_type = typename nested_vector_traits<T>::scalar_type;
static constexpr int depth = nested_vector_traits<T>::depth + 1;
};
template <class Nested>
void nested_vector_shape(const Nested& values, std::vector<int>& shape) {
if constexpr (nested_vector_traits<Nested>::depth > 0) {
assert(!values.empty());
assert(values.size() <= std::size_t(std::numeric_limits<int>::max()));
shape.push_back(int(values.size()));
nested_vector_shape(values.front(), shape);
}
}
template <class Nested, class Mint>
void flatten_nested_vector(
const Nested& values,
const std::vector<int>& shape,
int level,
std::vector<Mint>& flattened
) {
if constexpr (nested_vector_traits<Nested>::depth == 0) {
flattened.push_back(values);
} else {
assert(level < int(shape.size()));
assert(int(values.size()) == shape[level]);
for (const auto& child : values) {
flatten_nested_vector(child, shape, level + 1, flattened);
}
}
}
template <class Nested, class Mint>
void rebuild_nested_vector(
Nested& values,
const std::vector<int>& shape,
int level,
const std::vector<Mint>& flattened,
int& position
) {
if constexpr (nested_vector_traits<Nested>::depth == 0) {
assert(position < int(flattened.size()));
values = flattened[position++];
} else {
assert(level < int(shape.size()));
values.resize(shape[level]);
for (auto& child : values) {
rebuild_nested_vector(child, shape, level + 1, flattened, position);
}
}
}
template <class Nested>
std::vector<int> flatten_multivariate_inputs(
const Nested& first,
const Nested& second,
std::vector<typename nested_vector_traits<Nested>::scalar_type>& flattened_first,
std::vector<typename nested_vector_traits<Nested>::scalar_type>& flattened_second
) {
std::vector<int> shape;
nested_vector_shape(first, shape);
assert(int(shape.size()) == nested_vector_traits<Nested>::depth);
std::vector<int> second_shape;
nested_vector_shape(second, second_shape);
assert(second_shape == shape);
flatten_nested_vector(first, shape, 0, flattened_first);
flatten_nested_vector(second, shape, 0, flattened_second);
std::reverse(shape.begin(), shape.end());
return shape;
}
template <class Nested>
Nested rebuild_multivariate_result(
std::vector<int> dimensions,
const std::vector<typename nested_vector_traits<Nested>::scalar_type>& flattened
) {
std::reverse(dimensions.begin(), dimensions.end());
Nested result;
int position = 0;
rebuild_nested_vector(result, dimensions, 0, flattened, position);
assert(position == int(flattened.size()));
return result;
}
inline int multivariate_coefficient_count(const std::vector<int>& dimensions) {
int64_t count = 1;
for (int dimension : dimensions) {
assert(dimension > 0);
count *= dimension;
assert(count <= std::numeric_limits<int>::max());
}
return int(count);
}
inline std::vector<int> multivariate_colors(const std::vector<int>& dimensions) {
const int variable_count = int(dimensions.size());
const int coefficient_count = multivariate_coefficient_count(dimensions);
std::vector<int> color(coefficient_count);
if (variable_count == 0) return color;
for (int index = 0; index < coefficient_count; index++) {
int sum = 0;
int stride = 1;
for (int variable = 0; variable + 1 < variable_count; variable++) {
stride *= dimensions[variable];
sum += index / stride;
}
color[index] = sum % variable_count;
}
return color;
}
template <class Mint>
std::vector<Mint> geometric_evaluation(
const std::vector<Mint>& polynomial, Mint ratio
) {
const int size = int(polynomial.size());
if (size <= 64) {
std::vector<Mint> result(size);
Mint point = 1;
for (int i = 0; i < size; i++) {
Mint power = 1;
for (const Mint& coefficient : polynomial) {
result[i] += coefficient * power;
power *= point;
}
point *= ratio;
}
return result;
}
auto triangular_powers = [](Mint base, int length) {
std::vector<Mint> result(length);
if (length == 0) return result;
result[0] = 1;
Mint power = 1;
for (int i = 0; i + 1 < length; i++) {
result[i + 1] = result[i] * power;
power *= base;
}
return result;
};
std::vector<Mint> positive = triangular_powers(ratio, 2 * size - 1);
std::vector<Mint> negative = triangular_powers(ratio.inv(), size);
std::vector<Mint> scaled(polynomial);
for (int i = 0; i < size; i++) scaled[i] *= negative[i];
std::reverse(scaled.begin(), scaled.end());
std::vector<Mint> product = fps::convolution(scaled, positive);
std::vector<Mint> result(size);
for (int i = 0; i < size; i++) result[i] = product[size - 1 + i] * negative[i];
return result;
}
template <class Mint>
std::vector<Mint> cyclic_fourier_transform(
std::vector<Mint> values, Mint ratio, bool inverse
) {
if constexpr (fps::internal::has_static_modulus<Mint>::value) {
const int size = int(values.size());
if ((size & (size - 1)) == 0) {
// Keep normalization outside the per-axis transforms, matching
// the arbitrary-length DFT path below.
fps::internal::ntt(values, inverse, false);
return values;
}
}
return geometric_evaluation(values, ratio);
}
} // namespace internal
template <class Mint>
std::vector<Mint> multivariate_convolution_truncated(
const std::vector<int>& dimensions,
const std::vector<Mint>& first,
const std::vector<Mint>& second
) {
static_assert(
fps::internal::has_static_modulus<Mint>::value,
"truncated multivariate convolution requires a static-modulus type"
);
const int variable_count = int(dimensions.size());
const int coefficient_count = internal::multivariate_coefficient_count(dimensions);
assert(int(first.size()) == coefficient_count);
assert(int(second.size()) == coefficient_count);
if (variable_count == 0) return {first[0] * second[0]};
int64_t transform_size_64 = 1;
while (transform_size_64 < 2LL * coefficient_count - 1) transform_size_64 <<= 1;
assert(transform_size_64 <= std::numeric_limits<int>::max());
const int transform_size = int(transform_size_64);
assert((Mint::mod() - 1) % uint32_t(transform_size) == 0);
const std::vector<int> color = internal::multivariate_colors(dimensions);
std::vector<std::vector<Mint>> transformed_first(
variable_count, std::vector<Mint>(transform_size)
);
std::vector<std::vector<Mint>> transformed_second(
variable_count, std::vector<Mint>(transform_size)
);
for (int i = 0; i < coefficient_count; i++) {
transformed_first[color[i]][i] = first[i];
transformed_second[color[i]][i] = second[i];
}
for (int group = 0; group < variable_count; group++) {
fps::internal::ntt(transformed_first[group], false);
fps::internal::ntt(transformed_second[group], false);
}
std::vector<std::vector<Mint>> transformed_result(
variable_count, std::vector<Mint>(transform_size)
);
for (int left = 0; left < variable_count; left++) {
for (int right = 0; right < variable_count; right++) {
std::vector<Mint>& destination =
transformed_result[(left + right) % variable_count];
const std::vector<Mint>& left_values = transformed_first[left];
const std::vector<Mint>& right_values = transformed_second[right];
for (int i = 0; i < transform_size; i++) {
destination[i] += left_values[i] * right_values[i];
}
}
}
for (int group = 0; group < variable_count; group++) {
fps::internal::ntt(transformed_result[group], true);
}
std::vector<Mint> result(coefficient_count);
for (int i = 0; i < coefficient_count; i++) {
result[i] = transformed_result[color[i]][i];
}
return result;
}
template <
class Nested,
std::enable_if_t<(internal::nested_vector_traits<Nested>::depth > 0), int> = 0
>
Nested multivariate_convolution_truncated(
const Nested& first,
const Nested& second
) {
using Mint = typename internal::nested_vector_traits<Nested>::scalar_type;
std::vector<Mint> flattened_first, flattened_second;
std::vector<int> dimensions = internal::flatten_multivariate_inputs(
first, second, flattened_first, flattened_second
);
std::vector<Mint> flattened_result = multivariate_convolution_truncated(
dimensions, flattened_first, flattened_second
);
return internal::rebuild_multivariate_result<Nested>(
std::move(dimensions), flattened_result
);
}
template <class Mint>
std::vector<Mint> multivariate_convolution_cyclic(
const std::vector<int>& dimensions,
const std::vector<Mint>& first,
const std::vector<Mint>& second
) {
const int coefficient_count = internal::multivariate_coefficient_count(dimensions);
assert(int(first.size()) == coefficient_count);
assert(int(second.size()) == coefficient_count);
if (dimensions.empty()) return {first[0] * second[0]};
const uint32_t modulus = Mint::mod();
bool has_all_roots = true;
for (int dimension : dimensions) {
if ((modulus - 1) % uint32_t(dimension) != 0) has_all_roots = false;
}
if (!has_all_roots) {
std::vector<int> reduced_dimensions;
for (int dimension : dimensions) {
if (dimension != 1) reduced_dimensions.push_back(dimension);
}
if (reduced_dimensions.empty()) return {first[0] * second[0]};
std::vector<int> widened_dimensions(reduced_dimensions.size());
for (int i = 0; i < int(reduced_dimensions.size()); i++) {
const int64_t widened = 2LL * reduced_dimensions[i] - 1;
assert(widened <= std::numeric_limits<int>::max());
widened_dimensions[i] = int(widened);
}
const int widened_count =
internal::multivariate_coefficient_count(widened_dimensions);
// The largest embedded input index uses coordinate dimension - 1 on
// every axis. Its double is widened_count - 1, so convolving arrays
// ending at this index produces exactly the widened mixed-radix box.
// In particular, fps::convolution chooses the smallest transform that
// contains widened_count coefficients, instead of one that contains
// 2 * widened_count - 1 coefficients due to trailing zeroes.
int64_t maximum_embedded_index = 0;
int64_t widened_stride = 1;
for (int variable = 0; variable < int(reduced_dimensions.size()); variable++) {
maximum_embedded_index +=
int64_t(reduced_dimensions[variable] - 1) * widened_stride;
widened_stride *= widened_dimensions[variable];
}
assert(widened_stride == widened_count);
assert(2 * maximum_embedded_index + 1 == widened_count);
assert(maximum_embedded_index < std::numeric_limits<int>::max());
const int embedded_input_count = int(maximum_embedded_index) + 1;
std::vector<Mint> widened_first(embedded_input_count);
std::vector<Mint> widened_second(embedded_input_count);
for (int index = 0; index < coefficient_count; index++) {
int remaining = index;
int widened_index = 0;
int embedding_stride = 1;
for (int variable = 0; variable < int(reduced_dimensions.size()); variable++) {
const int coordinate = remaining % reduced_dimensions[variable];
remaining /= reduced_dimensions[variable];
widened_index += coordinate * embedding_stride;
embedding_stride *= widened_dimensions[variable];
}
widened_first[widened_index] = first[index];
widened_second[widened_index] = second[index];
}
std::vector<Mint> widened_product =
fps::convolution(widened_first, widened_second);
assert(int(widened_product.size()) == widened_count);
std::vector<Mint> result(coefficient_count);
for (int widened_index = 0; widened_index < widened_count; widened_index++) {
int remaining = widened_index;
int index = 0;
int stride = 1;
for (int variable = 0; variable < int(reduced_dimensions.size()); variable++) {
const int coordinate = remaining % widened_dimensions[variable];
remaining /= widened_dimensions[variable];
index += (coordinate % reduced_dimensions[variable]) * stride;
stride *= reduced_dimensions[variable];
}
result[index] += widened_product[widened_index];
}
return result;
}
const uint64_t generator = primitive_root(modulus);
assert(generator != 0);
std::vector<Mint> transformed_first(first);
std::vector<Mint> transformed_second(second);
int stride = 1;
for (int dimension : dimensions) {
assert((modulus - 1) % uint32_t(dimension) == 0);
const Mint root = Mint(generator).pow((modulus - 1) / dimension);
for (int block = 0; block < coefficient_count; block += stride * dimension) {
for (int offset = 0; offset < stride; offset++) {
std::vector<Mint> first_line(dimension);
std::vector<Mint> second_line(dimension);
for (int i = 0; i < dimension; i++) {
first_line[i] = transformed_first[block + offset + stride * i];
second_line[i] = transformed_second[block + offset + stride * i];
}
first_line = internal::cyclic_fourier_transform(
std::move(first_line), root, false
);
second_line = internal::cyclic_fourier_transform(
std::move(second_line), root, false
);
for (int i = 0; i < dimension; i++) {
transformed_first[block + offset + stride * i] = first_line[i];
transformed_second[block + offset + stride * i] = second_line[i];
}
}
}
stride *= dimension;
}
for (int i = 0; i < coefficient_count; i++) {
transformed_first[i] *= transformed_second[i];
}
stride = 1;
for (int dimension : dimensions) {
const Mint inverse_root =
Mint(generator).pow((modulus - 1) / dimension).inv();
for (int block = 0; block < coefficient_count; block += stride * dimension) {
for (int offset = 0; offset < stride; offset++) {
std::vector<Mint> line(dimension);
for (int i = 0; i < dimension; i++) {
line[i] = transformed_first[block + offset + stride * i];
}
line = internal::cyclic_fourier_transform(
std::move(line), inverse_root, true
);
for (int i = 0; i < dimension; i++) {
transformed_first[block + offset + stride * i] = line[i];
}
}
}
stride *= dimension;
}
const Mint inverse_size = Mint(coefficient_count).inv();
for (Mint& value : transformed_first) value *= inverse_size;
return transformed_first;
}
template <
class Nested,
std::enable_if_t<(internal::nested_vector_traits<Nested>::depth > 0), int> = 0
>
Nested multivariate_convolution_cyclic(
const Nested& first,
const Nested& second
) {
using Mint = typename internal::nested_vector_traits<Nested>::scalar_type;
std::vector<Mint> flattened_first, flattened_second;
std::vector<int> dimensions = internal::flatten_multivariate_inputs(
first, second, flattened_first, flattened_second
);
std::vector<Mint> flattened_result = multivariate_convolution_cyclic(
dimensions, flattened_first, flattened_second
);
return internal::rebuild_multivariate_result<Nested>(
std::move(dimensions), flattened_result
);
}
} // namespace math
} // namespace m1une