Tree Distance Frequency
(graph/tree/distance_frequency.hpp)
- View this file on GitHub
- Last update: 2026-08-13 01:41:40+09:00
- Include:
#include "graph/tree/distance_frequency.hpp"
Overview
tree_distance_frequency counts unordered pairs of vertices at every distance
in an unweighted tree. It is useful when a problem asks for the complete
distribution of path lengths rather than individual distance queries.
#include "graph/tree/distance_frequency.hpp"
The function is in m1une::tree and accepts m1une::graph::Graph<T>.
Interface
template <class T>
std::vector<long long> tree_distance_frequency(
const m1une::graph::Graph<T>& tree
);
| Function | Description | Complexity |
|---|---|---|
tree_distance_frequency(tree) |
Returns the number of unordered vertex pairs at each edge distance. | $O(N \log^2 N)$ time and $O(N)$ additional memory |
For a tree with N vertices, the result has length N:
-
result[0]isN, counting each vertex paired with itself; - for
distance > 0,result[distance]is the number of pairs(u, v)withu < vwhose path contains exactlydistanceedges.
The empty tree produces an empty vector. Otherwise, the graph must be a
connected undirected tree with N - 1 active edges. Edge costs are ignored.
The result is exact rather than reduced modulo a number.
Algorithm
At each centroid, depth histograms are collected for the whole current component and for every component obtained by removing the centroid. Squaring the whole histogram and subtracting the component squares counts ordered pairs whose path passes through that centroid. Centroid decomposition makes every pair appear at exactly one level.
The histogram squares use NTT convolution under two prime moduli. Chinese remaindering reconstructs the exact ordered-pair counts, which are then divided by two for positive distances.
Example
#include "graph/graph.hpp"
#include "graph/tree/distance_frequency.hpp"
#include <iostream>
#include <vector>
int main() {
m1une::graph::Graph<int> tree(4);
tree.add_edge(0, 1);
tree.add_edge(1, 2);
tree.add_edge(2, 3);
std::vector<long long> frequency =
m1une::tree::tree_distance_frequency(tree);
// frequency is {4, 3, 2, 1}.
std::cout << frequency[2] << '\n';
}
Depends on
Graph
(graph/graph.hpp)
Centroid Decomposition
(graph/tree/centroid_decomposition.hpp)
Convolution
(math/fps/convolution.hpp)
math/fps/internal/ntt998_faster.hpp
ModInt
(math/modint.hpp)
ModInt
(math/modint.hpp)
Required by
Verified with
verify/graph/cow_game.test.cpp
verify/graph/graph_algorithms.test.cpp
verify/graph/range_edge_graph.test.cpp
verify/graph/tree/distance_frequency.test.cpp
verify/graph/tree/tree_algorithms.test.cpp
Code
#ifndef M1UNE_TREE_DISTANCE_FREQUENCY_HPP
#define M1UNE_TREE_DISTANCE_FREQUENCY_HPP 1
#include <algorithm>
#include <cassert>
#include <cstddef>
#include <cstdint>
#include <utility>
#include <vector>
#include "../../math/fps/convolution.hpp"
#include "../../math/modint.hpp"
#include "centroid_decomposition.hpp"
namespace m1une {
namespace tree {
namespace distance_frequency_detail {
template <class Mint, class T>
std::vector<Mint> count_ordered_pairs(
const m1une::graph::Graph<T>& tree,
const CentroidDecomposition<T>& decomposition
) {
const int size = tree.size();
std::vector<Mint> count(static_cast<std::size_t>(size));
std::vector<char> removed(std::size_t(size), false);
std::vector<Mint> histogram;
std::vector<std::pair<int, int>> stack;
std::vector<int> parent(std::size_t(size), -1);
for (int centroid : decomposition.order) {
std::vector<Mint> total(1, Mint(1));
for (const auto& edge : tree[centroid]) {
if (!edge.alive || removed[std::size_t(edge.to)]) continue;
histogram.clear();
stack.clear();
stack.emplace_back(edge.to, 1);
parent[std::size_t(edge.to)] = centroid;
while (!stack.empty()) {
const auto [vertex, distance] = stack.back();
stack.pop_back();
if (int(histogram.size()) <= distance) {
histogram.resize(std::size_t(distance + 1));
}
histogram[std::size_t(distance)] += Mint(1);
for (const auto& next : tree[vertex]) {
if (!next.alive || removed[std::size_t(next.to)]) continue;
if (next.to == parent[std::size_t(vertex)]) continue;
parent[std::size_t(next.to)] = vertex;
stack.emplace_back(next.to, distance + 1);
}
}
if (total.size() < histogram.size()) {
total.resize(histogram.size());
}
for (std::size_t distance = 0; distance < histogram.size(); distance++) {
total[distance] += histogram[distance];
}
const std::vector<Mint> within_component =
m1une::fps::convolution(histogram, histogram);
const std::size_t limit = std::min(count.size(), within_component.size());
for (std::size_t distance = 0; distance < limit; distance++) {
count[distance] -= within_component[distance];
}
}
const std::vector<Mint> through_centroid =
m1une::fps::convolution(total, total);
const std::size_t limit = std::min(count.size(), through_centroid.size());
for (std::size_t distance = 0; distance < limit; distance++) {
count[distance] += through_centroid[distance];
}
removed[std::size_t(centroid)] = true;
}
return count;
}
inline std::uint64_t combine_residues(std::uint32_t first, std::uint32_t second) {
using First = m1une::math::ModInt<998244353>;
using Second = m1une::math::ModInt<924844033>;
static const std::uint64_t inverse = Second(First::mod()).inv().val();
const std::uint64_t offset =
(std::uint64_t(second) + Second::mod() - first % Second::mod()) %
Second::mod();
const std::uint64_t multiplier = offset * inverse % Second::mod();
return std::uint64_t(first) + std::uint64_t(First::mod()) * multiplier;
}
} // namespace distance_frequency_detail
template <class T>
std::vector<long long> tree_distance_frequency(
const m1une::graph::Graph<T>& tree
) {
const int size = tree.size();
assert(tree.edge_count() == std::max(0, size - 1));
if (size == 0) return {};
const CentroidDecomposition<T> decomposition(tree);
assert(decomposition.roots.size() == 1);
using First = m1une::math::ModInt<998244353>;
using Second = m1une::math::ModInt<924844033>;
assert(
std::uint64_t(size) * std::uint64_t(size - 1) <
std::uint64_t(First::mod()) * Second::mod()
);
const std::vector<First> first =
distance_frequency_detail::count_ordered_pairs<First>(
tree,
decomposition
);
const std::vector<Second> second =
distance_frequency_detail::count_ordered_pairs<Second>(
tree,
decomposition
);
std::vector<long long> result(static_cast<std::size_t>(size));
result[0] = size;
for (int distance = 1; distance < size; distance++) {
const std::uint64_t ordered =
distance_frequency_detail::combine_residues(
first[std::size_t(distance)].val(),
second[std::size_t(distance)].val()
);
assert((ordered & 1) == 0);
result[std::size_t(distance)] = static_cast<long long>(ordered / 2);
}
return result;
}
} // namespace tree
} // namespace m1une
#endif // M1UNE_TREE_DISTANCE_FREQUENCY_HPP#line 1 "graph/tree/distance_frequency.hpp"
#include <algorithm>
#include <cassert>
#include <cstddef>
#include <cstdint>
#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>
#include <type_traits>
#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 "graph/tree/centroid_decomposition.hpp"
#line 6 "graph/tree/centroid_decomposition.hpp"
#line 1 "graph/graph.hpp"
#line 8 "graph/graph.hpp"
namespace m1une {
namespace graph {
template <class T = int>
struct Edge {
using cost_type = T;
int from;
int to;
T cost;
int id;
bool alive;
Edge() : from(-1), to(-1), cost(T()), id(-1), alive(true) {}
Edge(int from_, int to_, T cost_ = T(1), int id_ = -1, bool alive_ = true)
: from(from_), to(to_), cost(cost_), id(id_), alive(alive_) {}
int other(int v) const {
assert(v == from || v == to);
return from ^ to ^ v;
}
};
template <class T = int>
struct Graph {
using edge_type = Edge<T>;
using cost_type = T;
private:
struct EdgePositions {
std::array<std::pair<int, int>, 2> value{};
int size = 0;
void push_back(std::pair<int, int> position) {
assert(size < 2);
value[size++] = position;
}
};
int _n;
int _edge_count;
std::vector<std::vector<edge_type>> _g;
std::vector<EdgePositions> _edge_positions;
public:
Graph() : _n(0), _edge_count(0) {}
explicit Graph(int n) : _n(n), _edge_count(0), _g(n) {
assert(0 <= n);
}
int size() const {
return _n;
}
bool empty() const {
return _n == 0;
}
int edge_count() const {
return _edge_count;
}
int add_vertex() {
_g.emplace_back();
return _n++;
}
int add_directed_edge(int from, int to, T cost = T(1)) {
assert(0 <= from && from < _n);
assert(0 <= to && to < _n);
int id = _edge_count++;
int idx = int(_g[from].size());
_g[from].push_back(edge_type(from, to, cost, id));
_edge_positions.emplace_back();
_edge_positions.back().push_back({from, idx});
return id;
}
int add_edge(int u, int v, T cost = T(1)) {
assert(0 <= u && u < _n);
assert(0 <= v && v < _n);
int id = _edge_count++;
int u_idx = int(_g[u].size());
_g[u].push_back(edge_type(u, v, cost, id));
int v_idx = int(_g[v].size());
_g[v].push_back(edge_type(v, u, cost, id));
_edge_positions.emplace_back();
_edge_positions.back().push_back({u, u_idx});
_edge_positions.back().push_back({v, v_idx});
return id;
}
void set_edge_alive(int id, bool alive) {
assert(0 <= id && id < _edge_count);
for (int i = 0; i < _edge_positions[id].size; ++i) {
auto [v, idx] = _edge_positions[id].value[i];
_g[v][idx].alive = alive;
}
}
void erase_edge(int id) {
set_edge_alive(id, false);
}
void revive_edge(int id) {
set_edge_alive(id, true);
}
bool is_edge_alive(int id) const {
assert(0 <= id && id < _edge_count);
assert(_edge_positions[id].size != 0);
auto [v, idx] = _edge_positions[id].value[0];
return _g[v][idx].alive;
}
const std::vector<edge_type>& operator[](int v) const {
assert(0 <= v && v < _n);
return _g[v];
}
std::vector<edge_type>& operator[](int v) {
assert(0 <= v && v < _n);
return _g[v];
}
const std::vector<std::vector<edge_type>>& adjacency() const {
return _g;
}
std::vector<std::vector<edge_type>>& adjacency() {
return _g;
}
std::vector<edge_type> edges(bool include_inactive = false) const {
std::vector<edge_type> result;
result.reserve(_edge_count);
std::vector<char> used(_edge_count, false);
for (int v = 0; v < _n; v++) {
for (const auto& e : _g[v]) {
if (!include_inactive && !e.alive) continue;
if (0 <= e.id && e.id < _edge_count) {
if (used[e.id]) continue;
used[e.id] = true;
}
result.push_back(e);
}
}
return result;
}
Graph reversed() const {
Graph result(_n);
result._edge_count = _edge_count;
result._edge_positions.assign(_edge_count, {});
for (int v = 0; v < _n; v++) {
for (const auto& e : _g[v]) {
int idx = int(result._g[e.to].size());
result._g[e.to].push_back(edge_type(e.to, e.from, e.cost, e.id, e.alive));
if (0 <= e.id && e.id < _edge_count) result._edge_positions[e.id].push_back({e.to, idx});
}
}
return result;
}
};
} // namespace graph
} // namespace m1une
#line 8 "graph/tree/centroid_decomposition.hpp"
namespace m1une {
namespace tree {
template <class T = int>
struct CentroidDecomposition {
int n;
std::vector<int> parent;
std::vector<int> depth;
std::vector<int> order;
std::vector<int> roots;
std::vector<std::vector<int>> children;
private:
std::vector<int> _subtree_size;
std::vector<int> _work_parent;
std::vector<char> _removed;
void build_component(const m1une::graph::Graph<T>& g, int start, int p, int d) {
std::vector<int> nodes;
std::vector<int> stack = {start};
_work_parent[start] = -2;
while (!stack.empty()) {
int v = stack.back();
stack.pop_back();
nodes.push_back(v);
for (const auto& e : g[v]) {
if (!e.alive || _removed[e.to]) continue;
if (_work_parent[e.to] != -1) continue;
_work_parent[e.to] = v;
stack.push_back(e.to);
}
}
for (int v : nodes) _subtree_size[v] = 1;
for (int i = int(nodes.size()) - 1; i >= 0; i--) {
int v = nodes[i];
if (_work_parent[v] >= 0) _subtree_size[_work_parent[v]] += _subtree_size[v];
}
int total = int(nodes.size());
int centroid = start;
int best = total + 1;
for (int v : nodes) {
int largest = total - _subtree_size[v];
for (const auto& e : g[v]) {
if (!e.alive || _removed[e.to]) continue;
if (_work_parent[e.to] == v) largest = std::max(largest, _subtree_size[e.to]);
}
if (largest < best) {
best = largest;
centroid = v;
}
}
for (int v : nodes) _work_parent[v] = -1;
parent[centroid] = p;
depth[centroid] = d;
order.push_back(centroid);
if (p == -1) {
roots.push_back(centroid);
} else {
children[p].push_back(centroid);
}
_removed[centroid] = true;
for (const auto& e : g[centroid]) {
if (!e.alive || _removed[e.to]) continue;
build_component(g, e.to, centroid, d + 1);
}
}
public:
CentroidDecomposition() : n(0) {}
explicit CentroidDecomposition(const m1une::graph::Graph<T>& g) {
build(g);
}
void build(const m1une::graph::Graph<T>& g) {
n = g.size();
parent.assign(n, -1);
depth.assign(n, -1);
order.clear();
order.reserve(n);
roots.clear();
children.assign(n, {});
_subtree_size.assign(n, 0);
_work_parent.assign(n, -1);
_removed.assign(n, false);
for (int v = 0; v < n; v++) {
if (depth[v] == -1) build_component(g, v, -1, 0);
}
}
int size() const {
return n;
}
bool empty() const {
return n == 0;
}
int root() const {
return roots.empty() ? -1 : roots[0];
}
};
} // namespace tree
} // namespace m1une
#line 14 "graph/tree/distance_frequency.hpp"
namespace m1une {
namespace tree {
namespace distance_frequency_detail {
template <class Mint, class T>
std::vector<Mint> count_ordered_pairs(
const m1une::graph::Graph<T>& tree,
const CentroidDecomposition<T>& decomposition
) {
const int size = tree.size();
std::vector<Mint> count(static_cast<std::size_t>(size));
std::vector<char> removed(std::size_t(size), false);
std::vector<Mint> histogram;
std::vector<std::pair<int, int>> stack;
std::vector<int> parent(std::size_t(size), -1);
for (int centroid : decomposition.order) {
std::vector<Mint> total(1, Mint(1));
for (const auto& edge : tree[centroid]) {
if (!edge.alive || removed[std::size_t(edge.to)]) continue;
histogram.clear();
stack.clear();
stack.emplace_back(edge.to, 1);
parent[std::size_t(edge.to)] = centroid;
while (!stack.empty()) {
const auto [vertex, distance] = stack.back();
stack.pop_back();
if (int(histogram.size()) <= distance) {
histogram.resize(std::size_t(distance + 1));
}
histogram[std::size_t(distance)] += Mint(1);
for (const auto& next : tree[vertex]) {
if (!next.alive || removed[std::size_t(next.to)]) continue;
if (next.to == parent[std::size_t(vertex)]) continue;
parent[std::size_t(next.to)] = vertex;
stack.emplace_back(next.to, distance + 1);
}
}
if (total.size() < histogram.size()) {
total.resize(histogram.size());
}
for (std::size_t distance = 0; distance < histogram.size(); distance++) {
total[distance] += histogram[distance];
}
const std::vector<Mint> within_component =
m1une::fps::convolution(histogram, histogram);
const std::size_t limit = std::min(count.size(), within_component.size());
for (std::size_t distance = 0; distance < limit; distance++) {
count[distance] -= within_component[distance];
}
}
const std::vector<Mint> through_centroid =
m1une::fps::convolution(total, total);
const std::size_t limit = std::min(count.size(), through_centroid.size());
for (std::size_t distance = 0; distance < limit; distance++) {
count[distance] += through_centroid[distance];
}
removed[std::size_t(centroid)] = true;
}
return count;
}
inline std::uint64_t combine_residues(std::uint32_t first, std::uint32_t second) {
using First = m1une::math::ModInt<998244353>;
using Second = m1une::math::ModInt<924844033>;
static const std::uint64_t inverse = Second(First::mod()).inv().val();
const std::uint64_t offset =
(std::uint64_t(second) + Second::mod() - first % Second::mod()) %
Second::mod();
const std::uint64_t multiplier = offset * inverse % Second::mod();
return std::uint64_t(first) + std::uint64_t(First::mod()) * multiplier;
}
} // namespace distance_frequency_detail
template <class T>
std::vector<long long> tree_distance_frequency(
const m1une::graph::Graph<T>& tree
) {
const int size = tree.size();
assert(tree.edge_count() == std::max(0, size - 1));
if (size == 0) return {};
const CentroidDecomposition<T> decomposition(tree);
assert(decomposition.roots.size() == 1);
using First = m1une::math::ModInt<998244353>;
using Second = m1une::math::ModInt<924844033>;
assert(
std::uint64_t(size) * std::uint64_t(size - 1) <
std::uint64_t(First::mod()) * Second::mod()
);
const std::vector<First> first =
distance_frequency_detail::count_ordered_pairs<First>(
tree,
decomposition
);
const std::vector<Second> second =
distance_frequency_detail::count_ordered_pairs<Second>(
tree,
decomposition
);
std::vector<long long> result(static_cast<std::size_t>(size));
result[0] = size;
for (int distance = 1; distance < size; distance++) {
const std::uint64_t ordered =
distance_frequency_detail::combine_residues(
first[std::size_t(distance)].val(),
second[std::size_t(distance)].val()
);
assert((ordered & 1) == 0);
result[std::size_t(distance)] = static_cast<long long>(ordered / 2);
}
return result;
}
} // namespace tree
} // namespace m1une