ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

反射计数问题如何用周期与中国剩余定理替代暴力模拟

反射计数问题如何用周期与中国剩余定理替代暴力模拟 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 基本必挂。第二是反射的细节。我第一次写模拟时用的是“先移动再判断越界越界后反向并往回走”结果在边界上的行为是错的。比如点位于 x0方向 dx-1正确行为应该是弹回 x1但错误的写法会让它停在原地或者位置错乱。这种问题普通样例根本测不出来一上特殊数据就暴露。所以网上虽然很多人贴过这题的模拟写法但真到评审机上跑很多版本过不了大数据。更稳妥的路线是先识别出路径的周期再用数论工具直接算出结果。下面的数学模型就是这个思路。1.3 从一个例子手算出周期规律只看 X 方向。网格有 n 列坐标范围 0 到 n-1。点在这个区间里来回跑跑一个完整来回走过的距离是 2(n-1)所以 X 方向坐标序列的周期是Px 2 * (n - 1)同理Y 方向的周期是Py 2 * (m - 1)在一个周期内X 方向上某个内部坐标比如 x3会被经过两次一次向右路过一次向左路过而边界坐标 0 和 n-1 只会被经过一次。这个“内部两个解、边界一个解”的差异正是后面代码里最容易漏掉的地方。拿样例手算n5m5起点 (2,2)方向 (1,1)目标点 (3,3)。X 方向周期 Px8x(t) 恰好等于 3 的时间满足t ≡ 1 (mod 8) 或者 t ≡ 3 (mod 8)Y 方向完全一样。两个条件都满足的时间在 [0, 10] 内只有 t1、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-1Sx {tx, Px - tx}tx 是边界点即 tx 0 或 tx n-1Sx {tx}因为 Px - tx 和 tx 在模 Px 意义下是同一个数于是得到同余方程dx * t ≡ s - x0 (mod Px)其中 s ∈ Sx因为 dx 只能是 1 或 -1dx 的逆元就是它自己所以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 一维退化和角点反射的特殊处理如果 n1 或者 m1周期公式里会出现 Px0 或 Py0没法套 CRT。这种情况必须单独分支n1 且 m1网格只有一个点。目标点若是 (0,0)答案是 T1否则是 0。n1X 方向永远停在 0。目标 tx 不为 0直接输出 0tx 为 0 时问题降维成 Y 方向上的一维往返只做一个方向的余数类计数。m1对称处理。角点反射不需要额外特判。在 (0,0) 且方向为 (-1,-1) 的瞬间两个方向同时越界dx、dy 同时取反点从角落弹向 (1,1)。这在模拟代码里天然成立在同余公式里也天然成立因为周期折叠函数把 0 和 Px 看成同一位置展开平面上角点就是网格镜像的拼接处。3. 四种语言落地Java、Python、Node.js、C3.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 ListLong genResidues(long pos, long dir, long period, long target, long limit) { ListLong 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); ListLong xs genResidues(x, dx, px, tx, n - 1); ListLong 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 到 6T 在 0 到 20用 Python 的周期法逐组对比结果完全一致。下面挑几个有代表性的用例说明。4.2 几组典型测试用例用例n m起点方向目标T模拟结果周期法结果基础往返5 52 21 13 31033角落弹射3 30 0-1 -11 11055一维退化1 50 21 10 21033不同周期3 41 21 -12 1511第一组就是样例。第二组验证角点反弹n3, m3起点 (0,0)方向 (-1,-1)目标 (1,1)T10。小球在角点反弹后反复经过中心点t1、3、5、7、9 共 5 次周期法算出 t ≡ 1 或 3 (mod 4)计数 5正确。第三组验证 n1 退化。网格只有一列五格小球在竖线上往返目标 (0,2) 在 t0、4、8 被经过答案 3。如果不处理 n1周期公式里会出现除零直接崩。第四组验证 Px4、Py6 这类不同周期场景。我手算过满足条件的时间只有一个 t1所以答案 1。CRT 正确合并了不同模的同余方程。4.3 最容易翻车的几个点第一个坑是边界点被当成内部点处理。tx0 或 txn-1 只能生成一个余数类如果按内部点生成了两个结果会多算。这个 bug 很隐蔽因为只要目标点落在边界且恰好某条伪解在 T 内出现答案就会偏大。代码里我用limit n - 1作为边界最大值再用0 target limit判断内部点就是为了从根上避开这个错误。第二个坑是 CRT 无解。Px 和 Py 不互质时会有一些余数组合找不出整数解必须跳过。漏掉这个判断的话C 和 Java 里可能算出错误的余数甚至因为模逆元不存在而崩溃。Px8、Py6 时x 方向余数 1 和 y 方向余数 3 就是无解组合因为两个同余条件自相矛盾。第三个坑是负余数。Java、C、JS 的取模结果可能是负数norm 函数必须加。我在四种语言里特意都写了 norm唯一不用改的是 Python。跨语言对拍时同样的逻辑在 Java 和 C 上跑错、在 Python 上跑对基本就是负数取模的问题。第四个坑是乘法溢出。C 和 Java 的 long 在极端数据下有风险。C 我已经用 __int128 兜底Java 如果怕溢出最快的改法是把 CRT 里的乘法换成 BigInteger或者限制一下输入范围。实际机试数据通常不会顶着上限出但心里要清楚这个边界。第五个坑是没有处理 n1 或 m1。如果不特判周期变量变成 0后面不是除零就是无限循环。这类“一维退化”的用例在评测机里几乎一定会出现属于必拿分的点。4.4 一个可以延伸思考的视角用展开平面看这道题会更直观反射等于把网格镜像展开光点永远走直线。目标点在展开平面上有无数个镜像问题变成“直线在某时刻是否落到某个镜像点上”。X 方向上目标镜像点的横坐标是 tx 2k(n-1) 和 -tx 2k(n-1)解同余方程得到的正好就是这组点。这也是为什么内部点有两个余数类、边界点只有一个——边界点在镜像展开后是自我重合的。如果遇到变体题比如“求反射次数而不是经过次数”或者“求首次到达某个点的时间”底层仍然是同一套周期模型。反射次数对应方向翻转次数首次到达时间对应同余方程的最小非负解改改判断条件就行。机器上对拍通过后我把四种语言版本都各自提交了一遍Java 和 C 用时几乎为 0Python 和 Node 也都在毫秒级。相比模拟法的 O(T)周期法的复杂度只有 O(log min(Px, Py)) 级别的常数差别是量级上的。最后说点实际的。我现在做这题会先把 n1、m1、目标点在边界这三种 corner case 写在草稿纸最上方然后才开始写代码。genResidues 和 crt 两个函数抽出来四个语言逻辑完全一样换语言只换语法不换结构。如果平台数据范围确实很小模拟能过但学会周期法之后碰到 T 特别大的“反射计数 Plus”就不会慌了。
返回列表