在线 O(1) 逆元

我们需要求出 \(x^{-1}\),转化一下,如果我们能对每个 \(x\) 找到一个非零的 \(u\),使得 \((xu)^{-1}\) 可以在有限时间内预处理(即 \(xu\) 的绝对值较小),那么直接根据 \(x^{-1}=(xu)^{-1}\times u\) 计算即可。下文默认 \(a\bmod p\) 为取模后在 \([0,p)\) 中的那个值。

\(B=\lceil p^{1/3}\rceil\),令 \(x=aB+c\),其中 \(c\in[0,B)\)。显然不能针对每个 \(x\) 预处理,现在我们希望对于每个 \(a\) 找到一个 \(u\in[0,B]\) 使得 \(aBu\bmod p=v\),且 \(|v|\le B^2\)。那么对于所有 \(x=aB+c\),有 \(xu\equiv v+cu\pmod p\),所以 \(|xu|<2B^2\),可以进行预处理。

然后我们说明 \(u\) 一定存在:固定一个 \(a\),考虑所有的 \(aBu,u\in[0,B]\) 在模 \(p\) 意义下从小到大排列,存在 \(B\) 个相邻间隔,根据抽屉原理,一定存在一个间隔 \(i,j\),使 \((aBi-aBj)\bmod p\le\frac pB\),由于 \(\frac pB\le B^2\),我们令 \(u=|i-j|\) 即可。

但是现在枚举 \(a,u\) 的话复杂度还是会炸,考虑枚举 \(u\) 贡献到 \(a\)。令 \(d=Bu\),则 \(u\)\(a\) 合法当且仅当存在 \(k\) 使得 \(|ad-kp|\le B^2\)。那么 \(a\in[\frac{kp-B^2}d,\frac{kp+B^2}d]\),暴力枚举 \(a\) 即可。而由于 \(aBu<pu\),所以离 \(aBu\) 最近的 \(k\) 的范围是 \(\mathcal{O}(u)\) 的,可以贡献到的 \(a\) 的区间长度 \(\frac{2B^2}{Bu}\)\(\mathcal{O}(\frac Bu)\) 的。于是对于每个 \(u\) 的复杂度是 \(\mathcal{O}(B)\) 的,预处理总复杂度为 \(\mathcal{O}(B^2)\)

时间复杂度 \(\mathcal{O}(p^{2/3})\sim\mathcal{O}(1)\)。实现时并不需要显式地进行上述枚举,只要枚举 \(u\) 后暴力跳,如果对于当前 \(a\) 合法就用 \(u\)\(-u\) 更新即可,时间复杂度仍然可以保证。

#include <bits/stdc++.h>
#include "inv.h"
using namespace std;
typedef long long ll;
typedef unsigned long long ull;
// typedef __int128 i128;
typedef pair<int, int> pii;
const int V = 1 << 21 | 1, B = 1024, mod = 998244353;
template<typename T>
void debug(const T &t) { cout << t << endl; }
template<typename Type, typename... Types>
void debug(const Type& arg, const Types&... args) {
    cout << arg << ' ';
    debug(args...);
}
#ifdef LOCAL
#define dbg(...) cout << "[" << #__VA_ARGS__ << "]: ", debug(__VA_ARGS__)
#else
#define dbg(...) 1
#endif
int k, pool[V << 1], *iv = pool + V;
struct Info {
    int u, v;
} mp[V];
void init(int _) {
    int lim = mod >> 10, _lim = mod - lim;
    for (int u = 1; u <= B; u++) {
        int cur = 0, d = u * B;
        for (int a = 0; a <= lim; ) {
            if (cur <= lim) mp[a] = {u, cur};
            else if (cur > _lim) mp[a] = {-u, mod - cur};
            else {
                int tmp = (_lim - cur) / d;
                cur += tmp * d; a += tmp;
            }
            cur += d; a++;
            if (cur >= mod) cur -= mod;
        }
    }
    iv[1] = 1;
    for (int i = 2; i < V; i++) iv[i] = (ll)iv[mod % i] * (mod - mod / i) % mod;
    for (int i = 1; i < V; i++) iv[-i] = mod - iv[i];
}
int inv(int x) {
    auto [u, v] = mp[x >> 10];
    int z = u * (x & 1023) + v;
    return (ll)(u + mod) * iv[z] % mod;
}
posted @ 2026-08-20 08:50  循环一号  阅读(8)  评论(0)    收藏  举报