先按题目给定优先级把表达式递归下降解析成语法树,再用多组模数代值比较签名,快速判断哪些选项与题干恒等。
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 的思路很直接:
- 先把表达式解析成语法树;
- 再把整棵树真的转成关于
a的多项式; - 最后精确比较多项式系数。
这个办法非常适合帮助理解题意,但正式数据里,表达式次数可能很大,真的展开会很麻烦。
所以正式解不去完整展开,而是先抓住两件事。
第一件事,是必须正确解析表达式。
这题的优先级和普通写法有一个容易错的点:
^不但优先级高于*- 而且它也是左结合
例如:
text
1^10^9在这题里要按:
text
(1^10)^9来理解。
因此正式解先用递归下降,把表达式按四层优先级解析成语法树:
expr处理+ -term处理*power处理^primary处理常数、变量、括号
第二件事,是没必要真的把多项式展开完。
因为指数始终是常数,所有合法表达式本质上都是关于 a 的多项式。
如果两个表达式不等价,那它们对应的多项式就不同。
对于不同的多项式,只要在若干个点上代值比较,极大概率就能把它们区分开。
所以做法就是:
- 解析题干表达式;
- 选若干组
(mod, a); - 在这些模数下求出题干签名;
- 对每个选项做同样的求值;
- 所有签名都相同,就认为它和题干等价。
另外,题目要求忽略非法表达式,所以解析阶段还要顺手做合法性检查。
尤其是 ^ 的右边,必须是一个不含变量的常量表达式,而且值在 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。
解析一次是
总结
这题有两个关键点:
- 先把题目的运算规则,尤其是左结合的
^,解析正确; - 再把“恒等比较”转成“多组代值签名比较”。
这样就不需要真正展开高次多项式,也能比较稳地完成判等。