cp-library

This documentation is automatically generated by online-judge-tools/verification-helper

View the Project on GitHub kobejean/cp-library

:heavy_check_mark: test/library-checker/set-power-series/power_projection_of_set_power_series.test.py

Depends on

Code

# verification-helper: PROBLEM https://judge.yosupo.jp/problem/power_projection_of_set_power_series

def main():
    N, M = rd()
    A = rdl(1<<N)
    W = rdl(1<<N)
    ans = sps_pow_proj_poly(W, A, M, 998244353)
    wtnl(ans)

from cp_library.math.sps.mod.sps_pow_proj_poly_fn import sps_pow_proj_poly
from cp_library.io.fast_io_fn import rd, rdl, wtnl

main()
# verification-helper: PROBLEM https://judge.yosupo.jp/problem/power_projection_of_set_power_series

def main():
    N, M = rd()
    A = rdl(1<<N)
    W = rdl(1<<N)
    ans = sps_pow_proj_poly(W, A, M, 998244353)
    wtnl(ans)

'''
╺━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╸
             https://kobejean.github.io/cp-library               
'''




from typing import Generic
from typing import TypeVar

_S = TypeVar('S'); _T = TypeVar('T'); _U = TypeVar('U'); _T1 = TypeVar('T1'); _T2 = TypeVar('T2'); _T3 = TypeVar('T3'); _T4 = TypeVar('T4'); _T5 = TypeVar('T5'); _T6 = TypeVar('T6')

import sys


def list_find(lst: list, value, start = 0, stop = sys.maxsize):
    try:
        return lst.index(value, start, stop)
    except:
        return -1


class view(Generic[_T]):
    __slots__ = 'A', 'l', 'r'
    def __init__(V, A: list[_T], l: int = 0, r: int = 0): V.A, V.l, V.r = A, l, r
    def __len__(V): return V.r - V.l
    def __getitem__(V, i: int): 
        if 0 <= i < V.r - V.l: return V.A[V.l+i]
        else: raise IndexError
    def __setitem__(V, i: int, v: _T): V.A[V.l+i] = v
    def __contains__(V, v: _T): return list_find(V.A, v, V.l, V.r) != -1
    def set_range(V, l: int, r: int): V.l, V.r = l, r
    def index(V, v: _T): return V.A.index(v, V.l, V.r) - V.l
    def reverse(V):
        l, r = V.l, V.r-1
        while l < r: V.A[l], V.A[r] = V.A[r], V.A[l]; l += 1; r -= 1
    def sort(V, /, *args, **kwargs):
        A = V.A[V.l:V.r]; A.sort(*args, **kwargs)
        for i,a in enumerate(A,V.l): V.A[i] = a
    def pop(V): V.r -= 1; return V.A[V.r]
    def append(V, v: _T): V.A[V.r] = v; V.r += 1
    def popleft(V): V.l += 1; return V.A[V.l-1]
    def appendleft(V, v: _T): V.l -= 1; V.A[V.l] = v; 
    def validate(V): return 0 <= V.l <= V.r <= len(V.A)


def popcnts(N):
    P = [0]*(1 << N)
    for i in range(N):
        for m in range(b := 1<<i):
            P[m^b] = P[m] + 1
    return P
'''
╺━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╸
    x₀ ────────●─●────────●───●────────●───────●────────► X₀
                ╳          ╲ ╱          ╲     ╱          
    x₄ ────────●─●────────●─╳─●────────●─╲───╱─●────────► X₁
                           ╳ ╳          ╲ ╲ ╱ ╱          
    x₂ ────────●─●────────●─╳─●────────●─╲─╳─╱─●────────► X₂
                ╳          ╱ ╲          ╲ ╳ ╳ ╱          
    x₆ ────────●─●────────●───●────────●─╳─╳─╳─●────────► X₃
                                        ╳ ╳ ╳ ╳         
    x₁ ────────●─●────────●───●────────●─╳─╳─╳─●────────► X₄
                ╳          ╲ ╱          ╱ ╳ ╳ ╲          
    x₅ ────────●─●────────●─╳─●────────●─╱─╳─╲─●────────► X₅
                           ╳ ╳          ╱ ╱ ╲ ╲          
    x₃ ────────●─●────────●─╳─●────────●─╱───╲─●────────► X₆
                ╳          ╱ ╲          ╱     ╲          
    x₇ ────────●─●────────●───●────────●───────●────────► X₇
╺━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╸
                      Math - Convolution                     
'''



def max2(a, b): return a if a > b else b


def ior_zeta_pair_ranked(A, B, N, M, Z):
    for i in range(0, Z, M):
        l, r = i+(1<<(i>>N))-1, i+M
        for j in range(N):
            m = l|(b := 1<<j)
            while m < r: A[m] += A[m^b]; B[m] += B[m^b]; m = m+1|b
    return A, B

def ior_mobius_ranked(A: list[int], N: int, M: int, Z: int):
    for i in range(0, Z, M):
        l, r = i, i+M-(1<<(N-(i>>N)))+1
        for j in range(N):
            m = l|(b := 1<<j)
            while m < r: A[m] -= A[m^b]; m = m+1|b
    return A

def isubset_conv_ranked(Ar, Br, N, M, Z, mod) -> list[int]:
    ior_zeta_pair_ranked(Ar, Br, N, M, Z)
    for i in range(Z): Ar[i], Br[i] = Ar[i]%mod, Br[i]%mod
    for ij in range(Z-M,-1,-M):
        for k in range(M): Ar[ij|k] = (Ar[ij|k] * Br[k]) % mod
        r = M-(1 << (N-(ij>>N)))+1
        for i in range(0,ij,M):
            j = ij-i; l = (1 << (max2(i,j)>>N))-1
            for k in range(l,r): Ar[ij|k] += Ar[i|k] * Br[j|k] % mod
    return ior_mobius_ranked(Ar, N, M, Z)

def subset_conv(A: list[int], B: list[int], N: int, mod: int) -> list[int]:
    Z = (N+1)*(M:=1<<N)
    Ar, Br, C, P = [0]*Z, [0]*Z, [0]*M, popcnts(N)
    for i, p in enumerate(P): Ar[p<<N|i], Br[p<<N|i] = A[i], B[i]
    isubset_conv_ranked(Ar, Br, N, M, Z, mod)
    for i, p in enumerate(P): C[i] = Ar[p<<N|i] % mod
    return C

def sps_pow_proj(A, B, M, mod):
    N = len(B).bit_length() - 1
    assert B[0] == 0, "B[0] must be 0 for sps_pow_proj"
    Aview, Bview, P = view(A := A[::-1]), view(B), []
    for i in range(N + 1):
        P.append(A[(1<<N)-1]); A[(1<<N)-1] = 0
        for m in range(N - i):
            i0 = (1<<N)-(1<<(m+1)); i1 = 1 << m; i2 = (1<<N)-(1<<m)
            Aview.set_range(i0, i0+(1<<m)); Bview.set_range(i1, i1+(1<<m))
            R = subset_conv(Aview, Bview, m, mod)
            for h in range(1 << m): A[i2+h], A[i0+h] = (A[i2+h]+R[h])%mod, 0
    return P[:M]

def sps_pow_proj_poly(A, B, M, mod):
    N = len(B).bit_length()-1; b = B[0]; B[0] = 0
    P, F, H = sps_pow_proj(A, B, N + 1, mod), [], [0]*(N+1); H[0] = 1
    for i in range(M):
        v = 0
        for j in range(N + 1):
            if j < len(P): v = (v+H[j]*P[j])%mod
        F.append(v)
        for j in range(N-1, -1, -1): H[j+1]=H[j]*(i+1)%mod
        H[0]=H[0]*b%mod
    return F

from os import read as os_read, write as os_write, fstat as os_fstat
from __pypy__.builders import StringBuilder

class IOBase:
    @property
    def char(io) -> bool: ...
    @property
    def writable(io) -> bool: ...
    def __next__(io) -> str: ...
    def write(io, s: str) -> None: ...
    def readline(io) -> str: ...
    def readtoken(io) -> str: ...
    def readtokens(io) -> list[str]: ...
    def readints(io) -> list[int]: ...
    def readdigits(io) -> list[int]: ...
    def readnums(io) -> list[int]: ...
    def readchar(io) -> str: ...
    def readchars(io) -> str: ...
    def readinto(io, lst: list[str]) -> list[str]: ...
    def readcharsinto(io, lst: list[str]) -> list[str]: ...
    def readtokensinto(io, lst: list[str]) -> list[str]: ...
    def readintsinto(io, lst: list[int]) -> list[int]: ...
    def readdigitsinto(io, lst: list[int]) -> list[int]: ...
    def readnumsinto(io, lst: list[int]) -> list[int]: ...
    def wait(io): ...
    def flush(io) -> None: ...
    def line(io) -> list[str]: ...

class IO(IOBase):
    BUFSIZE = 1 << 16; stdin: 'IO'; stdout: 'IO'
    __slots__ = 'f', 'file', 'B', 'O', 'V', 'S', 'l', 'p', 'char', 'sz', 'st', 'ist', 'writable', 'encoding', 'errors'
    def __init__(io, file):
        io.file = file
        try: io.f = file.fileno(); io.sz, io.writable = max2(io.BUFSIZE, os_fstat(io.f).st_size), ('x' in file.mode or 'r' not in file.mode)
        except: io.f, io.sz, io.writable = -1, io.BUFSIZE, False
        io.B, io.O, io.S = bytearray(), [], StringBuilder(); io.V = memoryview(io.B); io.l = io.p = 0
        io.char, io.st, io.ist, io.encoding, io.errors = False, [], [], 'ascii', 'ignore'
    def _dec(io, l, r): return io.V[l:r].tobytes().decode(io.encoding, io.errors)
    def readbytes(io, sz): return os_read(io.f, sz)
    def load(io):
        while io.l >= len(io.O):
            if not (b := io.readbytes(io.sz)):
                if io.O[-1] < len(io.B): io.O.append(len(io.B))
                break
            pos = len(io.B); io.B.extend(b)
            while ~(pos := io.B.find(b'\n', pos)): io.O.append(pos := pos+1)
    def __next__(io):
        if io.char: return io.readchar()
        else: return io.readtoken()
    def readchar(io):
        io.load(); r = io.O[io.l]
        c = chr(io.B[io.p])
        if io.p >= r-1: io.p = r; io.l += 1
        else: io.p += 1
        return c
    def write(io, s: str): io.S.append(s)
    def readline(io): io.load(); l, io.p = io.p, io.O[io.l]; io.l += 1; return io._dec(l, io.p)
    def readtoken(io):
        io.load(); r = io.O[io.l]
        if ~(p := io.B.find(b' ', io.p, r)): s = io._dec(io.p, p); io.p = p+1
        else: s = io._dec(io.p, r-1); io.p = r; io.l += 1
        return s
    def readtokens(io): io.st.clear(); return io.readtokensinto(io.st)
    def readints(io): io.ist.clear(); return io.readintsinto(io.ist)
    def readdigits(io): io.ist.clear(); return io.readdigitsinto(io.ist)
    def readnums(io): io.ist.clear(); return io.readnumsinto(io.ist)
    def readchars(io): io.load(); l, io.p = io.p, io.O[io.l]; io.l += 1; return io._dec(l, io.p-1)
    def readinto(io, lst):
        if io.char: return io.readcharsinto(lst)
        else: return io.readtokensinto(lst)
    def readcharsinto(io, lst): lst.extend(io.readchars()); return lst
    def readtokensinto(io, lst): 
        io.load(); r = io.O[io.l]
        while ~(p := io.B.find(b' ', io.p, r)): lst.append(io._dec(io.p, p)); io.p = p+1
        lst.append(io._dec(io.p, r-1)); io.p = r; io.l += 1; return lst
    def _readint(io, r):
        while io.p < r and io.B[io.p] <= 32: io.p += 1
        if io.p >= r: return None
        minus = x = 0
        if io.B[io.p] == 45: minus = 1; io.p += 1
        while io.p < r and io.B[io.p] >= 48: x = x * 10 + (io.B[io.p] & 15); io.p += 1
        io.p += 1
        return -x if minus else x
    def readintsinto(io, lst):
        io.load(); r = io.O[io.l]
        while io.p < r and (x := io._readint(r)) is not None: lst.append(x)
        io.l += 1; return lst
    def _readdigit(io): d = io.B[io.p] & 15; io.p += 1; return d
    def readdigitsinto(io, lst):
        io.load(); r = io.O[io.l]
        while io.p < r and io.B[io.p] > 32: lst.append(io._readdigit())
        if io.B[io.p] == 10: io.l += 1
        io.p += 1
        return lst
    def readnumsinto(io, lst):
        if io.char: return io.readdigitsinto(lst)
        else: return io.readintsinto(lst)
    def line(io): io.st.clear(); return io.readinto(io.st)
    def wait(io):
        io.load(); r = io.O[io.l]
        while io.p < r: yield
    def flush(io):
        if io.writable: os_write(io.f, io.S.build().encode(io.encoding, io.errors)); io.S = StringBuilder()
sys.stdin = IO.stdin = IO(sys.stdin); sys.stdout = IO.stdout = IO(sys.stdout)
def rd(): return IO.stdin.readints()
def rds(): return IO.stdin.__next__()
def rdl(n): return IO.stdin.readintsinto(elist(n))
def wt(s): IO.stdout.write(s)
def wtn(s): IO.stdout.write(f'{s}\n')
def wtnl(l): IO.stdout.write(' '.join(map(str, l)))

def elist(hint: int) -> list: ...
try:
    from __pypy__ import newlist_hint
except:
    def newlist_hint(hint): return []
elist = newlist_hint
    

main()
Back to top page