アルゴリズムは女の子#1

アルゴリズムは女の子#1

アルゴリズムの気持ちがわかりたい.

アルゴリズム覚えるだけが増えているので少しメモを書いていきます.

お品書き

趣旨は 実装を速攻で書けるように, です.

OJ作ったならば実装道場みたいなの作りたいですね.

  1. FFT
  2. KMP
  3. Manacher
  4. Z-algorithm

ようこそ精進道場へ

Warshall-Floyd

なんとなくDPなのはわかりますがほぼ暗記するだけみたいなところありますね.
k以下の頂点しかない場合の全点対最短路だと考えると, といった感じでしょうか.
間違えないので書きません.

FFT

これだけで記事一つになりそう…

実際実装するなら毎回導出している時間はないですね.
というわけで実装を暗記する,ということも必要ですね.
long doubleだと心配なので,NTT, CRTでの復元は使えるようになっておきたいです.
任意modはよくわからないです.

基本

complex<double> で書いてみます.
原始根には exp(2πi/n)exp(2 \pi i / n)を使います.
高校数学の複素数さえわかればガウス平面で図形的に直行性を感じることができます(また(有限)等比級数の和の公式からも示せます)
NTTを考えるにしても,この環が見えているとだいぶ理解しやすいと思います!

// a.size() is 2 ^ k
vector<comp> fft(vector<comp> a, bool inverse = false) {
  int n = a.size();
  if(n == 1) return a;
  vector<comp> a0(n/2), a1(n/2);
  for(int i = 0; i < n/2; i++) a0[i] = a[i * 2], a1[i] = a[i * 2 + 1];
  a0 = fft(a0); a1 = fft(a1);
  comp zeta = comp(cos(2 * PI / n), sin(2 * PI / n));
  comp powZeta = comp(1, 0);
  int n2 = n / 2;
  for(int i = 0; i < n; i++) {
    int ii = i;
    if(inverse) ii = n - i;
    a[i] = a0[ii % n2] + powZeta * a1[ii % n2];
    if(inverse) a[i] *= comp(1.0 / n, 0);
    if(inverse) powZeta /= zeta;
    else powZeta *= zeta;
  }
  return a;
}

template<typename T>
vector<comp> conv(vector<T> at, vector<T> bt) {
  vector<comp> a(at.size()), b(bt.size());
  for(int i = 0; i < (int) a.size(); i++) a[i] = at[i];
  for(int i = 0; i < (int) b.size(); i++) b[i] = bt[i];
  int n = a.size() + b.size();
  int m = 1;
  while(m < n) m <<= 1;
  a.resize(m); b.resize(m);
  a = fft(a); b = fft(b);
  vector<comp> c(m);
  for(int i = 0; i < m; i++) c[i] = a[i] * b[i];
  return fft(c, true);
}

ちなみに直行性というものを知らなかったのですが、調べてみると、“関数は実はベクトルなんです” と書いてあり興味が出てきた. ふしゃー!

精度に関してですが

次数dd,入力の係数の最大値uuとすると,
精度xxなら du2<xdu^2 \lt x でなければいけないようです(kirikaさんのPDFのFFTの項を読んでみてください)
ATC001-Cの場合d=2×105d = 2 \times 10^5,u=100u = 100でまあdu2<253du^2 \lt 2^{53}な気がしますが実装が雑なせいかpowZetaあたりで誤差がガッツリ出るのかな,long doubleじゃないと通りませんでした.

NTT

NTT(Number-Theoretic Transform; 数論変換)(または FMT(Fast Modulo Transform)).
整数環を使います, 完.

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

ll extgcd(ll a, ll b, ll &x, ll &y) {
  if(b == 0) {
    x = 1; y = 0;
    return a;
  }
  ll d = extgcd(b, a % b, y, x);
  y -= a / b * x;
  return d;
}

ll modinv(ll a, ll mod) {
  ll x = 0, y = 0;
  extgcd(a, mod, x, y);
  return (x % mod + mod) % mod;
}

ll modpow(ll a, ll b, ll mod) {
  a = (a % mod + mod) + mod;
  ll r = 1;
  while(b) {
    if(b & 1) r = r * a % mod;
    a = a * a % mod;
    b >>= 1;
  }
  return r;
}

template<ll mod>
struct ModInt{
  ll val;

  ModInt() : val(0) {}
  ModInt(ll val) : val((val % mod + mod) % mod) {}
  ll get() const { return val; }

  ModInt operator+(ModInt<mod> rhs) {
    return ModInt<mod>(val + rhs.val);
  }
  ModInt operator*(ModInt<mod> rhs) {
    return ModInt<mod>(val * rhs.val);
  }
  ModInt operator/(ModInt<mod> rhs) {
    return ModInt<mod>(val * rhs.inv().val);
  }
  ModInt &operator+=(ModInt<mod> rhs) {
    val = ((val + rhs.val) % mod + mod) % mod;
    return *this;
  }
  ModInt &operator*=(ModInt<mod> rhs) {
    val = (val * rhs.val % mod + mod) % mod;
    return *this;
  }
  ModInt &operator/=(ModInt<mod> rhs) {
    val = (val * rhs.inv().val % mod + mod) % mod;
    return *this;
  }
  ModInt inv() {
    return ModInt<mod>(modinv(val, mod));
  }
};

template<ll mod, ll primitive>
struct NTT {
  const ll Mod = mod;
  using Int = ModInt<mod>;

  vector<Int> fft(vector<Int> a, bool inverse = false) {
    int n = a.size(), n2 = n / 2;
    if(n == 1) return a;
    vector<Int> a0(n2), a1(n2);
    for(int i = 0; i < n2; i++) a0[i] = a[i*2], a1[i] = a[i*2+1];
    a0 = fft(a0); a1 = fft(a1);
    Int zeta(modpow(primitive, (mod - 1) / n, mod)); // todo
    if(inverse) zeta = zeta.inv();
    Int powZeta(1);
    for(int i = 0; i < n; i++) {
      int ii = i;
      if(inverse) ii = n - i;
      a[i] = a0[ii % n2] + powZeta * a1[ii % n2];
      if(inverse) a[i] /= Int(n);
      powZeta *= zeta;
    }
    return a;
  }

  template<typename T>
  vector<Int> conv(vector<T> at, vector<T> bt) {
    int deg = at.size() + bt.size();
    int n = 1;
    while(n < deg) n <<= 1;
    vector<Int> a(n), b(n);
    for(int i = 0; i < (int) at.size(); i++) a[i] = Int(at[i]);
    for(int i = 0; i < (int) bt.size(); i++) b[i] = Int(bt[i]);
    a = fft(a); b = fft(b);
    vector<Int> c(n);
    for(int i = 0; i < n; i++) c[i] = a[i] * b[i];
    return fft(c, true);
  }
};

NTT<(1 << 25) * 5 + 1, 3> ntt1;
NTT<(1 << 26) * 7 + 1, 3> ntt2;

template<typename T>
vector<ll> conv(vector<T> a, vector<T> b) {
  auto c1 = ntt1.conv(a, b);
  auto c2 = ntt2.conv(a, b);
  vector<ll> c(c1.size());
  for(int i = 0; i < (int) c.size(); i++) {
    // garner
    ll x = c1[i].val;
    ll v1 = (c2[i].val - c1[i].val);
    v1 = (v1 % ntt2.Mod + ntt2.Mod) % ntt2.Mod;
    v1 = v1 * modinv(ntt1.Mod, ntt2.Mod) % ntt2.Mod;
    x += v1 * ntt1.Mod;
    c[i] = x;
  }
  return c;
}

int main() {
  ios::sync_with_stdio(false), cin.tie(0);
  int n; cin >> n;
  vector<int> a(n + 1), b(n + 1);
  for(int i = 0; i < n; i++) cin >> a[i + 1] >> b[i + 1];
  auto c = conv(a, b);
  for(int i = 1; i <= n * 2; i++) {
    cout << c[i] << endl;
  }
}

ATC001-Cの提出そのまんまですはい.

復元に関してですが

xn1modm1xn2modm2 \begin{aligned} x \equiv n_1 \mod m_1 \\ x \equiv n_2 \mod m_2 \end{aligned}
みたいな形であらわされるxxを復元するというもの.

ATCに限っては1つのNTTで足りますが,あえて小さめのmodで復元をしてみました.

最初CRTでやろうと思ったのですが,二つのmod,たとえばm1m_1m2m_2の積のmodなので,倍精度だと足りないのでやめました.

3変数以上になると必要な精度も増えてくるので,Garnerのアルゴリズムを使いました.

Garnerのアルゴリズム

参考 : math314さんとか, kirika_compさんとか, こういうのとか

これはよくて,計算途中がせいぜいmim_iのmodなのではやいですし任意のほかのmodに変換することもできます.

実装はちょっとずつ練習していこうと思います.

extgcd

FFTと似た理由で,やっぱり覚えておいたほうがいいですね.

ll extgcd(ll a, ll b, ll &x, ll &y) {
	if(b == 0) {
		x = 1, y = 0;
		return a;
	}
	int d = extgcd(b, a % b, y, x);
	y -= x * (a / b);
	return d;
}

文字列アルゴリズム

そのアルゴリズムの気持ちを知るには試しに手で動かしてみることが一番いいです.
例示は理解の試金石,ですね!1

アルゴリズム名 参考
KMP (MP) aabaabaaa Link
Manacher abcbabcba Link
Z-algorithm aaabaaaab Link

以上の例を自分で手で処理できるようになれば,実装に迷うこともなくなってくるかもしれません!

参考はすべてSnukeさんの記事です.

MP

KMP法はKnuth-Morris-Pratt法の略です.

// size of longest common suffix and prefix in s[0:i-1]
vector<int> MP(string s) {
  int n = s.size();
  vector<int> A(n + 1);
  A[0] = -1;
  int j = -1;
  for(int i = 0; i < n; i++) {
    while(j >= 0 && s[i] != s[j]) j = A[j];
    A[i + 1] = ++j;
  }
  return A;
}

最小周期長

MPを利用すると,最小周期長も求められます.

vector<int> cycle(string s) {
  auto mp = MP(s);
  vector<int> len(s.size());
  for(int i = 0; i < (int) s.size(); i++)
    len[i] = i + 1 - mp[i + 1];
  return len;
}

いい関数名が浮かびません,cycle(string)以外に何かあれば教えてくだしゃい.

例題

ARC060-F: 最良表現

KMP

あなたのMPにKnuthパワー.

基本アイデアは以下のコードのような感じ.

vector<int> KMP(string s) {
  int n = s.size();
  vector<int> kmp(n + 1), mp(n + 1);
  kmp[0] = mp[0] = -1;
  int j = -1;
  for(int i = 0; i < n; i++) {
    while(j >= 0 && s[i] != s[j]) j = kmp[j];
    kmp[i + 1] = mp[i + 1] = ++j;
    if(i + 1 < n && s[i + 1] == s[j]) kmp[i + 1] = kmp[j];
  }
  return mp;
}

ある瞬間(iについて)の最悪計算量はO(N)からO(logN)になりますが,
全体で均し計算量O(N)なので,問題はこのアイデアがうまく使えるかどうか.

参考 : ポテチさんのブログ : MP法とKMP法の違い

TODO : xmas contest 2015 D - Destroy the Duplicated Poem

verify : F - 最良表現

案の定間違ってたりしたのでverifyしてよかった.

Manacher

vector<int> Manacher(string s) {
  int n = s.size();
  int i = 0, j = 0;
  vector<int> R(n);
  while(i < n) {
    while(i - j >= 0 && i + j < n &&
        s[i - j] == s[i + j]) ++j;
    R[i] = j;
    int k = 1;
    while(i - k >= 0 && R[i - k] < j - k) R[i + k] = R[i - k], ++k;
    i += k, j -= k;
  }
  return R;
}

たぷんおっけいじゃないでぷか?
TODO : verify

Z algorithm

// size of longest common prefix between s and s[i...]
vector<int> Zalgorithm(string s) {
  int n = s.size();
  vector<int> Z(n);
  Z[0] = n;
  int i = 1, j = 0;
  while(i < n) {
    while(i + j < n && s[j] == s[i + j]) ++j;
    Z[i] = j;
    if(j == 0) { ++i; continue; }
    int k = 1;
    while(i + k < n && Z[k] < j - k) Z[i + k] = Z[k], ++k;
    i += k, j -= k;
  }
  return Z;
}

verify : F - 最良表現

最良表現使いまくってるね.(解説PDFの各種証明もじっくり読みたいところ)


第二回も作ろ~

[追記]

ぶっちゃけ, 実装全然覚えてられない!w


  1. 数学ガール,これはいい. ↩︎

このブログの人気の投稿

YouTube Iridiumの紹介

うくこん

TDPC - T フィボナッチ