[NOIP 2005 提高组] 等价表达式

GitHub跳转原题关系图返回列表

先按题目给定优先级把表达式递归下降解析成语法树,再用多组模数代值比较签名,快速判断哪些选项与题干恒等。

OJ: luogu

题目 ID: P1054

难度:提高+/省选-

标签:字符串递归数学思维

日期: 2026-06-20 19:26

题意

给出一个只含变量 a 的代数表达式作为题干。

再给出 n 个选项表达式,要求找出哪些选项与题干表达式恒等。

如果某个选项本身是非法表达式,就直接忽略它。

输出所有等价选项对应的字母,按顺序拼起来,不加空格。

思路

先看一个可以直接验证想法的朴素解:

cpp
#include <bits/stdc++.h>
using namespace std;

// brute.cpp:小数据精确解。
// 先把表达式解析成语法树,再把它转成关于 a 的多项式。
// 这个做法非常直观,也适合帮助读者理解“等价表达式”的本质,
// 但当表达式次数很大时会变慢,所以只适合作为小数据验证代码。

const int NODE_CONST = 1;
const int NODE_VAR = 2;
const int NODE_ADD = 3;
const int NODE_SUB = 4;
const int NODE_MUL = 5;
const int NODE_POW = 6;

const long long LIMIT = 1000000;
typedef __int128 i128;

struct Node {
    int type;
    long long value;
    int power;
    int left;
    int right;
};

struct Poly {
    vector<i128> c; // c[i] 表示 a^i 的系数
};

struct Parser {
    string s;
    int n;
    int pos;
    bool ok;
    vector<Node> nodes;

    void init(const string &str) {
        s = str;
        n = (int)s.size();
        pos = 0;
        ok = true;
        nodes.clear();
        nodes.push_back(Node());
    }

    void skip_space() {
        while (pos < n && s[pos] == ' ') {
            pos++;
        }
    }

    int new_node(int type, long long value, int power, int left, int right) {
        Node cur;
        cur.type = type;
        cur.value = value;
        cur.power = power;
        cur.left = left;
        cur.right = right;
        nodes.push_back(cur);
        return (int)nodes.size() - 1;
    }

    long long clamp_value(__int128 x) {
        if (x > LIMIT) {
            return LIMIT;
        }
        if (x < -LIMIT) {
            return -LIMIT;
        }
        return (long long)x;
    }

    long long safe_add(long long a, long long b) {
        return clamp_value((__int128)a + (__int128)b);
    }

    long long safe_sub(long long a, long long b) {
        return clamp_value((__int128)a - (__int128)b);
    }

    long long safe_mul(long long a, long long b) {
        return clamp_value((__int128)a * (__int128)b);
    }

    long long safe_pow(long long a, int b) {
        long long ans = 1;
        for (int i = 0; i < b; i++) {
            ans = safe_mul(ans, a);
        }
        return ans;
    }

    bool eval_constant(int u, long long &ret) {
        Node &cur = nodes[u];
        if (cur.type == NODE_CONST) {
            ret = cur.value;
            return true;
        }
        if (cur.type == NODE_VAR) {
            return false;
        }

        long long left_value = 0;
        long long right_value = 0;

        if (!eval_constant(cur.left, left_value)) {
            return false;
        }

        if (cur.type == NODE_POW) {
            ret = safe_pow(left_value, cur.power);
            return true;
        }

        if (!eval_constant(cur.right, right_value)) {
            return false;
        }

        if (cur.type == NODE_ADD) {
            ret = safe_add(left_value, right_value);
            return true;
        }
        if (cur.type == NODE_SUB) {
            ret = safe_sub(left_value, right_value);
            return true;
        }
        if (cur.type == NODE_MUL) {
            ret = safe_mul(left_value, right_value);
            return true;
        }

        return false;
    }

    int parse_expr() {
        int left = parse_term();
        if (!ok) {
            return 0;
        }

        while (true) {
            skip_space();
            if (pos >= n || (s[pos] != '+' && s[pos] != '-')) {
                break;
            }

            char op = s[pos];
            pos++;
            int right = parse_term();
            if (!ok || right == 0) {
                ok = false;
                return 0;
            }

            if (op == '+') {
                left = new_node(NODE_ADD, 0, 0, left, right);
            } else {
                left = new_node(NODE_SUB, 0, 0, left, right);
            }
        }

        return left;
    }

    int parse_term() {
        int left = parse_power();
        if (!ok) {
            return 0;
        }

        while (true) {
            skip_space();
            if (pos >= n || s[pos] != '*') {
                break;
            }

            pos++;
            int right = parse_power();
            if (!ok || right == 0) {
                ok = false;
                return 0;
            }

            left = new_node(NODE_MUL, 0, 0, left, right);
        }

        return left;
    }

    int parse_power() {
        int left = parse_primary();
        if (!ok) {
            return 0;
        }

        while (true) {
            skip_space();
            if (pos >= n || s[pos] != '^') {
                break;
            }

            pos++;
            int right = parse_primary();
            if (!ok || right == 0) {
                ok = false;
                return 0;
            }

            long long exponent = 0;
            if (!eval_constant(right, exponent)) {
                ok = false;
                return 0;
            }
            if (exponent < 1 || exponent > 10) {
                ok = false;
                return 0;
            }

            left = new_node(NODE_POW, 0, (int)exponent, left, 0);
        }

        return left;
    }

    int parse_number() {
        skip_space();
        if (pos >= n || !isdigit(s[pos])) {
            ok = false;
            return 0;
        }

        long long value = 0;
        while (pos < n && isdigit(s[pos])) {
            value = value * 10 + (s[pos] - '0');
            pos++;
        }

        return new_node(NODE_CONST, value, 0, 0, 0);
    }

    int parse_primary() {
        skip_space();
        if (pos >= n) {
            ok = false;
            return 0;
        }

        if (s[pos] == 'a') {
            pos++;
            return new_node(NODE_VAR, 0, 0, 0, 0);
        }
        if (isdigit(s[pos])) {
            return parse_number();
        }
        if (s[pos] == '(') {
            pos++;
            int inside = parse_expr();
            skip_space();
            if (!ok || pos >= n || s[pos] != ')') {
                ok = false;
                return 0;
            }
            pos++;
            return inside;
        }

        ok = false;
        return 0;
    }

    bool build_tree(const string &str, int &root) {
        init(str);
        root = parse_expr();
        skip_space();
        if (!ok || root == 0 || pos != n) {
            return false;
        }
        return true;
    }
};

void trim_poly(Poly &p) {
    while (!p.c.empty() && p.c.back() == 0) {
        p.c.pop_back();
    }
}

Poly make_const(long long x) {
    Poly p;
    p.c.resize(1);
    p.c[0] = (i128)x;
    trim_poly(p);
    return p;
}

Poly make_var() {
    Poly p;
    p.c.resize(2);
    p.c[0] = 0;
    p.c[1] = 1;
    trim_poly(p);
    return p;
}

Poly poly_add(const Poly &a, const Poly &b) {
    Poly c;
    int size_a = (int)a.c.size();
    int size_b = (int)b.c.size();
    int size_c = max(size_a, size_b);

    c.c.assign(size_c, 0);
    for (int i = 0; i < size_a; i++) {
        c.c[i] += a.c[i];
    }
    for (int i = 0; i < size_b; i++) {
        c.c[i] += b.c[i];
    }

    trim_poly(c);
    return c;
}

Poly poly_sub(const Poly &a, const Poly &b) {
    Poly c;
    int size_a = (int)a.c.size();
    int size_b = (int)b.c.size();
    int size_c = max(size_a, size_b);

    c.c.assign(size_c, 0);
    for (int i = 0; i < size_a; i++) {
        c.c[i] += a.c[i];
    }
    for (int i = 0; i < size_b; i++) {
        c.c[i] -= b.c[i];
    }

    trim_poly(c);
    return c;
}

Poly poly_mul(const Poly &a, const Poly &b) {
    Poly c;
    if (a.c.empty() || b.c.empty()) {
        return c;
    }

    int deg_a = (int)a.c.size() - 1;
    int deg_b = (int)b.c.size() - 1;
    c.c.assign(deg_a + deg_b + 1, 0);

    for (int i = 0; i <= deg_a; i++) {
        for (int j = 0; j <= deg_b; j++) {
            c.c[i + j] += a.c[i] * b.c[j];
        }
    }

    trim_poly(c);
    return c;
}

Poly poly_pow(Poly base, int exponent) {
    Poly ans = make_const(1);
    while (exponent > 0) {
        if (exponent & 1) {
            ans = poly_mul(ans, base);
        }
        exponent >>= 1;
        if (exponent > 0) {
            base = poly_mul(base, base);
        }
    }
    trim_poly(ans);
    return ans;
}

Poly build_poly(const vector<Node> &nodes, int u) {
    const Node &cur = nodes[u];

    if (cur.type == NODE_CONST) {
        return make_const(cur.value);
    }
    if (cur.type == NODE_VAR) {
        return make_var();
    }
    if (cur.type == NODE_ADD) {
        return poly_add(build_poly(nodes, cur.left), build_poly(nodes, cur.right));
    }
    if (cur.type == NODE_SUB) {
        return poly_sub(build_poly(nodes, cur.left), build_poly(nodes, cur.right));
    }
    if (cur.type == NODE_MUL) {
        return poly_mul(build_poly(nodes, cur.left), build_poly(nodes, cur.right));
    }

    return poly_pow(build_poly(nodes, cur.left), cur.power);
}

bool same_poly(const Poly &a, const Poly &b) {
    int size_a = (int)a.c.size();
    int size_b = (int)b.c.size();
    int size_c = max(size_a, size_b);

    for (int i = 0; i < size_c; i++) {
        i128 x = 0;
        i128 y = 0;
        if (i < size_a) {
            x = a.c[i];
        }
        if (i < size_b) {
            y = b.c[i];
        }
        if (x != y) {
            return false;
        }
    }
    return true;
}

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

    string target_expr;
    getline(cin, target_expr);

    int n;
    cin >> n;
    string useless;
    getline(cin, useless);

    Parser parser;
    int target_root = 0;
    parser.build_tree(target_expr, target_root);
    Poly target_poly = build_poly(parser.nodes, target_root);

    string answer = "";
    for (int i = 0; i < n; i++) {
        string expr;
        getline(cin, expr);

        int root = 0;
        if (!parser.build_tree(expr, root)) {
            continue;
        }

        Poly cur_poly = build_poly(parser.nodes, root);
        if (same_poly(target_poly, cur_poly)) {
            answer.push_back(char('A' + i));
        }
    }

    cout << answer << '\n';

    return 0;
}

brute.cpp 的思路很直接:

  1. 先把表达式解析成语法树;
  2. 再把整棵树真的转成关于 a 的多项式;
  3. 最后精确比较多项式系数。

这个办法非常适合帮助理解题意,但正式数据里,表达式次数可能很大,真的展开会很麻烦。

所以正式解不去完整展开,而是先抓住两件事。

第一件事,是必须正确解析表达式。

这题的优先级和普通写法有一个容易错的点:

  • ^ 不但优先级高于 *
  • 而且它也是左结合

例如:

text
1^10^9

在这题里要按:

text
(1^10)^9

来理解。

因此正式解先用递归下降,把表达式按四层优先级解析成语法树:

  • expr 处理 + -
  • term 处理 *
  • power 处理 ^
  • primary 处理常数、变量、括号

第二件事,是没必要真的把多项式展开完。

因为指数始终是常数,所有合法表达式本质上都是关于 a 的多项式。 如果两个表达式不等价,那它们对应的多项式就不同。

对于不同的多项式,只要在若干个点上代值比较,极大概率就能把它们区分开。

所以做法就是:

  1. 解析题干表达式;
  2. 选若干组 (mod, a)
  3. 在这些模数下求出题干签名;
  4. 对每个选项做同样的求值;
  5. 所有签名都相同,就认为它和题干等价。

另外,题目要求忽略非法表达式,所以解析阶段还要顺手做合法性检查。

尤其是 ^ 的右边,必须是一个不含变量的常量表达式,而且值在 1..10 之间。 如果不满足,就直接把这个选项判成非法并跳过。

代码

cpp
#include <bits/stdc++.h>
using namespace std;

const int NODE_CONST = 1;
const int NODE_VAR = 2;
const int NODE_ADD = 3;
const int NODE_SUB = 4;
const int NODE_MUL = 5;
const int NODE_POW = 6;

const long long LIMIT = 1000000;
const int TEST_CNT = 4;
const long long MODS[TEST_CNT] = {1000000007LL, 1000000009LL, 998244353LL, 1004535809LL};
const long long VALUES[TEST_CNT] = {911382323LL, 972663749LL, 911382323LL, 97266353LL};

struct Node {
    int type;
    long long value; // 常数节点的值
    int power;       // 乘方节点的指数,保证是 1..10
    int left;
    int right;
};

struct Parser {
    string s;
    int n;
    int pos;
    bool ok;
    vector<Node> nodes;

    Parser() {
        n = 0;
        pos = 0;
        ok = true;
    }

    void init(const string &str) {
        s = str;
        n = (int)s.size();
        pos = 0;
        ok = true;
        nodes.clear();
        nodes.push_back(Node());
    }

    void skip_space() {
        while (pos < n && s[pos] == ' ') {
            pos++;
        }
    }

    int new_node(int type, long long value, int power, int left, int right) {
        Node cur;
        cur.type = type;
        cur.value = value;
        cur.power = power;
        cur.left = left;
        cur.right = right;
        nodes.push_back(cur);
        return (int)nodes.size() - 1;
    }

    long long clamp_value(__int128 x) {
        if (x > LIMIT) {
            return LIMIT;
        }
        if (x < -LIMIT) {
            return -LIMIT;
        }
        return (long long)x;
    }

    long long safe_add(long long a, long long b) {
        return clamp_value((__int128)a + (__int128)b);
    }

    long long safe_sub(long long a, long long b) {
        return clamp_value((__int128)a - (__int128)b);
    }

    long long safe_mul(long long a, long long b) {
        return clamp_value((__int128)a * (__int128)b);
    }

    long long safe_pow(long long a, int b) {
        long long ans = 1;
        for (int i = 0; i < b; i++) {
            ans = safe_mul(ans, a);
        }
        return ans;
    }

    // 判断一棵子树是否是不含变量的常量表达式。
    bool eval_constant(int u, long long &ret) {
        Node &cur = nodes[u];

        if (cur.type == NODE_CONST) {
            ret = cur.value;
            return true;
        }
        if (cur.type == NODE_VAR) {
            return false;
        }

        long long left_value = 0;
        long long right_value = 0;

        if (!eval_constant(cur.left, left_value)) {
            return false;
        }

        if (cur.type == NODE_POW) {
            ret = safe_pow(left_value, cur.power);
            return true;
        }

        if (!eval_constant(cur.right, right_value)) {
            return false;
        }

        if (cur.type == NODE_ADD) {
            ret = safe_add(left_value, right_value);
            return true;
        }
        if (cur.type == NODE_SUB) {
            ret = safe_sub(left_value, right_value);
            return true;
        }
        if (cur.type == NODE_MUL) {
            ret = safe_mul(left_value, right_value);
            return true;
        }

        return false;
    }

    int parse_expr() {
        int left = parse_term();
        if (!ok) {
            return 0;
        }

        while (true) {
            skip_space();
            if (pos >= n) {
                break;
            }

            if (s[pos] != '+' && s[pos] != '-') {
                break;
            }

            char op = s[pos];
            pos++;
            int right = parse_term();
            if (!ok || right == 0) {
                ok = false;
                return 0;
            }

            if (op == '+') {
                left = new_node(NODE_ADD, 0, 0, left, right);
            } else {
                left = new_node(NODE_SUB, 0, 0, left, right);
            }
        }

        return left;
    }

    int parse_term() {
        int left = parse_power();
        if (!ok) {
            return 0;
        }

        while (true) {
            skip_space();
            if (pos >= n || s[pos] != '*') {
                break;
            }

            pos++;
            int right = parse_power();
            if (!ok || right == 0) {
                ok = false;
                return 0;
            }

            left = new_node(NODE_MUL, 0, 0, left, right);
        }

        return left;
    }

    int parse_power() {
        int left = parse_primary();
        if (!ok) {
            return 0;
        }

        while (true) {
            skip_space();
            if (pos >= n || s[pos] != '^') {
                break;
            }

            pos++;
            int right = parse_primary();
            if (!ok || right == 0) {
                ok = false;
                return 0;
            }

            long long exponent = 0;
            if (!eval_constant(right, exponent)) {
                ok = false;
                return 0;
            }
            if (exponent < 1 || exponent > 10) {
                ok = false;
                return 0;
            }

            left = new_node(NODE_POW, 0, (int)exponent, left, 0);
        }

        return left;
    }

    int parse_number() {
        skip_space();
        if (pos >= n || !isdigit(s[pos])) {
            ok = false;
            return 0;
        }

        long long value = 0;
        while (pos < n && isdigit(s[pos])) {
            value = value * 10 + (s[pos] - '0');
            pos++;
        }

        return new_node(NODE_CONST, value, 0, 0, 0);
    }

    int parse_primary() {
        skip_space();
        if (pos >= n) {
            ok = false;
            return 0;
        }

        if (s[pos] == 'a') {
            pos++;
            return new_node(NODE_VAR, 0, 0, 0, 0);
        }

        if (isdigit(s[pos])) {
            return parse_number();
        }

        if (s[pos] == '(') {
            pos++;
            int inside = parse_expr();
            skip_space();
            if (!ok || pos >= n || s[pos] != ')') {
                ok = false;
                return 0;
            }
            pos++;
            return inside;
        }

        ok = false;
        return 0;
    }

    bool build_tree(const string &str, int &root) {
        init(str);
        root = parse_expr();
        skip_space();
        if (!ok || root == 0 || pos != n) {
            return false;
        }
        return true;
    }
};

long long mod_pow(long long a, int b, long long mod) {
    long long ans = 1 % mod;
    long long base = a % mod;

    while (b > 0) {
        if (b & 1) {
            ans = (__int128)ans * base % mod;
        }
        base = (__int128)base * base % mod;
        b >>= 1;
    }

    return ans;
}

long long norm_mod(long long x, long long mod) {
    x %= mod;
    if (x < 0) {
        x += mod;
    }
    return x;
}

long long eval_mod(const vector<Node> &nodes, int u, long long a_value, long long mod) {
    const Node &cur = nodes[u];

    if (cur.type == NODE_CONST) {
        return norm_mod(cur.value, mod);
    }
    if (cur.type == NODE_VAR) {
        return a_value % mod;
    }
    if (cur.type == NODE_ADD) {
        long long left_value = eval_mod(nodes, cur.left, a_value, mod);
        long long right_value = eval_mod(nodes, cur.right, a_value, mod);
        return (left_value + right_value) % mod;
    }
    if (cur.type == NODE_SUB) {
        long long left_value = eval_mod(nodes, cur.left, a_value, mod);
        long long right_value = eval_mod(nodes, cur.right, a_value, mod);
        return norm_mod(left_value - right_value, mod);
    }
    if (cur.type == NODE_MUL) {
        long long left_value = eval_mod(nodes, cur.left, a_value, mod);
        long long right_value = eval_mod(nodes, cur.right, a_value, mod);
        return (__int128)left_value * right_value % mod;
    }

    long long left_value = eval_mod(nodes, cur.left, a_value, mod);
    return mod_pow(left_value, cur.power, mod);
}

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

    string target_expr;
    getline(cin, target_expr);

    int n;
    cin >> n;
    string useless;
    getline(cin, useless);

    Parser parser;
    int target_root = 0;

    parser.build_tree(target_expr, target_root);

    vector<long long> target_signature(TEST_CNT, 0);
    for (int i = 0; i < TEST_CNT; i++) {
        target_signature[i] = eval_mod(parser.nodes, target_root, VALUES[i], MODS[i]);
    }

    string answer = "";
    for (int i = 0; i < n; i++) {
        string expr;
        getline(cin, expr);

        int root = 0;
        if (!parser.build_tree(expr, root)) {
            continue;
        }

        bool same = true;
        for (int j = 0; j < TEST_CNT; j++) {
            long long cur_value = eval_mod(parser.nodes, root, VALUES[j], MODS[j]);
            if (cur_value != target_signature[j]) {
                same = false;
                break;
            }
        }

        if (same) {
            answer.push_back(char('A' + i));
        }
    }

    cout << answer << '\n';

    return 0;
}

复杂度

设单个表达式长度为 L

解析一次是 O(L)O(L),每组模数代值也是 O(L)O(L)。 由于模数组数是常数,所以总时间复杂度是 O(nL)O(nL),空间复杂度是 O(L)O(L)

总结

这题有两个关键点:

  1. 先把题目的运算规则,尤其是左结合的 ^,解析正确;
  2. 再把“恒等比较”转成“多组代值签名比较”。

这样就不需要真正展开高次多项式,也能比较稳地完成判等。