1. 先把题目读明白:不是模拟,而是周期
1.1 题面还原与输入输出格式
最近在刷题群和社区里频繁看到有人在问“反射计数”这道题,分值 200,提交语言支持 Java、JS、Python、C,几乎每家在线评测系统都会收录。我参考的题面是这样的:
一个 n 行 m 列的网格,坐标从 0 开始。一个点从 (x, y) 出发,每单位时间沿方向向量 (dx, dy) 移动 1 格,dx、dy 只能取 1 或 -1。碰到边界时,对应方向取反:X 方向越界则 dx 变号,Y 方向越界则 dy 变号,角点上两个方向同时变号。给定目标点 (tx, ty) 和总时间 T,第 0 秒和第 T 秒都算在内,求 0 到 T 秒内该点正好位于目标点的次数。
输入格式如下,不同平台可能只调整数据排列顺序,解析逻辑改一下就行:
n m x y dx dy tx ty T样例:
5 5 2 2 1 1 3 3 10输出:
3为什么是 3?我后面会用手算验证。第一次看到这题时,我第一反应和大多数人一样:写个循环一秒一秒模拟,反正无非是坐标加减和方向取反。但冷静下来一想,T 如果开到 10^9,模拟必然超时。而且边界反射的代码看着简单,真正写对边界情况也不容易。这题放在 200 分档,考察的就是能不能跳出“模拟”这个舒适区,找到周期规律。
1.2 为什么直接模拟容易翻车
模拟的思路非常直白:每个时间步判断当前坐标是否等于目标点,然后按反射规则更新位置。你甚至可以很快写出一个能过样例的版本。但翻车点有两个。
第一是复杂度。T 只要到 10^8 以上,O(T) 的模拟就开始吃力,如果到 10^9 基本必挂。第二是反射的细节。我第一次写模拟时用的是“先移动,再判断越界,越界后反向并往回走”,结果在边界上的行为是错的。比如点位于 x=0,方向 dx=-1,正确行为应该是弹回 x=1,但错误的写法会让它停在原地或者位置错乱。这种问题普通样例根本测不出来,一上特殊数据就暴露。
所以网上虽然很多人贴过这题的模拟写法,但真到评审机上跑,很多版本过不了大数据。更稳妥的路线是:先识别出路径的周期,再用数论工具直接算出结果。下面的数学模型就是这个思路。
1.3 从一个例子手算出周期规律
只看 X 方向。网格有 n 列,坐标范围 0 到 n-1。点在这个区间里来回跑,跑一个完整来回走过的距离是 2(n-1),所以 X 方向坐标序列的周期是:
Px = 2 * (n - 1)同理,Y 方向的周期是:
Py = 2 * (m - 1)在一个周期内,X 方向上某个内部坐标(比如 x=3)会被经过两次:一次向右路过,一次向左路过;而边界坐标 0 和 n-1 只会被经过一次。这个“内部两个解、边界一个解”的差异,正是后面代码里最容易漏掉的地方。
拿样例手算:n=5,m=5,起点 (2,2),方向 (1,1),目标点 (3,3)。X 方向周期 Px=8,x(t) 恰好等于 3 的时间满足:
t ≡ 1 (mod 8) 或者 t ≡ 3 (mod 8)Y 方向完全一样。两个条件都满足的时间在 [0, 10] 内只有 t=1、3、9,所以答案是 3。
这里“合并两个同余条件”的工作,就是中国剩余定理(CRT)要干的活。
2. 核心数学模型:一维折返 + 中国剩余定理
2.1 一维折返怎么变成同余方程
把 X 方向单独拿出来,定义:
Px = 2 * (n - 1)在展开平面上,点其实是沿着直线匀速前进的,真实网格坐标只是对这条直线位置做了一个“折叠”映射。写成公式:
真实坐标 x(t) = fold((x0 + dx * t) mod Px)fold(z) 的含义是:如果 z ≤ n-1,取 z;如果 z > n-1,取 Px - z。
要判断 x(t) 是否等于目标 tx,相当于要求 (x0 + dx * t) mod Px 落在某个集合 Sx 上。这里分两种情况:
- tx 是内部点,即 0 < tx < n-1:Sx = {tx, Px - tx}
- tx 是边界点,即 tx = 0 或 tx = n-1:Sx = {tx},因为 Px - tx 和 tx 在模 Px 意义下是同一个数
于是得到同余方程:
dx * t ≡ s - x0 (mod Px),其中 s ∈ Sx因为 dx 只能是 1 或 -1,dx 的逆元就是它自己,所以:
t ≡ dx * (s - x0) (mod Px)Y 方向完全对称:
t ≡ dy * (s - y0) (mod Py),其中 Py = 2 * (m - 1),s ∈ Sy这个推导看起来很数学,实际写代码时只需要一个 gen_residues 函数:输入初始位置、方向、周期、目标坐标和边界坐标,输出一个包含 1 到 2 个余数的列表。
2.2 二维合并:先 CRT 再计数
现在 X 方向给出了一组可能的余数,Y 方向也给出了一组。对每一对组合:
t ≡ a (mod Px) t ≡ b (mod Py)用中国剩余定理合并。合并前先算 g = gcd(Px, Py),如果 (b - a) 不能被 g 整除,说明这对组合在整数时间里没有解,直接跳过。有解时合并模是 L = lcm(Px, Py),在 [0, L-1] 里唯一对应一个解 r。
最后统计从 r 开始,以 L 为步长,不超过 T 的项数。一个完整整体周期 L 内最多贡献 4 个命中时刻,因为每个方向最多两个余数类,组合最多 2 × 2 = 4 个。统计公式:
如果 r > T:贡献 0 否则贡献 (T - r) / L + 1为什么整体周期是 lcm(Px, Py)?因为 X 方向要回到初始位置并且方向也回到初始方向,需要经过 Px 的整数倍时间;Y 方向要同时复原,必须再过 Py 的整数倍。两个条件同时满足的最小正时间就是最小公倍数。这也是“周期法”的核心所在:不需要真的等到天荒地老,在数学上直接把周期算出来。
2.3 一维退化和角点反射的特殊处理
如果 n=1 或者 m=1,周期公式里会出现 Px=0 或 Py=0,没法套 CRT。这种情况必须单独分支:
- n=1 且 m=1:网格只有一个点。目标点若是 (0,0),答案是 T+1,否则是 0。
- n=1:X 方向永远停在 0。目标 tx 不为 0,直接输出 0;tx 为 0 时,问题降维成 Y 方向上的一维往返,只做一个方向的余数类计数。
- m=1:对称处理。
角点反射不需要额外特判。在 (0,0) 且方向为 (-1,-1) 的瞬间,两个方向同时越界,dx、dy 同时取反,点从角落弹向 (1,1)。这在模拟代码里天然成立,在同余公式里也天然成立,因为周期折叠函数把 0 和 Px 看成同一位置,展开平面上角点就是网格镜像的拼接处。
3. 四种语言落地:Java、Python、Node.js、C
3.1 公共工具函数的设计思路
不管用什么语言,代码骨架都是一样的:
- gen_residues:生成某个方向的余数类列表
- crt:合并两个同余方程
- count_in_range:统计等差数列里不超过 T 的元素个数
- 主流程:处理退化情况,再枚举余数组合累加答案
CRT 合并时最核心的一步是:先把两个模除以 gcd,得到两个互质的数,再对其中一个求模逆元。当 Px 和 Py 不互质时,不能直接对 Px 求逆元。正确写法是:
g = gcd(Px, Py) diff = b - a 如果 diff % g != 0,无解 m2g = Py / g 关键:要求 Px/g 在模 m2g 下的逆元 k = diff/g 乘以该逆元,再对 m2g 取模 合并模 L = lcm(Px, Py) 最终解 r = a + Px * k,对 L 取模这里有个跨语言的坑:负数取模。Java 和 C 的%结果是负数或零,Python 会自动转非负,JavaScript 的 BigInt 也是负数保留。所以 Java、C、JS 必须自己写 norm 函数:
norm(v, mod) = ((v % mod) + mod) % mod我实际调试时,四个语言跑同一组数据,结果不一致,最后排查发现就是负数取模的差异。
3.2 Java 完整实现
import java.util.*; public class Main { static long gcd(long a, long b) { return b == 0 ? a : gcd(b, a % b); } static long lcm(long a, long b) { return a / gcd(a, b) * b; } static long exgcd(long a, long b, long[] xy) { if (b == 0) { xy[0] = 1; xy[1] = 0; return a; } long g = exgcd(b, a % b, xy); long t = xy[0]; xy[0] = xy[1]; xy[1] = t - (a / b) * xy[1]; return g; } static long invMod(long a, long mod) { long[] xy = new long[2]; exgcd(a, mod, xy); return (xy[0] % mod + mod) % mod; } static long norm(long v, long mod) { return ((v % mod) + mod) % mod; } // 生成单方向上的余数类:内部点 2 个,边界点 1 个 static List<Long> genResidues(long pos, long dir, long period, long target, long limit) { List<Long> res = new ArrayList<>(); res.add(norm((target - pos) * dir, period)); if (target > 0 && target < limit) { long another = norm((period - target - pos) * dir, period); if (!res.contains(another)) { res.add(another); } } return res; } // 合并两个同余方程,返回 [解, 模],无解返回 null static long[] crt(long a, long m1, long b, long m2) { long g = gcd(m1, m2); long diff = b - a; if (diff % g != 0) { return null; } long m2g = m2 / g; long coeff = norm(m1 / g, m2g); long invCoeff = invMod(coeff, m2g); long k = norm(diff / g, m2g) * invCoeff % m2g; long mod = lcm(m1, m2); long ans = norm(norm(a, mod) + (m1 % mod) * k, mod); return new long[]{ans, mod}; } static long countInRange(long r, long step, long T) { if (r > T) { return 0; } return (T - r) / step + 1; } public static void main(String[] args) { Scanner sc = new Scanner(System.in); long n = sc.nextLong(), m = sc.nextLong(); long x = sc.nextLong(), y = sc.nextLong(); long dx = sc.nextLong(), dy = sc.nextLong(); long tx = sc.nextLong(), ty = sc.nextLong(); long T = sc.nextLong(); long ans = 0; if (n == 1 && m == 1) { System.out.println(tx == 0 && ty == 0 ? T + 1 : 0); return; } if (n == 1) { if (tx != 0) { System.out.println(0); return; } long py = 2 * (m - 1); for (long b : genResidues(y, dy, py, ty, m - 1)) { ans += countInRange(b, py, T); } System.out.println(ans); return; } if (m == 1) { if (ty != 0) { System.out.println(0); return; } long px = 2 * (n - 1); for (long a : genResidues(x, dx, px, tx, n - 1)) { ans += countInRange(a, px, T); } System.out.println(ans); return; } long px = 2 * (n - 1); long py = 2 * (m - 1); List<Long> xs = genResidues(x, dx, px, tx, n - 1); List<Long> ys = genResidues(y, dy, py, ty, m - 1); for (long a : xs) { for (long b : ys) { long[] res = crt(a, px, b, py); if (res == null) { continue; } ans += countInRange(res[0], res[1], T); } } System.out.println(ans); } }Java 版本需要注意两点:一是 Scanner 处理多行输入时,按空格和换行自动分词,直接用nextLong()读取即可;二是极端数据下(m1 % mod) * k可能溢出 long。本题常规数据范围没事,如果平台把 n、m 都开到 10^9 级别,建议把乘法换成 BigInteger 或快速乘法。我在第 4 节会再讲这个。
3.3 Python 完整实现
import sys from math import gcd def norm(v, mod): return v % mod def gen_residues(pos, d, period, target, limit): res = [norm((target - pos) * d, period)] if 0 < target < limit: another = norm((period - target - pos) * d, period) if another not in res: res.append(another) return res def crt(a, m1, b, m2): g = gcd(m1, m2) diff = b - a if diff % g != 0: return None m2g = m2 // g coeff = (m1 // g) % m2g inv_coeff = pow(coeff, -1, m2g) k = (diff // g) * inv_coeff % m2g mod = m1 // g * m2 # 等于 lcm(m1, m2) ans = (a % mod + (m1 % mod) * k) % mod return ans, mod def count_in_range(r, step, T): if r > T: return 0 return (T - r) // step + 1 def solve(): data = list(map(int, sys.stdin.read().split())) if not data: return n, m, x, y, dx, dy, tx, ty, T = data[:9] ans = 0 if n == 1 and m == 1: print(T + 1 if tx == 0 and ty == 0 else 0) return if n == 1: if tx != 0: print(0) return py = 2 * (m - 1) for b in gen_residues(y, dy, py, ty, m - 1): ans += count_in_range(b, py, T) print(ans) return if m == 1: if ty != 0: print(0) return px = 2 * (n - 1) for a in gen_residues(x, dx, px, tx, n - 1): ans += count_in_range(a, px, T) print(ans) return px = 2 * (n - 1) py = 2 * (m - 1) for a in gen_residues(x, dx, px, tx, n - 1): for b in gen_residues(y, dy, py, ty, m - 1): res = crt(a, px, b, py) if res is None: continue r, step = res ans += count_in_range(r, step, T) print(ans) if __name__ == "__main__": solve()Python 3.8 以上的pow(a, -1, mod)可以直接求模逆元,省去手写 exgcd 的麻烦。它要求 a 和 mod 互质,而我们传入的 coeff 满足这个条件,所以可以放心用。Python 的%天然返回非负值,这让我少踩了一半负数取模的坑。另外 Python 整数没有位数限制,CRT 里的大乘法也不会溢出,这是它在这种数论题里最舒服的地方。
3.4 JavaScript(Node.js) 完整实现
const readline = require('readline'); const rl = readline.createInterface({ input: process.stdin }); let input = ''; rl.on('line', line => { input += line + '\n'; }); rl.on('close', () => { const nums = input.trim().split(/\s+/).map(BigInt); const n = nums[0], m = nums[1]; const x = nums[2], y = nums[3]; const dx = nums[4], dy = nums[5]; const tx = nums[6], ty = nums[7]; const T = nums[8]; function norm(v, mod) { return ((v % mod) + mod) % mod; } function gcd(a, b) { while (b !== 0n) { [a, b] = [b, a % b]; } return a; } function lcm(a, b) { return a / gcd(a, b) * b; } function exgcd(a, b) { if (b === 0n) return [a, 1n, 0n]; const [g, x1, y1] = exgcd(b, a % b); return [g, y1, x1 - (a / b) * y1]; } function invMod(a, mod) { const [g, x] = exgcd(a, mod); return norm(x, mod); } function genResidues(pos, d, period, target, limit) { const res = [norm((target - pos) * d, period)]; if (target > 0n && target < limit) { const another = norm((period - target - pos) * d, period); if (!res.some(v => v === another)) res.push(another); } return res; } function crt(a, m1, b, m2) { const g = gcd(m1, m2); const diff = b - a; if (diff % g !== 0n) return null; const m2g = m2 / g; const coeff = norm(m1 / g, m2g); const invCoeff = invMod(coeff, m2g); const k = norm(diff / g, m2g) * invCoeff % m2g; const mod = lcm(m1, m2); const ans = norm(norm(a, mod) + (m1 % mod) * k % mod, mod); return [ans, mod]; } function countInRange(r, step, T) { if (r > T) return 0n; return (T - r) / step + 1n; } let ans = 0n; if (n === 1n && m === 1n) { console.log((tx === 0n && ty === 0n) ? T + 1n : 0n); return; } if (n === 1n) { if (tx !== 0n) { console.log(0n); return; } const py = 2n * (m - 1n); for (const b of genResidues(y, dy, py, ty, m - 1n)) { ans += countInRange(b, py, T); } console.log(ans); return; } if (m === 1n) { if (ty !== 0n) { console.log(0n); return; } const px = 2n * (n - 1n); for (const a of genResidues(x, dx, px, tx, n - 1n)) { ans += countInRange(a, px, T); } console.log(ans); return; } const px = 2n * (n - 1n); const py = 2n * (m - 1n); for (const a of genResidues(x, dx, px, tx, n - 1n)) { for (const b of genResidues(y, dy, py, ty, m - 1n)) { const res = crt(a, px, b, py); if (res) ans += countInRange(res[0], res[1], T); } } console.log(ans); });JavaScript 版本我直接用 BigInt,因为普通 Number 在超过 2^53 时会丢精度。这道题里周期、T 都可能很大,用 Number 写出来的版本会冷不丁在极端用例上出错,而且出错很隐蔽。BigInt 所有运算都要显示写n后缀或者用 BigInt 方法,代码略啰嗦,但换来的是稳。输入解析用 readline 逐行拼字符串,再统一 split 成 BigInt 数组。注意 BigInt 的console.log会输出十进制字符串,不需要手动转换。
3.5 C 完整实现
#include <stdio.h> typedef long long ll; ll gcd(ll a, ll b) { return b ? gcd(b, a % b) : a; } ll exgcd(ll a, ll b, ll *x, ll *y) { if (b == 0) { *x = 1; *y = 0; return a; } ll x1, y1; ll g = exgcd(b, a % b, &x1, &y1); *x = y1; *y = x1 - (a / b) * y1; return g; } ll norm(ll v, ll mod) { return (v % mod + mod) % mod; } ll invMod(ll a, ll mod) { ll x, y; exgcd(a, mod, &x, &y); return norm(x, mod); } ll mulMod(ll a, ll b, ll mod) { return (ll)((__int128)a * b % mod); } int genResidues(ll pos, ll d, ll period, ll target, ll limit, ll out[2]) { out[0] = norm((target - pos) * d, period); int cnt = 1; if (target > 0 && target < limit) { out[1] = norm((period - target - pos) * d, period); if (out[1] != out[0]) { cnt = 2; } } return cnt; } int crt(ll a, ll m1, ll b, ll m2, ll *res, ll *mod) { ll g = gcd(m1, m2); ll diff = b - a; if (diff % g != 0) { return 0; } ll m2g = m2 / g; ll coeff = norm(m1 / g, m2g); ll invCoeff = invMod(coeff, m2g); ll k = mulMod(norm(diff / g, m2g), invCoeff, m2g); *mod = m1 / g * m2; ll ans = norm(norm(a, *mod) + mulMod(m1, k, *mod), *mod); *res = ans; return 1; } ll countInRange(ll r, ll step, ll T) { if (r > T) { return 0; } return (T - r) / step + 1; } int main() { ll n, m, x, y, dx, dy, tx, ty, T; scanf("%lld %lld", &n, &m); scanf("%lld %lld", &x, &y); scanf("%lld %lld", &dx, &dy); scanf("%lld %lld", &tx, &ty); scanf("%lld", &T); ll ans = 0; if (n == 1 && m == 1) { printf("%lld\n", (tx == 0 && ty == 0) ? T + 1 : 0); return 0; } if (n == 1) { if (tx != 0) { puts("0"); return 0; } ll py = 2 * (m - 1); ll residues[2]; int cnt = genResidues(y, dy, py, ty, m - 1, residues); for (int i = 0; i < cnt; ++i) { ans += countInRange(residues[i], py, T); } printf("%lld\n", ans); return 0; } if (m == 1) { if (ty != 0) { puts("0"); return 0; } ll px = 2 * (n - 1); ll residues[2]; int cnt = genResidues(x, dx, px, tx, n - 1, residues); for (int i = 0; i < cnt; ++i) { ans += countInRange(residues[i], px, T); } printf("%lld\n", ans); return 0; } ll px = 2 * (n - 1); ll py = 2 * (m - 1); ll resX[2], resY[2]; int cx = genResidues(x, dx, px, tx, n - 1, resX); int cy = genResidues(y, dy, py, ty, m - 1, resY); for (int i = 0; i < cx; ++i) { for (int j = 0; j < cy; ++j) { ll r, mod; if (crt(resX[i], px, resY[j], py, &r, &mod)) { ans += countInRange(r, mod, T); } } } printf("%lld\n", ans); return 0; }C 版本我用__int128包了一个 mulMod,专门处理 CRT 里的乘法取模。m1 * k在最坏情况下可能达到 4e18,已经逼近 long long 的极限,普通乘法一不留神就溢出。用__int128做中间乘法,结果再转回 long long,代价小又安全。多数在线评测的 C 编译器都支持__int128,如果遇到不支持的平台,可以退回到快速乘法。
另一个 C 特有的坑是scanf读入 long long 要用%lld,写%d的话大数据直接读乱。别问我怎么知道的,都是泪。
4. 对拍验证与踩坑记录
4.1 一个正确的模拟器长什么样
周期法写完,一定要用模拟法对拍。但模拟器本身也得写对。我见过太多“模拟结果错误,导致误以为周期法写错”的情况。正确的单步移动逻辑是:先判断下一步是否越界,越界就先反向,然后再移动。
def simulate(n, m, x, y, dx, dy, tx, ty, T): cnt = 0 for t in range(T + 1): if x == tx and y == ty: cnt += 1 if t == T: break if x + dx < 0 or x + dx >= n: dx = -dx if y + dy < 0 or y + dy >= m: dy = -dy x += dx y += dy return cnt注意这里必须用 x + dx 判断,而不是先把 x 改掉再判断。先移动再判断会让你在边界上得到错误位置。角点反射在这个写法里自然成立,因为两个方向分别判断,都越界就都反向,位置一步走到对角。
我用这个模拟器随机生成了几百组小数据(n、m 在 2 到 6,T 在 0 到 20),用 Python 的周期法逐组对比,结果完全一致。下面挑几个有代表性的用例说明。
4.2 几组典型测试用例
| 用例 | n m | 起点 | 方向 | 目标 | T | 模拟结果 | 周期法结果 |
|---|---|---|---|---|---|---|---|
| 基础往返 | 5 5 | 2 2 | 1 1 | 3 3 | 10 | 3 | 3 |
| 角落弹射 | 3 3 | 0 0 | -1 -1 | 1 1 | 10 | 5 | 5 |
| 一维退化 | 1 5 | 0 2 | 1 1 | 0 2 | 10 | 3 | 3 |
| 不同周期 | 3 4 | 1 2 | 1 -1 | 2 1 | 5 | 1 | 1 |
第一组就是样例。第二组验证角点反弹:n=3, m=3,起点 (0,0),方向 (-1,-1),目标 (1,1),T=10。小球在角点反弹后反复经过中心点,t=1、3、5、7、9 共 5 次,周期法算出 t ≡ 1 或 3 (mod 4),计数 5,正确。
第三组验证 n=1 退化。网格只有一列五格,小球在竖线上往返,目标 (0,2) 在 t=0、4、8 被经过,答案 3。如果不处理 n=1,周期公式里会出现除零,直接崩。
第四组验证 Px=4、Py=6 这类不同周期场景。我手算过:满足条件的时间只有一个 t=1,所以答案 1。CRT 正确合并了不同模的同余方程。
4.3 最容易翻车的几个点
第一个坑是边界点被当成内部点处理。tx=0 或 tx=n-1 只能生成一个余数类,如果按内部点生成了两个,结果会多算。这个 bug 很隐蔽,因为只要目标点落在边界且恰好某条伪解在 T 内出现,答案就会偏大。代码里我用limit = n - 1作为边界最大值,再用0 < target < limit判断内部点,就是为了从根上避开这个错误。
第二个坑是 CRT 无解。Px 和 Py 不互质时,会有一些余数组合找不出整数解,必须跳过。漏掉这个判断的话,C 和 Java 里可能算出错误的余数,甚至因为模逆元不存在而崩溃。Px=8、Py=6 时,x 方向余数 1 和 y 方向余数 3 就是无解组合,因为两个同余条件自相矛盾。
第三个坑是负余数。Java、C、JS 的取模结果可能是负数,norm 函数必须加。我在四种语言里特意都写了 norm,唯一不用改的是 Python。跨语言对拍时,同样的逻辑在 Java 和 C 上跑错、在 Python 上跑对,基本就是负数取模的问题。
第四个坑是乘法溢出。C 和 Java 的 long 在极端数据下有风险。C 我已经用 __int128 兜底;Java 如果怕溢出,最快的改法是把 CRT 里的乘法换成 BigInteger,或者限制一下输入范围。实际机试数据通常不会顶着上限出,但心里要清楚这个边界。
第五个坑是没有处理 n=1 或 m=1。如果不特判,周期变量变成 0,后面不是除零就是无限循环。这类“一维退化”的用例在评测机里几乎一定会出现,属于必拿分的点。
4.4 一个可以延伸思考的视角
用展开平面看这道题会更直观:反射等于把网格镜像展开,光点永远走直线。目标点在展开平面上有无数个镜像,问题变成“直线在某时刻是否落到某个镜像点上”。X 方向上目标镜像点的横坐标是 tx + 2k(n-1) 和 -tx + 2k(n-1),解同余方程得到的正好就是这组点。这也是为什么内部点有两个余数类、边界点只有一个——边界点在镜像展开后是自我重合的。
如果遇到变体题,比如“求反射次数而不是经过次数”,或者“求首次到达某个点的时间”,底层仍然是同一套周期模型。反射次数对应方向翻转次数,首次到达时间对应同余方程的最小非负解,改改判断条件就行。
机器上对拍通过后,我把四种语言版本都各自提交了一遍,Java 和 C 用时几乎为 0,Python 和 Node 也都在毫秒级。相比模拟法的 O(T),周期法的复杂度只有 O(log min(Px, Py)) 级别的常数,差别是量级上的。
最后说点实际的。我现在做这题,会先把 n=1、m=1、目标点在边界这三种 corner case 写在草稿纸最上方,然后才开始写代码。genResidues 和 crt 两个函数抽出来,四个语言逻辑完全一样,换语言只换语法不换结构。如果平台数据范围确实很小,模拟能过,但学会周期法之后,碰到 T 特别大的“反射计数 Plus”就不会慌了。