60pts求调
查看原帖
60pts求调
533915
Acee楼主2023/7/5 23:10

code:

#include <iostream>
#include <cstdio>
#include <algorithm>
#include <cstring>
#include <climits>
#include <cstdlib>
//#include <assert.h>
//#include <bitset>
//#include <set>
//#include <map>
//#include <stack>
//#include <queue>
//#include <vector>
//#include <unordered_map>
#include<ext/pb_ds/assoc_container.hpp>
using namespace std;
using ll = long long;
using namespace __gnu_pbds;
#define int long long
namespace Main {
	const int MAXN = 5e6, N = 5e6 + 5;
	int mod, n;
	int prime[MAXN + 5];
	bool pri[MAXN + 5];
	int cur, inv6, inv2;
	int f[N], phi[N];
	int sum_phi[N];
	gp_hash_table<int, int> vis;
	int qpow(int a, int b) {
		int ans = 1 % mod;
		while (b) {
			if (b & 1) ans = 1ll * ans * a % mod;
			b >>= 1;
			a = 1ll * a * a % mod;
		}
		return ans;
	}
	int sum1(int x) {
		x %= mod;
		return 1ll * x % mod * (x + 1) % mod * inv2 % mod;
	}
	int sum2(int x) {
		x %= mod;
		return 1ll * x % mod * (x + 1) % mod * (2 * x + 1) % mod;
	}
	void euler() {
		phi[1] = 1;
		sum_phi[1] = 1;
		for (int i = 2; i <= MAXN; ++i) {
			if (!pri[i]) {
				prime[++cur] = i;
				f[i] = i;
				phi[i] = i - 1;
			}
			for (int j = 1; j <= cur && 1ll * i * prime[j] <= MAXN; ++j) {
				pri[i * prime[j]] = true;
				f[i * prime[j]] = prime[j];
				if (f[i] == prime[j]) phi[i * prime[j]] = phi[i] * prime[j];
				else phi[i * prime[j]] = phi[i] * (prime[j] - 1);
				if (i % prime[j] == 0) break;
			}
			sum_phi[i] = (sum_phi[i - 1] + 1ll * phi[i] * i % mod * i % mod) % mod;
		}
	}
	int ask(int n) {
		if (n <= MAXN) return sum_phi[n];
		if (vis[n]) return vis[n];
		int sum = sum1(n % mod) % mod, ans = 0;
		for (int l = 2, r, x; l <= n; l = r + 1) {
			x = n / l;
			r = n / x;
			ans = (ans + 1ll * (1ll * sum2(r) - sum2(l - 1) + mod) % mod * ask(x) % mod) % mod;
		}
		return vis[n] = ((sum - ans + mod) % mod);
	}
	int main() {
		ios :: sync_with_stdio(false);
		cin.tie(0), cout.tie(0);
		cin >> mod >> n;
		inv2 = qpow(2, mod - 2);
		inv6 = qpow(6, mod - 2);
		euler();
		int ans = 0;
		for (int l = 1, r, x; l <= n; l = r + 1) {
			x = n / l;
			r = n / x;
			ans = (ans + 1ll * sum1(x) % mod * sum1(x) % mod * ((ask(r) - ask(l - 1)) + mod) % mod) % mod;
		}
		cout << ans << '\n';
		return 0;
	}
}
signed main() {
	Main :: main();
	return 0;
}
/*
5 6
1 2 1
1 3 2
2 4 3
3 5 4
3 4 3
4 5 6
*/
2023/7/5 23:10
加载中...