「JZOI-2」猜数列 的题解


记住只在没有思路时使用题解,不要从它复制粘贴代码。请尊重题目和题解的作者。
在解题之前提交题解的代码会导致封禁。

作者: admin

概述

本题是交互题:给定一个乱序等差数列的长度 n、取值询问次数上限 q 与模数 p = 10^9+7,要求求出公差 d,两种询问合计不得超过 2n 次。核心解法是用比较询问数出两个已知取值的位置在等差数列中相差多少步 k,再由同余方程 A_2 - A_1 \equiv k d \pmod p 用扩展欧几里得解出 d。

分析
核心观察

第一种询问给出的是两个元素在等差数列中的先后关系,也就是它们的序号(第几项)谁更大,与模 p 之后的数值无关;第二种询问给出的才是模意义下的具体数值。只要知道任意两个元素的取值之差,以及它们在等差数列中的序号之差 k,就能列出同余方程

\displaystyle  A_y - A_x \equiv k \cdot d \pmod p

由于 1 \le k \le n - 1 < p 且 p 是素数,k 与 p 互质,方程有唯一解,用扩展欧几里得求 k 的逆元即可。

思路

先固定位置 1 与位置 2 两个元素,目标是求出它们的序号之差 k。等差数列的 n 个元素序号恰好取遍 0, 1, \ldots, n-1,所以「序号在两者之间」的元素个数正好是 k - 1。

对每个 i \ge 3,分别询问 > 1 i 与 > 2 i。两个回答不同,说明 A_i 与 A_1 的先后关系和它与 A_2 的先后关系相反,也就是 A_i 的序号介于两者之间;回答相同则说明 A_i 在两者同侧。把回答不同的下标个数记为 k',就有 k = k' + 1。

这一步用掉 2(n-2) 次比较询问。接着用两次取值询问拿到 A_1 与 A_2 的模值,再用一次比较询问 > 1 2 判断两者在等差数列中的先后:若返回 1 说明 A_1 的序号更大,把两个取值交换,使 x 对应序号较小者、y 对应序号较大者,于是

\displaystyle  y - x \equiv k \cdot d \pmod p

直接相减可能得到负数(模回绕),故若 x > y 就让 y 加上一个 p,使差值为正且同余关系不变。最后用扩展欧几里得解出 k^{-1},答案即 k^{-1}(y - x) \bmod p。

询问总数为 2(n-2) + 2 + 1 = 2n - 1 \le 2n,其中取值询问恰好 2 次,满足 q \ge 2 的限制。

具体示例

以样例 n = 3、q = 6 为例,数列为 A = [1, 2, 3],其公差 d = 1。

i = 3 时询问 > 1 3 得 0(A_1 < A_3),询问 > 2 3 得 0(A_2 < A_3),两次回答相同,故 k' = 0,k = 1。随后 ? 1 得 1、? 2 得 2,> 1 2 得 0 表示 A_1 序号更小,无需交换,且 x = 1 < y = 2 不必加 p。由 1 \cdot d \equiv 2 - 1 \pmod p 得 d = 1,输出 ! 1。

若把样例换成 A = [3, 1, 2](同一等差数列的另一种排列),> 1 2 会返回 1,交换后 x = 1、y = 3,而 k' 仍为 0、k = 1,同样解得 d = 1——可见答案与排列无关,只取决于两个元素的取值差与序号差。

算法步骤
  1. 读入 n 与 q。
  2. 对 i 从 3 到 n:输出 > 1 i 并读入结果存入 cmp1[i],输出 > 2 i 并读入结果存入 cmp2[i],每行输出后刷新缓冲区。
  3. 统计 cmp1[i] != cmp2[i] 的个数 between,令 k = between + 1。
  4. 询问 ? 1 与 ? 2 得到 x、y,再询问 > 1 2 得到 z。
  5. 若 z == 1 则交换 x 与 y;若 x > y 则令 y += p。
  6. 用扩展欧几里得求 k 在模 p 下的逆元 inv,输出 ! inv * (y - x) % p。
复杂度分析
  • 时间:询问 2n - 1 次,每次 O(1);扩展欧几里得为 O(\log p)。总时间 O(n + \log p)。
  • 空间:O(n)。需要两个长度为 n + 1 的数组记录比较询问的回答(也可以边问边统计,只留两个计数器,空间降到 O(1))。
实现注意事项
  • 每次输出后必须刷新缓冲区(C++ 用 std::endl 或 std::flush,Python 用 flush=True),否则交互双方互相等待,表现为超时。
  • k = k' + 1 中的加一是关键:等差序号在两者之间的元素个数是 k - 1,不是 k。
  • 比较询问的判断方向要与题面一致:返回 1 表示前者在等差数列中更大,即前者序号更大。
  • p 是素数保证了 k 的逆元一定存在;k \le n - 1 < p 也保证了 k \not\equiv 0 \pmod p。
  • y - x 为负时要加一个 p 再取模;直接对负数取模在 C++ 里会得到负值。
  • 取值询问只有两次,恰好用满 q \ge 2 的额度,不要多发。
  • 本题交互密集(n 可达 10^5,询问近 2 \times 10^5 次),Python 的逐行 I/O 开销较大,建议使用 C++ 提交;下面的 Python 版本给出的是同一算法的直译。
源代码
#include <algorithm>
#include <iostream>
#include <vector>

namespace {

const long long P = 1000000007LL;

long long exgcd(long long a, long long b, long long& x, long long& y) {
    if (b == 0) {
        x = 1;
        y = 0;
        return a;
    }
    long long g = exgcd(b, a % b, y, x);
    y -= a / b * x;
    return g;
}

}  // namespace

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

    int n, q;
    std::cin >> n >> q;

    // 对 i >= 3,比较位置 i 的数与位置 1、位置 2 的数的先后关系;
    // 两者答案不同的 i 恰好落在位置 1 与位置 2 之间的那一段里。
    std::vector<int> cmp1(n + 1, 0), cmp2(n + 1, 0);
    for (int i = 3; i <= n; ++i) {
        std::cout << "> 1 " << i << std::endl;
        std::cin >> cmp1[i];
        std::cout << "> 2 " << i << std::endl;
        std::cin >> cmp2[i];
    }

    int between = 0;
    for (int i = 3; i <= n; ++i) {
        if (cmp1[i] != cmp2[i]) {
            ++between;
        }
    }
    long long k = between + 1;  // 位置 1 与位置 2 的数的序号之差

    long long x, y;
    int z;
    std::cout << "? 1" << std::endl;
    std::cin >> x;
    std::cout << "? 2" << std::endl;
    std::cin >> y;
    std::cout << "> 1 2" << std::endl;
    std::cin >> z;
    if (z == 1) {
        std::swap(x, y);  // 令 x 对应序号较小者
    }
    if (x > y) {
        y += P;  // 模回绕:补一个周期,使 y - x 直接等于 k * d
    }

    long long inv, t;
    exgcd(k, P, inv, t);
    inv = (inv % P + P) % P;
    std::cout << "! " << inv * (y - x) % P << std::endl;
    return 0;
}
import sys


def main():
    P = 10 ** 9 + 7
    data = sys.stdin.buffer

    def ask(line):
        sys.stdout.write(line + "\n")
        sys.stdout.flush()
        return data.readline().split()

    n, q = map(int, data.readline().split())

    same_side = 0
    for i in range(3, n + 1):
        c1 = int(ask("> 1 %d" % i)[0])
        c2 = int(ask("> 2 %d" % i)[0])
        if c1 != c2:
            same_side += 1
    k = same_side + 1  # 位置 1 与位置 2 的数的序号之差

    x = int(ask("? 1")[0])
    y = int(ask("? 2")[0])
    z = int(ask("> 1 2")[0])
    if z == 1:
        x, y = y, x
    if x > y:
        y += P  # 模回绕:补一个周期

    inv = pow(k, P - 2, P)  # p 为素数,用费马小定理求逆元
    sys.stdout.write("! %d\n" % (inv * (y - x) % P))
    sys.stdout.flush()


main()

评论

目前没有评论。