代码虽长,但通俗易懂QAQ。
#include<bits/stdc++.h>
using namespace std;
const int N=5e5+100;
int n,m;
double ax,ay,bx,by,cx,cy,dx,dy;
struct node{double x,y;}a[N];
bool cmpp(node A,node B){return A.x<B.x;}
struct edge{int u,v;double w;}e[N];int cnt;
bool cmp(edge A,edge B){return A.w<B.w;}
double pw(double x){return x*x;}
double dis(int i,int j){return sqrt(pw(a[i].x-a[j].x)+pw(a[i].y-a[j].y));}
void insert(int i,int j){
cnt++;
e[cnt].u=i,e[cnt].v=j;
e[cnt].w=dis(i,j);
}
int fa[N];
int find(int x){return fa[x]==x?x:fa[x]=find(fa[x]);}
void krual(){
for(int i=1;i<=n+m;i++) fa[i]=i;
double ans=0;int all=0;
sort(e+1,e+1+cnt,cmp);
for(int i=1;i<=cnt;i++){
int u=e[i].u,v=e[i].v;
int X=find(u),Y=find(v);
if(X==Y) continue;
fa[Y]=X,ans+=e[i].w,all++;
if(all==n+m-1) break;
}
printf("%.3lf",ans);
}
int main(){
cin>>n>>m;
cin>>ax>>ay>>bx>>by;
cin>>cx>>cy>>dx>>dy;
for(int i=1;i<=n;i++){
double t;cin>>t;
a[i]=(node){ax*t+bx*(1-t),ay*t+by*(1-t)};
}
sort(a+1,a+1+n,cmpp);
for(int i=1+n;i<=m+n;i++){
double t;cin>>t;
a[i]=(node){cx*t+dx*(1-t),cy*t+dy*(1-t)};
}
sort(a+n+1,a+1+n+m,cmpp);
for(int i=1;i<=n;i++){
int bb=i+1;
if(i!=n) insert(i,bb);
int l=n+2,r=n+m,pos=n+1;
while(l<=r){//二分i点与另一条直线上哪个点距离最短
int mid=(l+r)>>1;
if(dis(i,mid)<dis(i,mid-1)) pos=mid,l=mid+1;
else r=mid-1;
}
insert(i,pos);
if(pos-1>=n+1) insert(i,pos-1);
if(pos+1<=n+m) insert(i,pos+1);
//保险起见,把相邻两个也加上去
}
for(int i=1+n;i<=m+n-1;i++) insert(i,i+1);
krual();
}