树状数组WA on #7 求助
查看原帖
树状数组WA on #7 求助
438461
liu_chen_hao楼主2022/8/12 17:04

RT,我的代码如下QWQ,麻烦大佬帮忙看看:

#include <bits/stdc++.h>
#define long long ll
#define pb(x) push_back(x)
using namespace std;
const int N=1e6+5,M=1e6+5;

struct node {
	int x,y,xx,yy,s1,s2,s3,s4;
}a[N];
struct line {
	int l,r,id;
	bool eg;
}tmp;
int n,m,x[N],y[N],xx[M],yy[M],cnt,nx,ny,mxx,mx,t[N];
vector<int> h[M];
vector<line> g[M];

int read() {
	int s=0,w=1;
	char ch=getchar();
	while(ch<'0' || ch>'9') {if(ch=='-') w=-1;ch=getchar();}
	while(ch>='0' && ch<='9') s=s*10+ch-'0',ch=getchar();
	return s*w;
}
void add(int p, int d) {
	for(; p<=mx; p+=p&-p) t[p]+=d;
}
int query(int p) {
	int rnt=0;
	for(; p; p-=p&-p) rnt+=t[p];
	return rnt;
}
int main()
{
	scanf("%d%d", &n, &m);
	for(int i=1; i<=n; i++)
	{
		x[i]=read()+1;
		y[i]=read()+1;  // +1是防止输入0,然后-1访问到数组的-1处RE,可能这个处理方式麻烦了
		xx[++cnt]=x[i];
		yy[cnt]=y[i];
	}
	for(int i=1; i<=m; i++)
	{
		a[i].x=read()+1;
		a[i].y=read()+1;
		a[i].xx=read()+1;
		a[i].yy=read()+1;
		xx[++cnt]=a[i].x;
		yy[cnt]=a[i].y;
		xx[++cnt]=a[i].xx;
		yy[cnt]=a[i].yy;
	}

	sort(xx+1,xx+cnt+1);
	sort(yy+1,yy+cnt+1);
	nx=unique(xx+1,xx+cnt+1)-xx-1;
	ny=unique(yy+1,yy+cnt+1)-yy-1;
	for(int i=1; i<=n; i++)
	{
		x[i]=lower_bound(xx+1,xx+nx+1,x[i])-xx;
		y[i]=lower_bound(yy+1,yy+ny+1,y[i])-yy;
		h[y[i]].pb(x[i]);
		mxx=max(mxx,y[i]);
		mx=max(mx,x[i]);
	}
	for(int i=1; i<=m; i++)
	{
		a[i].x=lower_bound(xx+1,xx+nx+1,a[i].x)-xx;
		a[i].xx=lower_bound(xx+1,xx+nx+1,a[i].xx)-xx;
		a[i].y=lower_bound(yy+1,yy+ny+1,a[i].y)-yy;
		a[i].yy=lower_bound(yy+1,yy+ny+1,a[i].yy)-yy;
		tmp.id=i;
		tmp.l=a[i].x;
		tmp.r=a[i].xx;
		tmp.eg=0;
		g[a[i].y-1].pb(tmp);
		tmp.eg=1;
		g[a[i].yy].pb(tmp);
		mxx=max(mxx,a[i].yy);
		mx=max(mx,a[i].xx);
	}

	for(int i=0; i<=mxx; i++)
	{
		for(int j=0; j<h[i].size(); j++) add(h[i][j],1);
		for(int j=0; j<g[i].size(); j++)
		{
			tmp=g[i][j];
			if(tmp.eg)
			{
				a[tmp.id].s3=query(tmp.l-1);
				a[tmp.id].s4=query(tmp.r);
			}
			else
			{
				a[tmp.id].s1=query(tmp.l-1);
				a[tmp.id].s2=query(tmp.r);
			}
		}
	}
	for(int i=1; i<=m; i++) printf("%d\n", a[i].s4-a[i].s3-a[i].s2+a[i].s1);
	
	return 0;
}
2022/8/12 17:04
加载中...