萌新求助倍增DP,WA3个点
查看原帖
萌新求助倍增DP,WA3个点
86896
rmxlinux楼主2022/8/26 10:00

评测记录

1个点问题2算错,2个点问题1算错

#include <iostream>
#include <set>
#include <cmath>
#include <vector>
#include <cstring>
using namespace std ;
#define INF 1e9
int n, m, x0 ;
int a[100005] ;
int ga[100005] ; // a choose second far
int gb[100005] ; // b choose farthest
int f[20][100005][2] ;//0=a 1=b
int da[20][100005][2] ;
int db[20][100005][2] ;
int len ;
int dist(int x, int y) {
	if(x >= 1 && x <= n && y >= 1 && y <= n)
		return abs(a[x] - a[y]) ;
	else return INF ;
}
void getgg() {
	set<pair<int, int> > s ;
	s.clear() ;
	for(int i = n; i >= 1; i--) {
		s.insert(make_pair(a[i], i)) ;
		set<pair<int, int> >::iterator ttmp = s.end() ; ttmp-- ;
		set<pair<int, int> > tmp ;
		set<pair<int, int> >::iterator tr = s.find(make_pair(a[i], i)) ;
		if(tr == s.begin()) tmp.insert(make_pair(INF, 0)) ;
		else --tr, tmp.insert(make_pair(abs((tr->first) - a[i]), tr->second)) ;
		if(tr == s.begin()) tmp.insert(make_pair(INF, 0)) ;
		else --tr, tmp.insert(make_pair(abs((tr->first) - a[i]), tr->second)) ;
		tr = s.find(make_pair(a[i], i)) ;
		if(tr == ttmp) tmp.insert(make_pair(INF, 0)) ;
		else ++tr, tmp.insert(make_pair(abs((tr->first) - a[i]), tr->second)) ;
		if(tr == ttmp) tmp.insert(make_pair(INF, 0)) ;
		else ++tr, tmp.insert(make_pair(abs((tr->first) - a[i]), tr->second)) ;
		int mn1 = tmp.begin()->first ;
		int id1 = tmp.begin()->second ;
		tmp.erase(tmp.begin()) ;
		int mn2 = tmp.begin()->first ;
		int id2 = tmp.begin()->second ;
		//cout<<i<<' '<<mn1<<' '<<mn2<<' '<<id1<<' '<<id2<<endl ;
		if(mn1 != mn2) ga[i] = id2, gb[i] = id1 ;
		else {
			gb[i] = (a[id1] < a[id2]) ? id1 : id2 ;
			if(gb[i] == id2) ga[i] = id1 ;
			else ga[i] = id2 ;
		}
	}
}
void calcf() {
	for(int i = 1; i <= n; i++) {
		f[0][i][0] = ga[i] ;
		f[0][i][1] = gb[i] ;
		f[1][i][0] = gb[f[0][i][0]] ;
		f[1][i][1] = ga[f[0][i][1]] ;
	}
	for(int i = 2; i <= len; i++) {
		for(int j = 1; j <= n; j++) {
			if(f[i - 1][j][0])
				f[i][j][0] = f[i - 1][f[i - 1][j][0]][0] ;
			if(f[i - 1][j][1])
				f[i][j][1] = f[i - 1][f[i - 1][j][1]][1] ;
		}
	}
}
void calcd() {
	for(int i = 1; i <= n; i++) {
		da[0][i][0] = dist(i, ga[i]) ;
		da[0][i][1] = 0 ;
		db[0][i][0] = 0 ;
		db[0][i][1] = dist(i, gb[i]) ;
		da[1][i][0] = dist(i, ga[i]) ;
		da[1][i][1] = dist(gb[i], ga[gb[i]]) ;
		db[1][i][0] = dist(ga[i], gb[ga[i]]) ;
		db[1][i][1] = dist(i, gb[i]) ;
	}
	for(int i = 2; i <= len; i++) {
		for(int j = 1; j <= n; j++) {
			if(f[i - 1][j][0] && da[i - 1][j][0] < INF)
				da[i][j][0] = da[i - 1][j][0] + da[i - 1][f[i - 1][j][0]][0] ;
			if(f[i - 1][j][1] && da[i - 1][j][1] < INF)
				da[i][j][1] = da[i - 1][j][1] + da[i - 1][f[i - 1][j][1]][1] ;
			if(f[i - 1][j][0] && db[i - 1][j][0] < INF)
				db[i][j][0] = db[i - 1][j][0] + db[i - 1][f[i - 1][j][0]][0] ;
			if(f[i - 1][j][1] && db[i - 1][j][1] < INF)
				db[i][j][1] = db[i - 1][j][1] + db[i - 1][f[i - 1][j][1]][1] ;
		}
	}
}
int calc(int s, int x, int &xa, int &xb) {
	xa = 0 ;
	xb = 0 ;
	int ans = s ;
	for(int i = len; i >= 0; i--) {
		if(x == 0) break ;
		if(f[i][ans][0]) {
			if(xa + xb + da[i][ans][0] + db[i][ans][0] <= x) {
				xa += da[i][ans][0] ;
				xb += db[i][ans][0] ;
				ans = f[i][ans][0] ;
			}
		}
	}
	return ans ;
}
int main() {
	freopen("test.in","r",stdin) ;
	freopen("test1.out","w",stdout) ;
	cin >> n ;
	len = log2(n) + 1;
	for(int i = 1; i <= n; i++)
		cin >> a[i] ;
	cin >> x0 ;
	cin >> m ;
	memset(da, 0x7f, sizeof(da)) ;
	memset(db, 0x7f, sizeof(db)) ;
	getgg() ;
	calcf() ;
	calcd() ;
	/*
	for(int i=1;i<=n;i++) {
		printf("ga[%d] = %d, gb[%d] = %d\n",i,ga[i],i,gb[i]) ;
	}
	/*                     
		for(int j=0;j<=len;j++)
	f	or(int i=1;i<=n;i++) {
		for(int j=0;j<=len;j++)
			printf("f[%d][%d][0] = %d, f[%d][%d][1] = %d\n",j,i,f[j][i][0],j,i,f[j][i][0]) ;
	}*/
	int u, v ;
	double bzmax = INF ;
	int ansi = 0 ;
	int tmpa, tmpb ;
	for(int i = 1; i <= n; i++) {
		calc(i, x0, tmpa, tmpb) ;
		double ans = (double)((double)tmpa / (double)tmpb) ;
		if(bzmax > ans) {
			bzmax = ans ;
			ansi = i ;
		} else if(bzmax == ans) { //一定要加!!!否则可能小于ans
			if(a[ansi] < a[i]) {
				bzmax = ans ;
				ansi = i ;
			}
		}
	}
	cout << ansi << endl ;
	for(int i = 1; i <= m; i++) {
		cin >> u >> v ;
		calc(u, v, tmpa, tmpb) ;
		cout << tmpa << ' ' << tmpb << endl ;
	}
	return 0 ;
}
2022/8/26 10:00
加载中...