Official

E - ネットワークの巡回点検 / Network Patrol Inspection Editorial by kyopro_friends


\(N=1\) のとき答えは \(1\) です。以下、\(N\geq 2\) とします。サーバー \(1\) の存在は無視し、サーバー \(2\) から点検を始めるとしてよいです。

アイディア

サーバー \(k\) の次に点検するサーバー \(k'\)がどれになるかを考えます。\(k\) の素因数の集合を \(P_k\) とするとき、\(k'\)\(P_k\) のいずれかの要素の倍数です。よって、「\(p\) の倍数であって、未点検のサーバーのうち、番号が \(x\) 以上の最小のものは?」を全ての素数 \(p\) について高速に求めることができれば良さそうです。

解法

\(N\) 以下の全ての数について、素因数の集合を求めておきます。これはエラトステネスの篩と同様にして \(O(N\log\log N)\) で求めることができます。また各素数 \(p\) についてordered set \(S_p\) を「\(p\) の倍数であって、未点検のサーバーの番号からなる集合」と定めます。全ての \(p\) について初期状態の \(S_p\) を求めることは適切な実装により \(O(N\log\log N)\) でできます。
操作の過程において、各 \(S_p\) から点検したサーバーの番号を削除することと「\(p\) の倍数であって、未点検のサーバーのうち、番号が \(x\) 以上の最小のものを求めること」によりordered set に対する操作を合計 \(O(N\log\log N)\) 回行うため、計算量は全体で \(O(N\log N \log\log N)\) となります。

なお、この解法では、以下の3種類の操作を行うために ordered set を用いました。

  • 指定された要素たちで初期化
  • 指定された要素を削除
  • \(x\) 以上の最小の要素を求める

途中で要素の追加が行われないことから、ordered setの代わりにDSUを用いることでもこれらの処理を行うことができます。
(削除の代わりに”右隣”の要素と連結し、連結成分内の最大要素を取得できるようにしておく)
この場合、ordered set の操作に掛かっていた \(O(\log N)\) 時間を \(O(\alpha(N))\) 時間にすることができ、全体の計算量を \(O(N\log\log N\alpha(N))\) としてこの問題を解くこともできます。

実装例 (C++)

#include<bits/stdc++.h>
using namespace std;

int main(){
  int n;
  cin >> n;
  if(n == 1){
    cout << 1 << endl;
    return 0;
  }

  vector<vector<int>> prime_factors(n+1);
  vector<set<int>> multi(n+1);
  for(int p=2; p<=n; p++){
    if(prime_factors[p].size() == 0){
      // p が素数
      for(int q=p; q<=n; q+=p){
        prime_factors[q].push_back(p);
        multi[p].insert(q);
      }
    }
  }

  vector<bool> visited(n+1);
  int ans = 0;
  for(int start=2; start<=n; start++){
    if(visited[start])continue;
    ans++;
    int crr = start;
    while(crr <= n){
      visited[crr] = true;
      int nxt = 1e9;
      for(auto p: prime_factors[crr]){
        auto it = multi[p].find(crr);
        it = multi[p].erase(it);
        if(it != multi[p].end()){
          nxt = min(nxt, *it);
        }
      }
      crr = nxt;
    }
  }
  cout << ans << endl;
}

実装例 (Python)

sortedcontainers の SortedSet ではTLEになったため、以下の実装例では tatyamさんの作成した SortedSet を利用しています。

# https://github.com/tatyam-prime/SortedSet/blob/main/SortedSet.py
import math
from bisect import bisect_left, bisect_right
from typing import Generic, Iterable, Iterator, TypeVar
T = TypeVar('T')

class SortedSet(Generic[T]):
    BUCKET_RATIO = 16
    SPLIT_RATIO = 24
    
    def __init__(self, a: Iterable[T] = []) -> None:
        "Make a new SortedSet from iterable. / O(N) if sorted and unique / O(N log N)"
        a = list(a)
        n = len(a)
        if any(a[i] > a[i + 1] for i in range(n - 1)):
            a.sort()
        if any(a[i] >= a[i + 1] for i in range(n - 1)):
            a, b = [], a
            for x in b:
                if not a or a[-1] != x:
                    a.append(x)
        n = self.size = len(a)
        num_bucket = int(math.ceil(math.sqrt(n / self.BUCKET_RATIO)))
        self.a = [a[n * i // num_bucket : n * (i + 1) // num_bucket] for i in range(num_bucket)]

    def __iter__(self) -> Iterator[T]:
        for i in self.a:
            for j in i: yield j

    def __reversed__(self) -> Iterator[T]:
        for i in reversed(self.a):
            for j in reversed(i): yield j
    
    def __eq__(self, other) -> bool:
        return list(self) == list(other)
    
    def __len__(self) -> int:
        return self.size
    
    def __repr__(self) -> str:
        return "SortedSet" + str(self.a)
    
    def __str__(self) -> str:
        s = str(list(self))
        return "{" + s[1 : len(s) - 1] + "}"

    def _position(self, x: T) -> tuple[list[T], int, int]:
        "return the bucket, index of the bucket and position in which x should be. self must not be empty."
        for i, a in enumerate(self.a):
            if x <= a[-1]: break
        return (a, i, bisect_left(a, x))

    def __contains__(self, x: T) -> bool:
        if self.size == 0: return False
        a, _, i = self._position(x)
        return i != len(a) and a[i] == x

    def add(self, x: T) -> bool:
        "Add an element and return True if added. / O(√N)"
        if self.size == 0:
            self.a = [[x]]
            self.size = 1
            return True
        a, b, i = self._position(x)
        if i != len(a) and a[i] == x: return False
        a.insert(i, x)
        self.size += 1
        if len(a) > len(self.a) * self.SPLIT_RATIO:
            mid = len(a) >> 1
            self.a[b:b+1] = [a[:mid], a[mid:]]
        return True
    
    def _pop(self, a: list[T], b: int, i: int) -> T:
        ans = a.pop(i)
        self.size -= 1
        if not a: del self.a[b]
        return ans

    def discard(self, x: T) -> bool:
        "Remove an element and return True if removed. / O(√N)"
        if self.size == 0: return False
        a, b, i = self._position(x)
        if i == len(a) or a[i] != x: return False
        self._pop(a, b, i)
        return True
    
    def lt(self, x: T) -> T | None:
        "Find the largest element < x, or None if it doesn't exist."
        for a in reversed(self.a):
            if a[0] < x:
                return a[bisect_left(a, x) - 1]

    def le(self, x: T) -> T | None:
        "Find the largest element <= x, or None if it doesn't exist."
        for a in reversed(self.a):
            if a[0] <= x:
                return a[bisect_right(a, x) - 1]

    def gt(self, x: T) -> T | None:
        "Find the smallest element > x, or None if it doesn't exist."
        for a in self.a:
            if a[-1] > x:
                return a[bisect_right(a, x)]

    def ge(self, x: T) -> T | None:
        "Find the smallest element >= x, or None if it doesn't exist."
        for a in self.a:
            if a[-1] >= x:
                return a[bisect_left(a, x)]
    
    def __getitem__(self, i: int) -> T:
        "Return the i-th element."
        if i < 0:
            for a in reversed(self.a):
                i += len(a)
                if i >= 0: return a[i]
        else:
            for a in self.a:
                if i < len(a): return a[i]
                i -= len(a)
        raise IndexError
    
    def pop(self, i: int = -1) -> T:
        "Pop and return the i-th element."
        if i < 0:
            for b, a in enumerate(reversed(self.a)):
                i += len(a)
                if i >= 0: return self._pop(a, ~b, i)
        else:
            for b, a in enumerate(self.a):
                if i < len(a): return self._pop(a, b, i)
                i -= len(a)
        raise IndexError
    
    def index(self, x: T) -> int:
        "Count the number of elements < x."
        ans = 0
        for a in self.a:
            if a[-1] >= x:
                return ans + bisect_left(a, x)
            ans += len(a)
        return ans

    def index_right(self, x: T) -> int:
        "Count the number of elements <= x."
        ans = 0
        for a in self.a:
            if a[-1] > x:
                return ans + bisect_right(a, x)
            ans += len(a)
        return ans

# ========================================

N = int(input())
if N == 1:
  print(1)
  exit()

prime_factors = [[] for _ in range(N+1)]
multi = [SortedSet() for _ in range(N+1)]
for p in range(2, N+1):
  if len(prime_factors[p]) == 0:
    multi[p] = SortedSet(range(p, N+1, p))
    for q in range(p, N+1, p):
      prime_factors[q].append(p)

visited = [False] * (N+1)
ans = 0
for start in range(2, N+1):
  if visited[start]:
    continue
  ans += 1
  crr = start
  while crr <= N:
    visited[crr] = True
    nxt = 10**9
    for p in prime_factors[crr]:
      multi[p].discard(crr)
      x = multi[p].gt(crr)
      if x is not None:
        nxt = min(nxt, x)
    crr = nxt

print(ans)

posted:
last update: