最大子段和-cnblog
最大子段和
线性序列最大子段和
线性序列上我们要维护最大子段和,只需要维护这4个值:
- \(pre\):这个节点的最大前缀和;
- \(suf\):这个节点的最大后缀和;
- \(sum\):节点的和;
- \(mx\):节点的最大子段和;
考虑合并两个节点\([l_1,r_1],[l_2,r_2]\)时,最大子段和只可能在这三个值中产生:\(lmx,rmx,lsuf+rpre\)。
因为每次合并节点时会用到节点的\(pre,suf\),而\(pre,suf\)的更新依赖\(sum\),所以我们维护这4个值:
struct node{
ll pre,suf,sum,mx,lz;
node():pre(0),suf(0),sum(0),mx(0),lz(1e18){}
};
node operator+(const node&l,const node&r){
node res;
res.sum=l.sum+r.sum;
res.pre=max(l.pre,l.sum+r.pre);
res.suf=max(r.suf,r.sum+l.suf);
res.mx=max({l.mx,r.mx,l.suf+r.pre});
res.lz=1e18;
}
在区间询问时,每次\(query\)返回一个\(node\),用于合并。
node query(int u,int l,int r,int ql,int qr){
if(l>=ql&&r<=qr){
return tr[u];
}
int mid=(l+r)/2;
pushdown(u,l,r);
if(qr<=mid) return query(ls(u),l,mid,ql,qr);
else if(ql>mid) return query(rs(u),mid+1,r,ql,qr);
else{
node left=query(ls(u),l,mid,ql,qr);
node right=query(rs(u),mid+1,r,ql,qr);
node res=left+right;
return res;
}
}
完整代码:
#include <bits/stdc++.h>
using namespace std;
#define endl '\n'
using ll=long long;
/**
* 实现了最大子段和,区间查询和区间修改
*/
struct Segment{
int n;
vector<ll> val;
struct node{
ll pre,suf,sum,mx,lz;//lz需要复制成可以判别的值
node():pre(0),suf(0),sum(0),mx(0),lz(1e18){}
};
vector<node> tr;
Segment(){}
Segment(int n): n(n){
val.resize(n+1);
tr.resize((n<<2)+10);
}
void init(int n){
this->n=n;
val.resize(n+1);
tr.resize((n<<2)+10);
}
int ls(int x){return x<<1;}
int rs(int x){return x<<1|1;}
friend node operator+(const node& l,const node& r){
node res;
res.sum=l.sum+r.sum;
res.pre=max(l.pre,l.sum+r.pre);
res.suf=max(r.suf,r.sum+l.suf);
res.mx=max({l.mx,r.mx,l.suf+r.pre});
res.lz=1e18;//这个没有实际作用
return res;
}
void pushup(int u){
tr[u]=tr[ls(u)]+tr[rs(u)];
}
void build(int u,int l,int r){
if(l==r){
tr[u].sum=tr[u].suf=tr[u].pre=tr[u].mx=val[l];
tr[u].lz=(ll)1e18;
return;
}
int mid=(l+r)/2;
build(ls(u),l,mid);
build(rs(u),mid+1,r);
pushup(u);
}
void apply(int u,int l,int r,int z){
tr[u].sum=1LL*(r-l+1)*z;
//如果z为负数,那么区间mx,pre,suf都为z,否则为sum
tr[u].pre=tr[u].suf=tr[u].mx=(z>=0?tr[u].sum:z);
tr[u].lz=z;
}
void pushdown(int u,int l,int r){
if(tr[u].lz!=1e18){
int mid=(l+r)/2;
apply(ls(u),l,mid,tr[u].lz);
apply(rs(u),mid+1,r,tr[u].lz);
tr[u].lz=1e18;
}
}
void update(int u,int l,int r,int ql,int qr,int z){
if(l>=ql&&r<=qr){
apply(u,l,r,z);
return;
}
pushdown(u,l,r);
int mid=(l+r)/2;
if(ql<=mid) update(ls(u),l,mid,ql,qr,z);
if(qr>mid) update(rs(u),mid+1,r,ql,qr,z);
pushup(u);
}
node query(int u,int l,int r,int ql,int qr){
if(l>=ql&&r<=qr){
return tr[u];
}
int mid=(l+r)/2;
pushdown(u,l,r);
if(qr<=mid) return query(ls(u),l,mid,ql,qr);
else if(ql>mid) return query(rs(u),mid+1,r,ql,qr);
else{
node left=query(ls(u),l,mid,ql,qr);
node right=query(rs(u),mid+1,r,ql,qr);
node res=left+right;
return res;
}
}
};
void solve() {
int n;
cin >> n;
vector<int> a(n+1);
for(int i=1;i<=n;i++){
cin >> a[i];
}
Segment sgt(n);
for(int i=1;i<=n;i++){
sgt.val[i]=a[i];
}
sgt.build(1,1,n);
int q;
cin >> q;
while(q--){
int op,l,r;
cin >> op;
if(op==0){
int x,y;
cin >> x >> y;
sgt.update(1,1,n,x,x,y);
}else if(op==1){
cin >> l >> r;
auto it=sgt.query(1,1,n,l,r);
cout << it.mx << endl;
}
}
}
/*
*/
int main() {
ios::sync_with_stdio(0);
cin.tie(0), cout.tie(0);
int t = 1;
// cin >> t;
while (t--) {
solve();
}
return 0;
}
树上路径最大子段和
前置知识:树链剖分。
如果学了这个,那么都知道对于路径询问,我们可以将其拆成若干段,每段内都是连续的,那么在\(query\)时,我们查询每一段,然后进行合并就好。
对于查询\((u,v)\),\(u,v\)都会向\(lca\)跳,我们用node left表示u -> lca,node right表示v -> lca,那么最后会得到的两段由于方向一致,都向上,所以需要将其中一个做操作:swap(pre,suf)。

但是我们的代码中的合并操作是l.suf+r.pre,这样就接不上了,所以交换其suf,pre就可以实现合并,总的来说也就是要保持方向一致。
注意点:我们在合并节点时,一定是查询段和left或者right合并,顺序不能反。
这样我们就可以拿下了。
#include <bits/stdc++.h>
using namespace std;
#define endl '\n'
using ll=long long;
const int N = 2e5 + 9;
vector<int> g[N];
int depth[N]; // 节点深度
int fa[N]; // 父节点
int top[N]; // 所在重链的顶端节点
int son[N]; // 重儿子(子树最大的儿子)
int sz[N]; // 子树大小
int dfn[N], rnk[N]; // dfn: dfs序, rnk: dfs序对应的节点编号
int w[N], a[N]; // w: dfs序上的权值, a: 原节点权值
int idx; // 当前dfs序时间戳
// 第一次dfs:求父节点、深度、子树大小、重儿子
void dfs1(int u, int p) {
fa[u] = p;
depth[u] = depth[p] + 1;
sz[u] = 1;
for (int v : g[u]) {
if (v == p) continue;
dfs1(v, u);
sz[u] += sz[v];
// 更新重儿子:当前子节点v的子树大小大于当前记录的重儿子
if (sz[son[u]] < sz[v]) son[u] = v;
}
}
// 第二次dfs:分配每条重链的顶端和dfs序
void dfs2(int u, int t) {
top[u] = t; // 记录节点u所在重链的顶端
dfn[u] = ++idx; // 分配dfs序
rnk[dfn[u]] = u; // 记录dfs序对应的原节点编号
if (son[u]) dfs2(son[u], t); // 优先处理重儿子,继承当前链顶
for (int v : g[u]) {
if (v != son[u] && v != fa[u]) {
dfs2(v, v); // 轻儿子自己作为新链的链顶
}
}
}
struct node{
ll pre,suf,sum,mx,lz;//lz需要复制成可以判别的值
node():pre(0),suf(0),sum(0),mx(0),lz((ll)1e18){}
};
node operator+(const node& l,const node& r){
node res;
res.sum=l.sum+r.sum;
res.pre=max(l.pre,l.sum+r.pre);
res.suf=max(r.suf,r.sum+l.suf);
res.mx=max({l.mx,r.mx,l.suf+r.pre});
res.lz=1e18;//这个没有实际作用
return res;
}
struct Segment{
int n;
vector<ll> val;
vector<node> tr;
Segment(){}
Segment(int n): n(n){
val.resize(n+1);
tr.resize((n<<2)+10);
}
void init(int n){
this->n=n;
val.resize(n+1);
tr.resize((n<<2)+10);
}
int ls(int x){return x<<1;}
int rs(int x){return x<<1|1;}
void pushup(int u){
tr[u]=tr[ls(u)]+tr[rs(u)];
}
void build(int u,int l,int r){
if(l==r){
tr[u].sum=val[l];
tr[u].suf=tr[u].pre=tr[u].mx=max(0LL,val[l]);
tr[u].lz=(ll)1e18;
return;
}
int mid=(l+r)/2;
build(ls(u),l,mid);
build(rs(u),mid+1,r);
pushup(u);
}
void apply(int u,int l,int r,int z){
tr[u].sum=1LL*(r-l+1)*z;
//如果z为负数,那么区间mx,pre,suf都为z,否则为sum
tr[u].pre=tr[u].suf=tr[u].mx=max(0LL,tr[u].sum);
tr[u].lz=z;
}
void pushdown(int u,int l,int r){
if(tr[u].lz!=1e18){
int mid=(l+r)/2;
apply(ls(u),l,mid,tr[u].lz);
apply(rs(u),mid+1,r,tr[u].lz);
tr[u].lz=1e18;
}
}
void update(int u,int l,int r,int ql,int qr,int z){
if(l>=ql&&r<=qr){
apply(u,l,r,z);
return;
}
pushdown(u,l,r);
int mid=(l+r)/2;
if(ql<=mid) update(ls(u),l,mid,ql,qr,z);
if(qr>mid) update(rs(u),mid+1,r,ql,qr,z);
pushup(u);
}
node query(int u,int l,int r,int ql,int qr){
if(l>=ql&&r<=qr){
return tr[u];
}
int mid=(l+r)/2;
pushdown(u,l,r);
if(qr<=mid) return query(ls(u),l,mid,ql,qr);
else if(ql>mid) return query(rs(u),mid+1,r,ql,qr);
else{
node left=query(ls(u),l,mid,ql,qr);
node right=query(rs(u),mid+1,r,ql,qr);
node res=left+right;
return res;
}
}
}sgt;
int n;
ll query_path(int u,int v){
node left,right;//left:u->lca right:v->lca,最后只需要其中一个swap一下就好
while(top[u]!=top[v]){
if(depth[top[u]]<depth[top[v]]){
right=sgt.query(1,1,n,dfn[top[v]],dfn[v])+right;
v=fa[top[v]];
}else{
left=sgt.query(1,1,n,dfn[top[u]],dfn[u])+left;
u=fa[top[u]];
}
}
// u和v在同一重链上
if (dfn[u] > dfn[v]) {
// u在下面,v在上面
left = sgt.query(1, 1, n, dfn[v], dfn[u]) + left;
} else {
// u在上面,v在下面
right = sgt.query(1, 1, n, dfn[u], dfn[v]) + right;
}
swap(left.pre,left.suf);
node res=left+right;
return max(0LL,res.mx);
}
void update_path(int u,int v,int z){
while(top[u]!=top[v]){
if(depth[top[u]]<depth[top[v]]) swap(u,v);
sgt.update(1,1,n,dfn[top[u]],dfn[u],z);
u=fa[top[u]];
}
if(dfn[u]>dfn[v]) swap(u,v);
sgt.update(1,1,n,dfn[u],dfn[v],z);
}
void solve() {
cin >> n;
for(int i=1;i<=n;i++){
cin >> a[i];
}
for(int i=1;i<n;i++){
int u,v;
cin >> u >> v;
g[u].push_back(v);
g[v].push_back(u);
}
dfs1(1,0);
dfs2(1,1);
sgt.init(n);
for(int i=1;i<=n;i++){
w[dfn[i]]=a[i];
sgt.val[dfn[i]]=a[i];
}
sgt.build(1,1,n);
int q;
cin >> q;
while(q--){
int op;
cin >> op;
if(op==1){
int u,v;
cin >> u >> v;
cout << query_path(u,v) << endl;
}else if(op==2){
int u,v,z;
cin >> u >> v >> z;
update_path(u,v,z);
}
}
}
/*
*/
int main() {
ios::sync_with_stdio(0);
cin.tie(0), cout.tie(0);
int t = 1;
// cin >> t;
while (t--) {
solve();
}
return 0;
}
题目链接:
SP6779 GSS7 - Can you answer these queries VII - 洛谷
SP2916 GSS5 - Can you answer these queries V - 洛谷
[P2572 SCOI2010] 序列操作 - 洛谷
第二题ac代码:
#include <bits/stdc++.h>
using namespace std;
#define endl '\n'
using ll=long long;
struct Segment{
int n;
vector<ll> val;
struct node{
ll pre,suf,sum,mx,lz;//lz需要复制成可以判别的值
node():pre(0),suf(0),sum(0),mx(0),lz(1e18){}
node(ll x):pre(x),suf(x),sum(x),mx(x),lz(1e18){}
};
vector<node> tr;
Segment(){}
Segment(int n): n(n){
val.resize(n+1);
tr.resize((n<<2)+10);
}
void init(int n){
this->n=n;
val.resize(n+1);
tr.resize((n<<2)+10);
}
int ls(int x){return x<<1;}
int rs(int x){return x<<1|1;}
friend node operator+(const node& l,const node& r){
node res;
res.sum=l.sum+r.sum;
res.pre=max(l.pre,l.sum+r.pre);
res.suf=max(r.suf,r.sum+l.suf);
res.mx=max({l.mx,r.mx,l.suf+r.pre});
res.lz=1e18;//这个没有实际作用
return res;
}
void pushup(int u){
tr[u]=tr[ls(u)]+tr[rs(u)];
}
void build(int u,int l,int r){
if(l==r){
tr[u].sum=tr[u].suf=tr[u].pre=tr[u].mx=val[l];
tr[u].lz=(ll)1e18;
return;
}
int mid=(l+r)/2;
build(ls(u),l,mid);
build(rs(u),mid+1,r);
pushup(u);
}
void apply(int u,int l,int r,int z){
tr[u].sum=1LL*(r-l+1)*z;
//如果z为负数,那么区间mx,pre,suf都为z,否则为sum
tr[u].pre=tr[u].suf=tr[u].mx=(z>=0?tr[u].sum:z);
tr[u].lz=z;
}
void pushdown(int u,int l,int r){
if(tr[u].lz!=1e18){
int mid=(l+r)/2;
apply(ls(u),l,mid,tr[u].lz);
apply(rs(u),mid+1,r,tr[u].lz);
tr[u].lz=1e18;
}
}
void update(int u,int l,int r,int ql,int qr,int z){
if(l>=ql&&r<=qr){
apply(u,l,r,z);
return;
}
pushdown(u,l,r);
int mid=(l+r)/2;
if(ql<=mid) update(ls(u),l,mid,ql,qr,z);
if(qr>mid) update(rs(u),mid+1,r,ql,qr,z);
pushup(u);
}
node query(int u,int l,int r,int ql,int qr){
if(ql>qr){
return node(-1e12);
}
if(l>=ql&&r<=qr){
return tr[u];
}
int mid=(l+r)/2;
pushdown(u,l,r);
if(qr<=mid) return query(ls(u),l,mid,ql,qr);
else if(ql>mid) return query(rs(u),mid+1,r,ql,qr);
else{
node left=query(ls(u),l,mid,ql,qr);
node right=query(rs(u),mid+1,r,ql,qr);
node res=left+right;
return res;
}
}
};
void solve() {
int n;
cin >> n;
vector<int> a(n+1),pre(n+1);
for(int i=1;i<=n;i++){
cin >> a[i];
pre[i]=pre[i-1]+a[i];
}
Segment sgt(n);
for(int i=1;i<=n;i++){
sgt.val[i]=a[i];
}
sgt.build(1,1,n);
int q;
cin >> q;
while(q--){
int l1,r1,l2,r2;
cin >> l1 >> r1 >> l2 >> r2;
if(r1<l2){
auto it1=sgt.query(1,1,n,l1,r1);
auto it2=sgt.query(1,1,n,l2,r2);
ll ans=pre[l2-1]-pre[r1]+it1.suf+it2.pre;
cout << ans << endl;
}else{
auto it1=sgt.query(1,1,n,l1,l2-1);
auto it2=sgt.query(1,1,n,l2,r1);
auto it3=sgt.query(1,1,n,r1+1,r2);
ll ans=max({it1.suf+pre[r1]-pre[l2-1]+it3.pre,it1.suf+it2.pre,it2.suf+it3.pre,it2.mx});
cout << ans << endl;
}
}
}
/*
*/
int main() {
ios::sync_with_stdio(0);
cin.tie(0), cout.tie(0);
int t = 1;
cin >> t;
while (t--) {
solve();
}
return 0;
}
第三题ac代码:
#include <bits/stdc++.h>
using namespace std;
#define endl '\n'
using ll=long long;
struct Segment{
int n;
vector<ll> val;
struct node{
ll pre0,suf0,pre1,suf1,sum,lnum,rnum,mx0,mx1,len,lz1,lz2;
node():pre0(0),suf0(0),pre1(0),suf1(0),sum(0),lnum(0),rnum(0),mx0(0),mx1(0),len(0),lz1(-1),lz2(-1){}
};
vector<node> tr;
Segment(){}
Segment(int n):n(n){
val.resize(n+1);
tr.resize((n<<2)+10);
}
void init(int n){
this->n=n;
val.resize(n+1);
tr.resize((n<<2)+10);
}
friend node operator+(const node&l,const node&r){
node res;
res.sum=l.sum+r.sum;
res.len=l.len+r.len;
res.pre0=l.pre0;
if(l.pre0==l.len) res.pre0=max(res.pre0,l.len+r.pre0);
res.pre1=l.pre1;
if(l.pre1==l.len) res.pre1=max(res.pre1,l.len+r.pre1);
res.suf0=r.suf0;
if(r.suf0==r.len) res.suf0=max(res.suf0,r.len+l.suf0);
res.suf1=r.suf1;
if(r.suf1==r.len) res.suf1=max(res.suf1,r.len+l.suf1);
res.lnum=l.lnum;
res.rnum=r.rnum;
res.mx0=max({l.mx0,r.mx0,l.suf0+r.pre0});
res.mx1=max({l.mx1,r.mx1,l.suf1+r.pre1});
res.lz1=-1;
res.lz2=-1;
return res;
}
int ls(int x){return x<<1;}
int rs(int x){return x<<1|1;}
void pushup(int u){
tr[u]=tr[ls(u)]+tr[rs(u)];
}
void apply1(int u,int l,int r,int z){
tr[u].sum=tr[u].len*z;
tr[u].mx1=(z==1?tr[u].len:0);
tr[u].mx0=(z==0?tr[u].len:0);
tr[u].pre0=tr[u].suf0=(z==0?tr[u].len:0);
tr[u].pre1=tr[u].suf1=(z==1?tr[u].len:0);
tr[u].lnum=tr[u].rnum=z;
tr[u].lz1=z;
tr[u].lz2=-1;
}
//区间取反
void apply2(int u,int l,int r,int z){
tr[u].sum=tr[u].len-tr[u].sum;
swap(tr[u].pre1,tr[u].pre0);
swap(tr[u].suf1,tr[u].suf0);
swap(tr[u].mx1,tr[u].mx0);
tr[u].lnum=1-tr[u].lnum;
tr[u].rnum=1-tr[u].rnum;
//如果lz1有标记,那么直接对lz1取反即可
//如果lz1没有标记,同时lz2有标记取反,再次取反就不会修改,所以赋值为-1
//如果二者都没有标记,赋值为z
if(tr[u].lz1!=-1){
tr[u].lz1=1-tr[u].lz1;
tr[u].lz2=-1;
}else if(tr[u].lz2!=-1){
tr[u].lz2=-1;
}else{
tr[u].lz2=z;
}
}
//lz1表示区间变成lz1,lz2表示区间取反
void pushdown(int u,int l,int r){
if(tr[u].lz1!=-1||tr[u].lz2!=-1){
int mid=(l+r)/2;
if(tr[u].lz1!=-1){
apply1(ls(u),l,mid,tr[u].lz1);
apply1(rs(u),mid+1,r,tr[u].lz1);
tr[u].lz1=-1;
}else{
apply2(ls(u),l,mid,tr[u].lz2);
apply2(rs(u),mid+1,r,tr[u].lz2);
tr[u].lz2=-1;
}
}
}
void build(int u,int l,int r){
if(l==r){
tr[u].len=1;
tr[u].sum=tr[u].lnum=tr[u].rnum=val[l];
tr[u].mx1=tr[u].pre1=tr[u].suf1=(val[l]==1);
tr[u].mx0=tr[u].pre0=tr[u].suf0=(val[l]==0);
tr[u].lz1=tr[u].lz2=-1;
return;
}
int mid=(l+r)/2;
build(ls(u),l,mid);
build(rs(u),mid+1,r);
pushup(u);
}
//区间推平
void update1(int u,int l,int r,int ql,int qr,int z){
if(l>=ql&&r<=qr){
apply1(u,l,r,z);
return;
}
int mid=(l+r)/2;
pushdown(u,l,r);
if(ql<=mid) update1(ls(u),l,mid,ql,qr,z);
if(qr>mid) update1(rs(u),mid+1,r,ql,qr,z);
pushup(u);
}
//区间取反
void update2(int u,int l,int r,int ql,int qr,int z){
if(l>=ql&&r<=qr){
apply2(u,l,r,z);
return;
}
int mid=(l+r)/2;
pushdown(u,l,r);
if(ql<=mid) update2(ls(u),l,mid,ql,qr,z);
if(qr>mid) update2(rs(u),mid+1,r,ql,qr,z);
pushup(u);
}
node query(int u,int l,int r,int ql,int qr){
if(l>=ql&&r<=qr){
return tr[u];
}
pushdown(u,l,r);
int mid=(l+r)/2;
if(qr<=mid) return query(ls(u),l,mid,ql,qr);
else if(ql>mid) return query(rs(u),mid+1,r,ql,qr);
else{
return query(ls(u),l,mid,ql,qr)+query(rs(u),mid+1,r,ql,qr);
}
}
};
void solve() {
int n,q;
cin >> n >> q;
vector<int> a(n+1);
for(int i=1;i<=n;i++){
cin >> a[i];
}
Segment sgt(n);
for(int i=1;i<=n;i++){
sgt.val[i]=a[i];
}
sgt.build(1,1,n);
while(q--){
int op,l,r;
cin >> op >> l >> r;
l++;
r++;
if(op==0){
sgt.update1(1,1,n,l,r,0);
}else if(op==1){
sgt.update1(1,1,n,l,r,1);
}else if(op==2){
sgt.update2(1,1,n,l,r,1);
}else if(op==3){
auto it=sgt.query(1,1,n,l,r);
cout << it.sum << endl;
}else if(op==4){
auto it=sgt.query(1,1,n,l,r);
cout << it.mx1 << endl;
}
}
}
/*
*/
int main() {
ios::sync_with_stdio(0);
cin.tie(0), cout.tie(0);
int t = 1;
// cin >> t;
while (t--) {
solve();
}
return 0;
}

浙公网安备 33010602011771号