数列色ぬり
Mier_Samuelle · · 题解
先转化一下题意,发现就是让你找一组不相交的下降子序列和上升子序列,使两序列的长度之和最大。
首先有一个结论,一个序列的 LIS(最长上升子序列)和 LDS(最长下降子序列)最多只有一个交点。这个是好证的,如果有两个交点,那肯定不能同时满足单增和单减。
因此,设 LIS 的长度为
直接 dp 似乎只能做到
可以依据定义简单地理解这个式子,左边就是所有 LIS 和 LDS 的组合数,右边就是相交的 LIS 和 LDS 的组合数,要是这两个相等了,那自然就不存在不相交的 LIS 和 LDS 了。
现在考虑怎么求
:::success[Code]{open}
#include <bits/stdc++.h>
#define lc (u << 1)
#define rc ((u << 1) | 1)
#define fi first
#define se second
#define mid ((l + r) >> 1)
#define int long long
using namespace std;
typedef pair<int, int> pii;
const int MAXN = 1e5 + 10;
const int MOD = 1e9 + 7;
int a[MAXN], toti[MAXN], totd[MAXN], n;
pii li[MAXN], ld[MAXN], ri[MAXN], rd[MAXN];
struct Fenwick_tree{
pii tr[MAXN];
void init(){
for (int i = 0; i <= n; i++){
tr[i] = {0, 0};
}
return;
}
pii merge(pii x, pii y){
if (x.fi == y.fi){
return {x.fi, (x.se + y.se) % MOD};
}
else{
if (x.fi > y.fi){
return x;
}
else{
return y;
}
}
}
int lowbit(int x){
return x & (-x);
}
void modify(int u, pii x){
while (u <= n){
tr[u] = merge(tr[u], x);
u += lowbit(u);
}
return;
}
pii query(int u){
pii res = {0, 0};
while (u){
res = merge(res, tr[u]);
u -= lowbit(u);
}
return res;
}
}tr;
void solve(){
cin >> n;
for (int i = 1; i <= n; i++){
cin >> a[i];
li[i] = ld[i] = ri[i] = rd[i] = {0, 0};
toti[i] = totd[i] = 0;
}
int mxli = 0, mxld = 0;
tr.init();
for (int i = 1; i <= n; i++){
pii res = tr.query(a[i] - 1);
if (res.fi == 0){
li[i] = {1, 1};
}
else{
li[i] = {res.fi + 1, res.se};
}
tr.modify(a[i], li[i]);
mxli = max(mxli, li[i].fi);
}
tr.init();
for (int i = 1; i <= n; i++){
pii res = tr.query(n - a[i]);
if (res.fi == 0){
ld[i] = {1, 1};
}
else{
ld[i] = {res.fi + 1, res.se};
}
tr.modify(n - a[i] + 1, ld[i]);
mxld = max(mxld, ld[i].fi);
}
tr.init();
for (int i = n; i >= 1; i--){
pii res = tr.query(a[i] - 1);
if (res.fi == 0){
ri[i] = {1, 1};
}
else{
ri[i] = {res.fi + 1, res.se};
}
tr.modify(a[i], ri[i]);
}
tr.init();
for (int i = n; i >= 1; i--){
pii res = tr.query(n - a[i]);
if (res.fi == 0){
rd[i] = {1, 1};
}
else{
rd[i] = {res.fi + 1, res.se};
}
tr.modify(n - a[i] + 1, rd[i]);
}
int cnti = 0, cntd = 0;
for (int i = 1; i <= n; i++){
if (li[i].fi == mxli){
cnti = (cnti + li[i].se) % MOD;
}
if (ld[i].fi == mxld){
cntd = (cntd + ld[i].se) % MOD;
}
}
for (int i = 1; i <= n; i++){
if (li[i].fi + rd[i].fi - 1 == mxli){
toti[i] = li[i].se * rd[i].se % MOD;
}
}
for (int i = 1; i <= n; i++){
if (ld[i].fi + ri[i].fi - 1 == mxld){
totd[i] = ld[i].se * ri[i].se % MOD;
}
}
int tot = 0;
for (int i = 1; i <= n; i++){
tot = (tot + toti[i] * totd[i] % MOD) % MOD;
}
if (tot == cnti * cntd % MOD){
cout << mxli + mxld - 1 << "\n";
}
else{
cout << mxli + mxld << "\n";
}
return;
}
signed main(){
ios::sync_with_stdio(false);
cin.tie(0);
cout.tie(0);
int t;
cin >> t;
while (t--){
solve();
}
return 0;
}
:::