AVL树60分求调
查看原帖
AVL树60分求调
72922
iterator_traits楼主2023/8/1 22:49

60分,错的询问好像都是3和4:

#include <iostream>
#include <cassert>
#include <climits>
#include <array>

struct Node {
    int num = 0, cnt = 1, h = 1, siz = 1;
    Node *ls = nullptr, *rs = nullptr;

    void update();

    // Balance defined as h(l) - h(r)
    int balance();
};

constexpr int MAXN = 100'005;
std::array<Node, MAXN> nodes;
int nodes_cnt = 0;

int get_h(Node* p) {
    return p ? p->h : 0;
}

int get_siz(Node* p) {
    return p ? p->siz : 0;
}

void Node::update() {
    siz = cnt + get_siz(ls) + get_siz(rs);
    h = 1 + std::max(get_h(ls), get_h(rs));
}

int Node::balance() {
    return get_h(ls) - get_h(rs);
}

// Handles the left left situation
void left_plus(Node*& curr) {
    auto ls = curr->ls;
    curr->ls = ls->rs;
    ls->rs = curr;
    curr->update();
    ls->update();
    curr = ls;
}

// Handles the right right situation
void right_plus(Node*& curr) {
    auto rs = curr->rs;
    curr->rs = rs->ls;
    rs->ls = curr;
    curr->update();
    rs->update();
    curr = rs;
}

// Handles left right
void left_right(Node*& curr) {
    right_plus(curr->ls);
    left_plus(curr);
}

// Handles right left
void right_left(Node*& curr) {
    left_plus(curr->rs);
    right_plus(curr);
}

// Maintain the AVL properties at the node and updates info
void maintain(Node*& curr) {
    curr->update();
    if (curr->balance() > 1) {
        if (curr->balance() >= 0)
            left_plus(curr);
        else
            left_right(curr);
    } else if (curr->balance() < -1) {
        if (curr->balance() <= 0)
            right_plus(curr);
        else
            right_left(curr);
    }
}

// Inserts number x into subtree
void insert(int x, Node*& curr) {
    if (!curr) {
        curr = &nodes[++nodes_cnt];
        curr->num = x;
        return;
    }
    if (x == curr->num) {
        curr->cnt++;
        curr->update();
        return;
    }
    if (x < curr->num)
        insert(x, curr->ls);
    else if (x > curr->num)
        insert(x, curr->rs);
    maintain(curr);
}

// Finds the max element in subtree, returns a copy of the node.
Node find_max(Node* curr) {
    while (curr->rs)
        curr = curr->rs;
    return *curr;
}

// Removes x from subtree mul times
void remove(int x, Node*& curr, int mul = 1) {
    if (!curr)
        return;
    if (x == curr->num) {
        curr->cnt -= mul;
        if (curr->cnt) {
            // No structural changes for curr
            curr->update();
            return;
        }
        // Now we know curr->cnt == 0
        if (!curr->ls && !curr->rs)
            // Leaf node
            curr = nullptr;
        else if (curr->ls && !curr->rs)
            curr = curr->ls;
        else if (!curr->ls && curr->rs)
            curr = curr->rs;
        else {
            auto lmax = find_max(curr->ls);
            curr->num = lmax.num;
            curr->cnt = lmax.cnt;
            remove(lmax.num, curr->ls, lmax.cnt);
        }
        if (curr)
            maintain(curr);
        return;
    } else if (x < curr->num) {
        remove(x, curr->ls);
    } else {
        remove(x, curr->rs);
    }
    maintain(curr);
}

// Get the cnt of numbers smaller than x
int small_cnt(int x, Node* curr) {
    if (!curr)
        return 0;
    return curr->num < x ? get_siz(curr->ls) + curr->cnt + small_cnt(x, curr->rs) : small_cnt(x, curr->ls);
}

// Get by rank
int by_rank(int rnk, Node* curr) {
    if (rnk <= get_siz(curr->ls))
        return by_rank(rnk, curr->ls);
    if (get_siz(curr->ls) + 1 <= rnk && rnk <= get_siz(curr->ls) + curr->cnt)
        return curr->num;
    if (rnk > get_siz(curr->ls) + curr->cnt)
        return by_rank(rnk - get_siz(curr->ls) - curr->cnt, curr->rs);
}

// Get predecessor of x
int get_pred(int x, Node* curr) {
    if (!curr)
        return INT_MIN;
    return curr->num >= x ? get_pred(x, curr->ls) : std::max(curr->num, get_pred(x, curr->rs));
}

// Get successor of x
int get_suc(int x, Node* curr) {
    if (!curr)
        return INT_MAX;
    return curr->num <= x ? get_suc(x, curr->rs) : std::min(curr->num, get_suc(x, curr->ls));
}

int main() {
    std::ios::sync_with_stdio(false);
    std::cin.tie(nullptr);
    std::cout.tie(nullptr);
    int n;
    Node* root = nullptr;
    std::cin >> n;
    while (n--) {
        int op, x;
        std::cin >> op >> x;
        switch (op) {
        case 1:
            insert(x, root);
            break;
        case 2:
            remove(x, root);
            break;
        case 3:
            std::cout << small_cnt(x, root) + 1 << '\n';
            break;
        case 4:
            std::cout << by_rank(x, root) << '\n';
            break;
        case 5:
            std::cout << get_pred(x, root) << '\n';
            break;
        case 6:
            std::cout << get_suc(x, root) << '\n';
            break;
        default:
            return 2;
        }
    }
}
2023/8/1 22:49
加载中...