如题,本人不会 AVL Tree,然后喂给 ChatGPT 题目,结果它写了 AVL Tree,希望哪个大佬帮 AI 调一下谢谢!!
#include <bits/stdc++.h>
using namespace std;
struct Node {
int val;
Node *left;
Node *right;
int height;
int size;
Node(int v) : val(v), left(nullptr), right(nullptr), height(1), size(1) {}
};
class AVLTree {
public:
AVLTree() : root(nullptr) {} ~AVLTree() {
deleteTree(root);
} void insert(int val) {
root = insert(root, val);
} void remove(int val) {
root = remove(root, val);
} int getRank(int val) {
return getRank(root, val);
} int getValByRank(int rank) {
return getValByRank(root, rank);
} int getPredecessor(int val) {
return getPredecessor(root, val);
} int getSuccessor(int val) {
return getSuccessor(root, val);
} private:
Node *root;
Node *insert(Node *node, int val) {
if (node == nullptr) {
return new Node(val);
}
if (val < node->val) {
node->left = insert(node->left, val);
} else {
node->right = insert(node->right, val);
}
node->height = max(getHeight(node->left), getHeight(node->right)) + 1;
node->size = getSize(node->left) + getSize(node->right) + 1;
return balance(node);
} Node *remove(Node *node, int val) {
if (node == nullptr) {
return nullptr;
}
if (val < node->val) {
node->left = remove(node->left, val);
} else if (val > node->val) {
node->right = remove(node->right, val);
} else {
if (node->left == nullptr) {
Node *right = node->right;
delete node;
return right;
} else if (node->right == nullptr) {
Node *left = node->left;
delete node;
return left;
} else {
Node *successor = getMin(node->right);
node->val = successor->val;
node->right = remove(node->right, successor->val);
}
}
node->height = max(getHeight(node->left), getHeight(node->right)) + 1;
node->size = getSize(node->left) + getSize(node->right) + 1;
return balance(node);
} int getRank(Node *node, int val) {
if (node == nullptr) {
return 0;
}
if (val == node->val) {
return getSize(node->left) + 1;
} else if (val < node->val) {
return getRank(node->left, val);
} else {
return getSize(node->left) + 1 + getRank(node->right, val);
}
} int getValByRank(Node *node, int rank) {
if (node == nullptr) {
return -1;
}
int leftSize = getSize(node->left);
if (rank == leftSize + 1) {
return node->val;
} else if (rank <= leftSize) {
return getValByRank(node->left, rank);
} else {
return getValByRank(node->right, rank - leftSize - 1);
}
} int getPredecessor(Node *node, int val) {
if (node == nullptr) {
return -1;
}
if (val <= node->val) {
return getPredecessor(node->left, val);
} else {
int rightPredecessor = getPredecessor(node->right, val);
if (rightPredecessor != -1) {
return rightPredecessor;
} else {
return node->val;
}
}
} int getSuccessor(Node *node, int val) {
if (node == nullptr) {
return -1;
}
if (val >= node->val) {
return getSuccessor(node->right, val);
} else {
int leftSuccessor = getSuccessor(node->left, val);
if (leftSuccessor != -1) {
return leftSuccessor;
} else {
return node->val;
}
}
} Node *getMin(Node *node) {
if (node->left == nullptr) {
return node;
}
return getMin(node->left);
} int getHeight(Node *node) {
if (node == nullptr) {
return 0;
}
return node->height;
} int getSize(Node *node) {
if (node == nullptr) {
return 0;
}
return node->size;
} Node *balance(Node *node) {
int balanceFactor = getHeight(node->left) - getHeight(node->right);
if (balanceFactor > 1) {
if (getHeight(node->left->left) >= getHeight(node->left->right)) {
return rightRotate(node);
} else {
node->left = leftRotate(node->left);
return rightRotate(node);
}
} else if (balanceFactor < -1) {
if (getHeight(node->right->right) >= getHeight(node->right->left)) {
return leftRotate(node);
} else {
node->right = rightRotate(node->right);
return leftRotate(node);
}
}
return node;
} Node *rightRotate(Node *node) {
Node *left = node->left;
node->left = left->right;
left->right = node;
node->height = max(getHeight(node->left), getHeight(node->right)) + 1;
left->height = max(getHeight(left->left), getHeight(left->right)) + 1;
node->size = getSize(node->left) + getSize(node->right) + 1;
left->size = getSize(left->left) + getSize(left->right) + 1;
return left;
} Node *leftRotate(Node *node) {
Node *right = node->right;
node->right = right->left;
right->left = node;
node->height = max(getHeight(node->left), getHeight(node->right)) + 1;
right->height = max(getHeight(right->left), getHeight(right->right)) + 1;
node->size = getSize(node->left) + getSize(node->right) + 1;
right->size = getSize(right->left) + getSize(right->right) + 1;
return right;
} void deleteTree(Node *node) {
if (node == nullptr) {
return;
}
deleteTree(node->left);
deleteTree(node->right);
delete node;
}
};
AVLTree tree;
int main() {
int n;
cin>>n;
int opt;
while(n--){
cin>>opt;
switch(opt){
case 1:{
int x;
cin>>x;
tree.insert(x);
break;
}
case 2:{
int x;
cin>>x;
tree.remove(x);
break;
}
case 3:{
int x;
cin>>x;
cout<<tree.getRank(x)<<endl;
break;
}
case 4:{
int rnk;
cin>>rnk;
cout<<tree.getValByRank(rnk)<<endl;
break;
}
case 5:{
int x;
cin>>x;
cout<<tree.getPredecessor(x)<<endl;
break;
}
case 6:{
int x;
cin>>x;
cout<<tree.getSuccessor(x)<<endl;
break;
}
}
}
return 0;
}