指针BST求调
查看原帖
指针BST求调
519384
Link_Cut_Y楼主2023/6/23 22:53

今天脑抽想写 BST 结果写炸了。

求求路过聚聚帮忙看看!

#include <iostream>
#include <cstring>
#include <cstdio>

using namespace std;

const int INF = 2e9;
struct node {
	node *ls, *rs;
	int val, size, cnt;
	node() {}
	node(int val, int size, int cnt):val(val), size(size), cnt(cnt) {
		ls = rs = nullptr;
	}
	void pushup() {
		size = ((ls != nullptr) ? (ls -> size) : 0) + 
			   ((rs != nullptr) ? (rs -> size) : 0) + 
			   cnt;
	}
}*root;

void insert(node *x, int val) {
	if (x -> val == val) {
		x -> size ++ ; x -> cnt ++ ;
		return;
	}
	if (x -> val < val) {
		if (x -> rs == nullptr) {
			x -> rs = new node(val, 1, 1);
			x -> pushup();
			return;
		}
		insert(x -> rs, val);
	}
	else {
		if (x -> ls == nullptr) {
			x -> ls = new node(val, 1, 1);
			x -> pushup();
			return;
		}
		insert(x -> ls, val);
	}
	x -> pushup();
}

void remove(node *x, int val) {
	if (x -> val == val) {
		x -> cnt -- ; x -> size -- ;
		return;
	}
	if (x -> val < val) remove(x -> rs, val);
	else remove(x -> ls, val);
	x -> pushup();
}

int get_rank(node *x, int val) {
	node *now = x;
	int rank = 0;
	while (true) {
		if (now -> val == val) {
			rank += now -> ls -> size;
			return rank;
		}
		if (now -> val < val) {
			rank += now -> ls -> size + now -> cnt;
			if (now -> rs != nullptr) now = now -> rs;
			else return rank;
		}
		else if (now -> val > val) {
			if (now -> ls != nullptr) now = now -> ls;
			else return rank;
		}
	}
	return rank;
}

int get_val(node *x, int rank) {
	node *now = x;
	while (true) {
		int l_size = (now -> ls == nullptr) ? 0 : now -> ls -> size;
		if (rank >= l_size + 1 and rank <= l_size + now -> cnt)
			return now -> val;
		if (rank <= l_size) now = now -> ls;
		else rank -= l_size + now -> cnt, now = now -> rs;
	}
}

int get_pre(node *x, int val) {
	node *now = x;
	int ans = -0x3f3f3f3f;
	while (true) {
		if (now -> val >= val) {
			if (now -> ls != nullptr) now = now -> ls;
			else return ans;
		}
		else {
			if (now -> cnt > 0) ans = max(ans, now -> val);
			if (now -> rs != nullptr) now = now -> rs;
			else return ans;
		}
	}
}

int get_next(node *x, int val) {
	node *now = x;
	int ans = 0x3f3f3f3f;
	while (true) {
		if (now -> val <= val) {
			if (now -> rs != nullptr) now = now -> rs;
			else return ans;
		}
		else {
			if (now -> cnt > 0) ans = min(ans, now -> val);
			if (now -> ls != nullptr) now = now -> ls;
			else return ans;
		}
	}
}

int main() {
	int n;
	scanf("%d", &n);
	root = new node(-INF, 1, 1);
	insert(root, INF);
	while (n -- ) {
		int op, x;
		scanf("%d%d", &op, &x);		
		if (op == 1) insert(root, x);
		else if (op == 2) remove(root, x);
		else if (op == 3) printf("%d\n", get_rank(root, x) - 1);
		else if (op == 4) printf("%d\n", get_val(root, x + 1));
		else if (op == 5) printf("%d\n", get_pre(root, x));
		else printf("%d\n", get_next(root, x));
	}
	
	return 0;
}
2023/6/23 22:53
加载中...