公式

F - Chebyshev Cafe 解説 by sheyasutaka


\(N^2\) 個あるマスそれぞれの \(f(i,j)\) の値を求めることを考えます.マス \((i,j)\) に住む参加者の人数を \(C_{i,j}\) と書くことにします.

マス \((i,j)\) から \(f(i+x,j+y)\) への寄与

\(C_{i,j}\) の \(f(i+x,\ j+y)\) に対する寄与 \(d[x][y]\) は,以下のように書けます.

  • \(d[x][y] = \min(\max(|x|,\ |y|),\ N) \times C_{i,j}\)

ここで,累積和によって \(d\) が得られるような \(2\) 次元数列 \(d'\) を考えます.つまり,任意の \((x,y)\) で \(\displaystyle d[x][y] = NC_{i,j} + \sum_{x' \leq x,\ y' \leq y} d'[x'][y']\) が成り立つような \(d'\) を考えます.これは \(d'[x][y] = d[x][y] - d[x-1][y] - d[x][y-1] + d[x-1][y-1]\) で得ることができ,具体的な値は以下のようになります.

  • \(-N+1 \leq k \leq N\) を満たす整数 \(k\) について,\(d'[k][k] = (-1) \times C_{i,j}\)
  • \(-N+1 \leq k \leq N\) を満たす整数 \(k\) について,\(d'[k][-k+1] = (+1) \times C_{i,j}\)
  • 上のいずれにも該当しない位置では \(d'[x][y] = 0\)

これは,\(2\) 種類の斜め方向の imos 法によって,定数個の位置への加算によってあらわすことができます.

すべてのマスの寄与

\(N^2\) 個すべてのマスの寄与は,マス \((i,j)\) の寄与 \(d[x][y]\) を位置 \((i+x,\ j+y)\) に加算することで求まります.よって,各マスに対応する \(d'\) を適切な位置に加算したものの \(2\) 次元累積和を取り,全体に \(N \sum_{i,j} C_{i,j}\) を加算することで,すべての \(f(i,j)\) の値が求まります.

実装例 (C++)

#include <iostream>
using std::cin;
using std::cout;
using std::cerr;
using std::endl;
#include <vector>
using std::vector;
using std::pair;
#include <map>
using std::map;
using std::max;
using std::min;


#ifdef DEBUG
const int debug = 1;
#else
const int debug = 0;
#endif

using ll = int64_t;
using P = pair<ll, ll>;

const ll FOD = 998244353;
#include <atcoder/modint>
using mint = atcoder::modint998244353;
#include <atcoder/convolution>
using atcoder::convolution;

ll n, k, h, w, m;
vector<ll> a, b;

inline void output (const ll x) {
	cout << x << "\n";
}

void solve() {
	ll c[n][n];
	for (ll i = 0; i < n; i++) {
		for (ll j = 0; j < n; j++) {
			c[i][j] = a[i] * b[j] % m;
		}
	}

	const ll nn = 3*n+1;
	const ll offset = n;
	ll sx[nn][nn], sy[nn][nn], s[nn][nn];
	for (ll i = 0; i < nn; i++) {
		for (ll j = 0; j < nn; j++) {
			sx[i][j] = 0;
			sy[i][j] = 0;
			s[i][j] = 0;
		}
	}

	ll outsum = 0;
	for (ll i = 0; i < n; i++) {
		for (ll j = 0; j < n; j++) {
			outsum += c[i][j] * n;

			sx[(offset+i) - (n-1) - 0][(offset+j) - (n-1) - 0] += c[i][j] * (-1);
			sx[(offset+i) + (n+0) + 1][(offset+j) + (n+0) + 1] -= c[i][j] * (-1);

			sy[(offset+i) - (n-1) - 0][(offset+j) + (n+0) + 0] += c[i][j] * (+1);
			sy[(offset+i) + (n+0) + 1][(offset+j) - (n-1) - 1] -= c[i][j] * (+1);

		}
	}

	// sx, sy
	for (ll i = 0; i < nn; i++) {
		for (ll j = 0; j < nn; j++) {
			if (i - 1 >= 0 && j - 1 >= 0) sx[i][j] += sx[i - 1][j - 1];
			if (i - 1 >= 0 && j + 1 < nn) sy[i][j] += sy[i - 1][j + 1];

			s[i][j] = sx[i][j] + sy[i][j];
		}
	}

	for (ll i = 0; i < nn; i++) {
		ll x = 0;
		for (ll j = 0; j < nn; j++) {
			s[i][j] += x;
			x = s[i][j];
		}
	}
	for (ll j = 0; j < nn; j++) {
		ll x = 0;
		for (ll i = 0; i < nn; i++) {
			s[i][j] += x;
			x = s[i][j];
		}
	}

	ll ans = 0;
	for (ll i = 0; i < n; i++) {
		for (ll j = 0; j < n; j++) {
			ll item = outsum + s[offset+i][offset+j];
			if (debug) {
				cerr << i << " " << j << ": " << item << endl;
			}

			ans ^= (item + i*n + j);
		}
	}

	output(ans);


	return;
}

int main (void) {
	std::cin.tie(nullptr);
	std::ios_base::sync_with_stdio(false);

	cin >> n;
	cin >> m;
	a.resize(n);
	b.resize(n);
	for (ll i = 0; i < n; i++) {
		cin >> a[i];
	}
	for (ll i = 0; i < n; i++) {
		cin >> b[i];
	}

	
	solve();

	return 0;
}

投稿日時:
最終更新: