20pts WA求助 悬赏1关注
查看原帖
20pts WA求助 悬赏1关注
743811
Shakespeare07楼主2023/4/7 19:44

rt.

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

#define int long long

int read(){
	int s=0,w=1; char c=getchar();
	while(!isdigit(c)){ if(c=='-') w=-1; c=getchar();}
	while(isdigit(c)){ s=(s<<3)+(s<<1)+(c^48); c=getchar();}
	return s*w;
}

const int N=1e6+5;

int n,m;
int a[N];
int Sum[N];

struct sgt{
	int add,tag,sum,mx;
	#define add(x) tr[x].add
	#define sum(x) tr[x].sum
	#define tag(x) tr[x].tag
	#define mx(x) tr[x].mx
	int l,r;
	#define l(x) tr[x].l
	#define r(x) tr[x].r
}tr[N<<2];

void pushup(int x){
	sum(x)=sum(x<<1)+sum(x<<1|1);
	mx(x)=mx(x<<1|1);
}

void build(int p,int l,int r){
	l(p)=l,r(p)=r,tag(p)=-1;
	if(l==r) return;
	int mid=l+r>>1;
	build(p<<1,l,mid);
	build(p<<1|1,mid+1,r);
}

void pushdown(int x){
	if(~tag(x)){
		int tmp=tag(x);
		tag(x<<1)=tag(x<<1|1)=tmp;
		sum(x<<1)=(r(x<<1)-l(x<<1)+1)*tmp;
		sum(x<<1|1)=(r(x<<1|1)-l(x<<1|1)+1)*tmp;
		mx(x<<1)=mx(x<<1|1)=tmp;
		add(x<<1)=add(x<<1|1)=0;
		tag(x)=-1;
		return; 
	}
	if(add(x)){
		int tmp=add(x);
		sum(x<<1)+=(Sum[r(x<<1)]-Sum[l(x<<1)-1])*tmp;
		sum(x<<1|1)+=(Sum[r(x<<1|1)]-Sum[l(x<<1|1)-1])*tmp;
		mx(x<<1)+=a[r(x<<1)]*tmp;
		mx(x<<1|1)+=a[r(x<<1|1)]*tmp;
		add(x<<1)+=tmp;
		add(x<<1|1)+=tmp;
		add(x)=0;
		return;
	}
}

int find(int p,int l,int r,int d){
	if(l==r) return l;
	pushdown(p);
	int mid=l+r>>1;
	if(mx(p<<1)>=d) return find(p<<1,l,mid,d);
	return find(p<<1|1,mid+1,r,d);
}

int cut(int p,int l,int r,int ql,int qr,int d){
	if(l>=ql && r<=qr){
		int yuan=sum(p);
		tag(p)=d;
		sum(p)=d*(r-l+1);
		mx(p)=d;
		add(p)=0;
		return yuan-sum(p);
	}
	pushdown(p);
	int mid=l+r>>1,res=0;
	if(ql<=mid) res+=cut(p<<1,l,mid,ql,qr,d);
	if(qr>mid) res+=cut(p<<1|1,mid+1,r,ql,qr,d);
	pushup(p);
	return res;
}

signed main(){
	n=read(),m=read();
	for(int i=1;i<=n;++i) a[i]=read();
	
	sort(a+1,a+n+1);
	for(int i=1;i<=n;++i) Sum[i]=Sum[i-1]+a[i];
	
	build(1,1,n);
	
	int lst=0;
	while(m--){
		int x=read(),y=read();
		
		pushdown(1);
		add(1)+=x-lst;
		sum(1)+=(x-lst)*Sum[n];
		mx(1)+=(x-lst)*a[n];
		
		if(mx(1)<y){
			puts("0");
			lst=x;
			continue;
		}
		
		int st=find(1,1,n,y);
		int tmp=cut(1,1,n,st,n,y);
		printf("%lld\n",tmp);
		lst=x;
	}
	
	return 0;
}
2023/4/7 19:44
加载中...