萌新求助k-d tree
查看原帖
萌新求助k-d tree
363529
ForLune_楼主2023/7/18 20:55

rt,原来交替建树,TLE #11 ,改为方差建树后不仅没过 #11 ,而且还 WA 了几个点(除 #5 #6 #7 三个点之外其他的点都没过)。

望各路大佬帮忙看一下问题!本人感激不尽!!

orz

交替建树代码:

#include <stack>
#include <cmath>
#include <cstdio>
#include <algorithm>
#define lint long long
#define x first
#define y second
using namespace std;
const lint inf=2e18;
int n,q; lint ans;
pair<lint,lint> data[600005];
struct Node
{
	int left,right,size;
	pair<lint,lint> pos,minn,maxn;
};
int root,TreeSize;
Node node[600005];
stack<int> Stack;
inline int read()
{
    int x=0,f=1;
    char ch=getchar();
    while(ch<'0'||ch>'9') {if(ch=='-') f=-1;ch=getchar();}
    while(ch>='0'&&ch<='9') {x=(x<<1)+(x<<3)+(ch^48);ch=getchar();}
    return x*f;
}
inline bool cmpx(const pair<lint,lint> A,const pair<lint,lint> B) {return A.x<B.x;}
inline bool cmpy(const pair<lint,lint> A,const pair<lint,lint> B) {return A.y<B.y;}
inline lint dist(const pair<lint,lint> A,const pair<lint,lint> B) {return abs(A.x-B.x)+abs(A.y-B.y);}
inline int NewNode()
{
	if(Stack.empty()) return ++TreeSize;
	int temp=Stack.top(); Stack.pop();
	return temp;
}
inline void PushUp(const int u)
{
	if(node[u].left)
	{
		node[u].size+=node[node[u].left].size;
		node[u].minn.x=min(node[u].minn.x,node[node[u].left].minn.x);
		node[u].minn.y=min(node[u].minn.y,node[node[u].left].minn.y);
		node[u].maxn.x=max(node[u].maxn.x,node[node[u].left].maxn.x);
		node[u].maxn.y=max(node[u].maxn.y,node[node[u].left].maxn.y);
	}
	if(node[u].right)
	{
		node[u].size+=node[node[u].right].size;
		node[u].minn.x=min(node[u].minn.x,node[node[u].right].minn.x);
		node[u].minn.y=min(node[u].minn.y,node[node[u].right].minn.y);
		node[u].maxn.x=max(node[u].maxn.x,node[node[u].right].maxn.x);
		node[u].maxn.y=max(node[u].maxn.y,node[node[u].right].maxn.y);
	}
}
inline void build(const int l,const int r,int &u,const bool mode)
{
	if(l>r) return;
	if(!u) u=NewNode();
	int mid=(l+r)>>1;
	nth_element(data+l,data+mid,data+r+1,(mode)?(cmpx):(cmpy));
	node[u]=(Node){0,0,0,data[mid],data[mid],data[mid]};
	build(l,mid-1,node[u].left,mode^1);
	build(mid+1,r,node[u].right,mode^1);
	PushUp(u);
}
inline void GetNode(const int now,const int u)
{
	if(node[u].left) GetNode(now,node[u].left);
	data[now+node[node[u].left].size+1]=node[u].pos,Stack.push(u);
	if(node[u].right) GetNode(now+node[node[u].left].size+1,node[u].right);
}
inline void check(int &u,const bool mode)
{
	if(node[node[u].left].size>node[node[u].right].size*4||node[node[u].left].size*4<node[node[u].right].size)
		GetNode(0,u),build(1,node[u].size,u,mode);
}
inline void insert(const pair<lint,lint> Value,int &u,const bool mode)
{
	if(!u) {u=NewNode(),node[u]=(Node){0,0,0,Value,Value,Value},PushUp(u); return;}
	if((mode)?(node[u].pos.x>=Value.x):(node[u].pos.y>=Value.y)) insert(Value,node[u].left,mode^1);
	if((mode)?(node[u].pos.x<Value.x):(node[u].pos.y<Value.y)) insert(Value,node[u].right,mode^1);
	PushUp(u),check(u,mode);
}
inline lint calc(const pair<lint,lint> target,const int u)
{
	if(!u) return inf;
	lint res=0;
	if(node[u].minn.x>target.x) res+=node[u].minn.x-target.x;
	if(node[u].minn.y>target.y) res+=node[u].minn.y-target.y;
	if(node[u].maxn.x<target.x) res+=target.x-node[u].maxn.x;
	if(node[u].maxn.y<target.y) res+=target.y-node[u].maxn.y;
	return res;
}
inline void query(const pair<lint,lint> target,const int u)
{
	ans=min(ans,dist(target,node[u].pos));
	lint lvalue=calc(target,node[u].left),rvalue=calc(target,node[u].right);
	if(lvalue<rvalue)
	{
		if(lvalue<ans) query(target,node[u].left);
		if(rvalue<ans) query(target,node[u].right);
	}
	if(lvalue>=rvalue)
	{
		if(rvalue<ans) query(target,node[u].right);
		if(lvalue<ans) query(target,node[u].left);
	}
}
int main()
{
	n=read(),q=read();
	for(int i=1;i<=n;++i) data[i].first=1ll*read(),data[i].second=1ll*read();
	build(1,n,root,false);
	while(q--)
	{
		int opt=read(); lint x=1ll*read(),y=1ll*read(); ans=inf;
		if(opt==1) insert(make_pair(x,y),root,false);
		if(opt==2) query(make_pair(x,y),root),printf("%lld\n",ans);
	}
    return 0;
}

方差建树代码:

#include <stack>
#include <cmath>
#include <cstdio>
#include <algorithm>
#define lint long long
#define x first
#define y second
using namespace std;
const lint inf=2e18;
int n,q; lint ans;
pair<lint,lint> data[600005];
struct Node
{
	int left,right,size,mode;
	pair<lint,lint> pos,minn,maxn;
};
int root,TreeSize;
Node node[600005];
stack<int> Stack;
inline int read()
{
    int x=0,f=1;
    char ch=getchar();
    while(ch<'0'||ch>'9') {if(ch=='-') f=-1;ch=getchar();}
    while(ch>='0'&&ch<='9') {x=(x<<1)+(x<<3)+(ch^48);ch=getchar();}
    return x*f;
}
inline bool cmpx(const pair<lint,lint> A,const pair<lint,lint> B) {return A.x<B.x;}
inline bool cmpy(const pair<lint,lint> A,const pair<lint,lint> B) {return A.y<B.y;}
inline lint dist(const pair<lint,lint> A,const pair<lint,lint> B) {return abs(A.x-B.x)+abs(A.y-B.y);}
inline int NewNode()
{
	if(Stack.empty()) return ++TreeSize;
	int temp=Stack.top(); Stack.pop();
	return temp;
}
inline int choose(const int L,const int R)
{
	lint x1=0,y1=0,x2=0,y2=0;
	for(int i=L;i<=R;++i) x1+=data[i].x,y1+=data[i].y;
	x1/=1ll*(R-L+1),y1/=1ll*(R-L+1);
	for(int i=L;i<=R;++i) x2+=(data[i].x-x1)*(data[i].x-x1),y2+=(data[i].y-y1)*(data[i].y-y1);
	return (x2>y2)?(1):(2);
}
inline void PushUp(const int u)
{
	if(node[u].left)
	{
		node[u].size+=node[node[u].left].size;
		node[u].minn.x=min(node[u].minn.x,node[node[u].left].minn.x);
		node[u].minn.y=min(node[u].minn.y,node[node[u].left].minn.y);
		node[u].maxn.x=max(node[u].maxn.x,node[node[u].left].maxn.x);
		node[u].maxn.y=max(node[u].maxn.y,node[node[u].left].maxn.y);
	}
	if(node[u].right)
	{
		node[u].size+=node[node[u].right].size;
		node[u].minn.x=min(node[u].minn.x,node[node[u].right].minn.x);
		node[u].minn.y=min(node[u].minn.y,node[node[u].right].minn.y);
		node[u].maxn.x=max(node[u].maxn.x,node[node[u].right].maxn.x);
		node[u].maxn.y=max(node[u].maxn.y,node[node[u].right].maxn.y);
	}
}
inline void build(const int l,const int r,int &u)
{
	if(l>r) return;
	if(!u) u=NewNode();
	int mid=(l+r)>>1,mode=choose(l,r);
	nth_element(data+l,data+mid,data+r+1,(mode==1)?(cmpx):(cmpy));
	node[u]=(Node){0,0,0,mode,data[mid],data[mid],data[mid]};
	build(l,mid-1,node[u].left);
	build(mid+1,r,node[u].right);
	PushUp(u);
}
inline void GetNode(const int now,const int u)
{
	if(node[u].left) GetNode(now,node[u].left);
	data[now+node[node[u].left].size+1]=node[u].pos,Stack.push(u);
	if(node[u].right) GetNode(now+node[node[u].left].size+1,node[u].right);
}
inline void check(int &u)
{
	if(node[node[u].left].size>node[node[u].right].size*3||node[node[u].left].size*3<node[node[u].right].size)
		GetNode(0,u),build(1,node[u].size,u);
}
inline void insert(const pair<lint,lint> Value,int &u)
{
	if(!u) {u=NewNode(),node[u]=(Node){0,0,0,1,Value,Value,Value},PushUp(u); return;}
	if((node[u].mode==1)?(node[u].pos.x>=Value.x):(node[u].pos.y>=Value.y)) insert(Value,node[u].left);
	if((node[u].mode==2)?(node[u].pos.x<Value.x):(node[u].pos.y<Value.y)) insert(Value,node[u].right);
	PushUp(u),check(u);
}
inline lint calc(const pair<lint,lint> target,const int u)
{
	if(!u) return inf;
	lint res=0;
	if(node[u].minn.x>target.x) res+=node[u].minn.x-target.x;
	if(node[u].minn.y>target.y) res+=node[u].minn.y-target.y;
	if(node[u].maxn.x<target.x) res+=target.x-node[u].maxn.x;
	if(node[u].maxn.y<target.y) res+=target.y-node[u].maxn.y;
	return res;
}
inline void query(const pair<lint,lint> target,const int u)
{
	ans=min(ans,dist(target,node[u].pos));
	lint lvalue=calc(target,node[u].left),rvalue=calc(target,node[u].right);
	if(lvalue<rvalue)
	{
		if(lvalue<ans) query(target,node[u].left);
		if(rvalue<ans) query(target,node[u].right);
	}
	if(lvalue>=rvalue)
	{
		if(rvalue<ans) query(target,node[u].right);
		if(lvalue<ans) query(target,node[u].left);
	}
}
int main()
{
	n=read(),q=read();
	for(int i=1;i<=n;++i) data[i].first=1ll*read(),data[i].second=1ll*read();
	build(1,n,root);
	while(q--)
	{
		int opt=read(); lint x=1ll*read(),y=1ll*read(); ans=inf;
		if(opt==1) insert(make_pair(x,y),root);
		if(opt==2) query(make_pair(x,y),root),printf("%lld\n",ans);
	}
    return 0;
}
2023/7/18 20:55
加载中...