1 条题解

  • 0
    @ 2026-1-28 9:30:03

    题目大意

    给两个长为 nn 的正整数序列 a,ca,c,以及一个长为 nn 的整数序列 bb

    定义 f(a)=ai=0bif(a)=\sum_{a_i=0}b_ig(a)=ai=0cig(a)=\prod_{a_i=0}c_i

    你可以对 aa 执行任意次以下操作:

    • 选择两个相邻的位置 i,ji,j,若 aiaja_i \leq a_j,则将 aja_j 改为 ajaia_j - a_i,同时将 aia_i 改为 00

    对于所有可能经过 00 次或更多次操作得到的 aa,求出 maxf(a)\max f(a)g(a)mod109+7\sum g(a)\mod 10^9+7

    解题思路

    一些定义与发现

    特别的,假若上述对 aa 的操作满足 i<ji<j,我们称为往右操作;反之 i>ji>j,我们称为往左操作。同时,定义 ai>aja_i>a_j 的情况为无法操作

    经过一次操作后,aia_iaja_{j} 中小的那个会变成 00。同时,对任意一组 min(ai,aj)=0\min(a_i,a_{j})=0 操作是无效的。所以,操作一次后两边独立

    我们更进一步地发现,最终序列每一个非 00 位置都对应了初始 aa 的一个区间。这个区间的操作形式为:从两边不断往内操作。

    第一问

    先做第一问。记 fif_i 表示,考虑了前 ii 个数的最大权值是多少。每次枚举一个 jj 转移过去。转移的时候还要枚举一个 kk 表示最终缩在哪个数上了(记得特判最最终整个区间都是 00 的情况),

    具体的,我们暴力 check 最终位置 kk 的合法性容易做到暴力 O(n)O(n)。 优化的话考虑预处理出 pli,pripl_i,pr_i 表示,每个位置作为左/右端点往另一端操作最远能到哪里。此时判断条件就容易写成 max(prj,i+1)kmin(pli+1,j)\max(pr_{j},i+1)\leq k\leq\min(pl_{i+1},j),时间复杂度 O(n3)O(n^3)

    事实上,我们不需要枚举 kk,只需要知道最优的 kk 对应的贡献即可。 由于我们发现每个 aia_i 对最终的 aka_k 的贡献只取决于 ii 的奇偶性,故而可以分类讨论。

    s0s0 表示 [l,r][l,r] 内所有偶数位置的 ai\sum a_i,记 s1s1 表示所有奇数为的 ai\sum a_i

    • s0=s1s0=s1,区间最后会变为全 00
    • s0>s1s0>s1kk 是上述合法区间 [max(prj,i+1),min(pli+1,j)][\max(pr_{j},i+1),\min(pl_{i+1},j)] 内的任意偶数位置
    • s0<s1s0<s1kk 是上述合法区间 [max(prj,i+1),min(pli+1,j)][\max(pr_{j},i+1),\min(pl_{i+1},j)] 内的任意奇数位置。

    预处理每个区间内的奇/偶位置最值,可以 O(1)O(1) 完成上面的转移。至此,我们 O(n2)O(n^2) 完成了第一问。

    第二问

    第二问同理。特别的,为了防止出现两段被拼起来重复计数(其中有一段得是最终全 00),我们将对 pl,prpl,pr 的定义进行一点修改:

    • plipl_i:从 ii 往右操作,第一次满足“操作后 aj=0a_j=0”的位置前停下。
    • pripr_i:从 ii 往左操作,第一次操作后 aj=0a_j=0 停下。

    这样可以不重不漏计数。最终这题总时间复杂度 O(n2)O(n^2),空间复杂度 O(n2)O(n^2)。如果使用 ST 表维护区间最值,空间上可以做到 O(nlogn)O(n\log n)

    代码实现

    #include <bits/stdc++.h>
    #define FL(i, a, b) for (int i = (a); i <= (b); ++i)
    #define FR(i, a, b) for (int i = (a); i >= (b); --i)
    using namespace std;
    typedef long long ll;
    const int N = 5e3 + 10;
    const int MOD = 1e9 + 7;
    const ll INFLL = 1e18;
    int n, a[N], b[N], c[N], ic[N];
    int pl[N], pr[N], v[N], iv[N];
    ll mi[2][N][N], ct[2][N][N];
    ll s[N], t[N], f[N], g[N];
    int Qpow(int a, int b) {
    	int ret = 1;
    	for (; b; b >>= 1) {
    		if (b & 1)
    			ret = (ll)ret * a % MOD;
    		a = (ll)a * a % MOD;
    	}
    	return ret;
    }
    int Inv(int x) {
    	return Qpow(x, MOD - 2);  
    }
    void Solve() {
    	scanf("%d", &n);
    	FL(i, 1, n) {
    		scanf("%d", &a[i]);
    		t[i] = t[i - 1] + (i & 1? -a[i] : a[i]);
    	}
    	FL(i, 1, n) {
    		scanf("%d", &b[i]);
    		s[i] = s[i - 1] + b[i];
    	}
    	v[0] = iv[0] = 1;
    	FL(i, 1, n) {
    		scanf("%d", &c[i]);
    		v[i] = (ll)v[i - 1] * c[i] % MOD;
    		ic[i] = Inv(c[i]);
    		iv[i] = Inv(v[i]);
    	}
    	FL(i, 1, n) {
    		mi[0][i][i - 1] = mi[1][i][i - 1] = INFLL;
    		ct[0][i][i - 1] = ct[1][i][i - 1] = 0;
    		FL(j, i, n) {
    			FL(k, 0, 1) {
    				mi[k][i][j] = mi[k][i][j - 1];
    				ct[k][i][j] = ct[k][i][j - 1];
    			}
    			mi[j & 1][i][j] = min(mi[j & 1][i][j], (ll)b[j]);
    			ct[j & 1][i][j] = (ct[j & 1][i][j] + ic[j]) % MOD;
    		}
    	}
    	FL(i, 1, n) {
    		int d = a[i];
    		FL(j, i, n) {
    			d = a[j + 1] - d;
    			if (j == n || d <= 0) {
    				pl[i] = j;
    				break;
    			}
    		}
    		d = a[i];
    		FR(j, i, 1) {
    			d = a[j - 1] - d;
    			if (j == 1 || d <= 0) {
    				pr[i] = (j > 1 && !d? j - 1 : j);
    				break;
    			}
    		}
    	}
    	fill(f, f + n + 1, -INFLL);
    	fill(g, g + n + 1, 0);
    	f[0] = 0, g[0] = 1;
    	FL(i, 0, n) {
    		FL(j, i + 1, n) {
    			int l = max(i + 1, pr[j]), r = min(j, pl[i + 1]);
    			if (l > r){
    				continue;
    			}
    
    			int h = (ll)g[i] * v[j] % MOD * iv[i] % MOD;
    			if (t[j] - t[i] == 0) {
    				f[j] = max(f[j], f[i] + (s[j] - s[i]));
    				g[j] = (g[j] + h) % MOD;
    			} else if (t[j] - t[i] > 0) {
    				f[j] = max(f[j], f[i] + (s[j] - s[i]) - mi[0][l][r]);
    				g[j] = (g[j] + (ll)h * ct[0][l][r]) % MOD;
    			} else {
    				f[j] = max(f[j], f[i] + (s[j] - s[i]) - mi[1][l][r]);
    				g[j] = (g[j] + (ll)h * ct[1][l][r]) % MOD;
    			}
    		}
    	}
    	printf("%lld %lld\n", f[n], g[n]);
    }
    int main() {
    	int T;
    	scanf("%*d %d", &T);
    	while (T--) {
    		Solve();
    	}
    	return 0;
    }
    
    • 1

    信息

    ID
    2386
    时间
    6000ms
    内存
    1024MiB
    难度
    10
    标签
    递交数
    7
    已通过
    1
    上传者