11.平衡二叉树.md 14 KB

11. 平衡二叉树

概念

平衡二叉树(AVL 树,Adelson-Velsky and Landis) 是一种自平衡的二叉排序树,它在 BST 性质(左<根<右)的基础上,额外保证树在插入、删除后仍然保持接近平衡,从而避免退化成链表。

平衡因子(Balance Factor, BF):某结点的平衡因子 = 该结点左子树高度 − 右子树高度(或反过来定义,两种约定等价,本文用 左高−右高)。

平衡条件:AVL 树要求每个结点的平衡因子绝对值不超过 1,即 |BF| <= 1。只要任意结点 |BF| > 1,就说明失衡,需要旋转调整。

为什么需要平衡:普通 BST 在最坏情况下(如有序插入)会退化为一条链,树高为 n,查找退化为 O(n)。AVL 树通过严格的平衡条件把树高约束在 O(log n),从而保证查找、插入、删除都稳定在 O(log n)。

四种旋转:当插入新结点导致某个祖先结点失衡时,根据"新结点在失衡结点哪个方向"分为四种情况,通过旋转恢复平衡:

  • LL(左左):在失衡结点左孩子的左子树插入 → 对失衡结点做一次右旋
  • RR(右右):在失衡结点右孩子的右子树插入 → 对失衡结点做一次左旋
  • LR(左右):在失衡结点左孩子的右子树插入 → 先对左孩子做左旋,再对失衡结点做右旋(两次旋转)。
  • RL(右左):在失衡结点右孩子的左子树插入 → 先对右孩子做右旋,再对失衡结点做左旋(两次旋转)。

以失衡结点 A(左高-右高>1 为例)的旋转示意:

LL:A 左子 B 的左子树过高,右旋 A
      A                B
     / \              / \
    B   T3    =>     T1  A
   / \                  / \
  T1  T2               T2  T3

RR:A 右子 B 的右子树过高,左旋 A
      A                 B
     / \               / \
    T1  B      =>     A   T3
       / \           / \
      T2  T3        T1  T2

LR:A 左子 B 的右子 C 过高,先左旋 B,再右旋 A
      A               A              C
     / \             / \            / \
    B   T4   =>     C   T4   =>    B   A
   / \             / \            / \ / \
  T1  C           B   T3         T1 T2 T3 T4
     / \         / \
    T2  T3      T1  T2

RL:A 右子 B 的左子 C 过高,先右旋 B,再左旋 A
      A               A                C
     / \             / \              / \
    T1  B     =>    T1  C     =>     A   B
       / \             / \          / \ / \
      C   T4          T2  B        T1 T2 T3 T4
     / \                 / \
    T2  T3              T3  T4

适用场景:需要频繁查找且插入/删除也频繁、对最坏性能有要求的动态查找表(如数据库索引、内存中的有序集合)。相比普通 BST,它牺牲少量插入/删除常数,换来稳定的 O(log n)。

核心操作

  • 插入:先按 BST 规则插入为叶结点,然后从插入点回溯到根,沿途更新各结点高度,一旦发现某个结点失衡(|BF|>1),就根据四种情况做相应旋转。
  • 求高度:结点高度 = max(左高, 右高) + 1;空结点高度为 0(或 -1,本文约定空为 0)。高度用于计算平衡因子。
  • 更新平衡因子:每次插入/旋转后都要重新计算受影响结点的高度与平衡因子。
  • 中序遍历:左-根-右,得到的仍是递增有序序列。
  • 删除(思路,不要求实现):按 BST 删除规则删除结点后,同样从删除点回溯调整平衡,可能需要多次旋转(可能向上传播失衡),比插入更复杂。

复杂度分析

操作 时间复杂度 空间复杂度
查找 O(log n) 递归 O(log n),迭代 O(1)
插入 O(log n) 递归 O(log n)
删除 O(log n) 递归 O(log n)
中序遍历 O(n) 递归 O(log n)

为什么:AVL 树保证任意结点左右子树高度差不超过 1,因此树高始终为 O(log n)。查找沿路径下降一层 O(1),故 O(log n);插入和删除在 BST 基础上沿路径回溯做旋转,每层 O(1),总共 O(log n) 次旋转;中序遍历访问全部 n 个结点,为 O(n)。与普通 BST 相比,关键区别在于 AVL 不会退化为 O(n)

语言实现

下面是 4 种语言的完整 AVL 实现,均演示:插入若干元素、中序遍历输出、输出根及部分结点的平衡因子(高度差)。删除仅说明思路不实现。

C

#include <stdio.h>
#include <stdlib.h>

// AVL 结点:含高度字段
typedef struct Node {
    int data;
    struct Node *left, *right;
    int height;   // 当前结点子树高度
} Node;

int height(Node *n) {
    return n == NULL ? 0 : n->height;
}

int max(int a, int b) { return a > b ? a : b; }

// 计算平衡因子:左高 - 右高
int balanceFactor(Node *n) {
    return n == NULL ? 0 : height(n->left) - height(n->right);
}

Node *newNode(int data) {
    Node *n = (Node *)malloc(sizeof(Node));
    n->data = data;
    n->left = n->right = NULL;
    n->height = 1;
    return n;
}

// 右旋(处理 LL 失衡)
Node *rotateRight(Node *y) {
    Node *x = y->left;
    Node *T2 = x->right;
    x->right = y;
    y->left = T2;
    y->height = max(height(y->left), height(y->right)) + 1;
    x->height = max(height(x->left), height(x->right)) + 1;
    return x;
}

// 左旋(处理 RR 失衡)
Node *rotateLeft(Node *x) {
    Node *y = x->right;
    Node *T2 = y->left;
    y->left = x;
    x->right = T2;
    x->height = max(height(x->left), height(x->right)) + 1;
    y->height = max(height(y->left), height(y->right)) + 1;
    return y;
}

// 插入并自动平衡
Node *insert(Node *n, int data) {
    if (n == NULL) return newNode(data);

    if (data < n->data)      n->left = insert(n->left, data);
    else if (data > n->data) n->right = insert(n->right, data);
    else                     return n;  // 重复值不插入

    n->height = max(height(n->left), height(n->right)) + 1;
    int bf = balanceFactor(n);

    // 四种失衡情形
    if (bf > 1 && data < n->left->data)      return rotateRight(n);          // LL
    if (bf < -1 && data > n->right->data)    return rotateLeft(n);           // RR
    if (bf > 1 && data > n->left->data) {    // LR
        n->left = rotateLeft(n->left);
        return rotateRight(n);
    }
    if (bf < -1 && data < n->right->data) {  // RL
        n->right = rotateRight(n->right);
        return rotateLeft(n);
    }
    return n;
}

void inorder(Node *n) {
    if (n == NULL) return;
    inorder(n->left);
    printf("%d ", n->data);
    inorder(n->right);
}

int main() {
    Node *root = NULL;
    // 故意按会触发失衡的顺序插入
    int a[] = {10, 20, 30, 40, 50, 25};
    for (int i = 0; i < 6; i++)
        root = insert(root, a[i]);

    printf("中序遍历(有序): ");
    inorder(root);
    printf("\n");

    printf("根结点值: %d, 平衡因子: %d\n",
           root->data, balanceFactor(root));
    printf("左子树高度: %d, 右子树高度: %d\n",
           height(root->left), height(root->right));
    printf("整棵树高度: %d\n", height(root));
    return 0;
}

C++

#include <iostream>
#include <algorithm>
using namespace std;

struct Node {
    int data, height;
    Node *left, *right;
    Node(int d) : data(d), height(1), left(nullptr), right(nullptr) {}
};

int height(Node *n) { return n ? n->height : 0; }
int bf(Node *n) { return n ? height(n->left) - height(n->right) : 0; }

Node *rotateRight(Node *y) {
    Node *x = y->left;
    Node *T2 = x->right;
    x->right = y;
    y->left = T2;
    y->height = max(height(y->left), height(y->right)) + 1;
    x->height = max(height(x->left), height(x->right)) + 1;
    return x;
}

Node *rotateLeft(Node *x) {
    Node *y = x->right;
    Node *T2 = y->left;
    y->left = x;
    x->right = T2;
    x->height = max(height(x->left), height(x->right)) + 1;
    y->height = max(height(y->left), height(y->right)) + 1;
    return y;
}

Node *insert(Node *n, int data) {
    if (!n) return new Node(data);
    if (data < n->data) n->left = insert(n->left, data);
    else if (data > n->data) n->right = insert(n->right, data);
    else return n;

    n->height = max(height(n->left), height(n->right)) + 1;
    int balance = bf(n);

    if (balance > 1 && data < n->left->data) return rotateRight(n);   // LL
    if (balance < -1 && data > n->right->data) return rotateLeft(n);  // RR
    if (balance > 1 && data > n->left->data) {                        // LR
        n->left = rotateLeft(n->left);
        return rotateRight(n);
    }
    if (balance < -1 && data < n->right->data) {                      // RL
        n->right = rotateRight(n->right);
        return rotateLeft(n);
    }
    return n;
}

void inorder(Node *n) {
    if (!n) return;
    inorder(n->left);
    cout << n->data << " ";
    inorder(n->right);
}

int main() {
    Node *root = nullptr;
    for (int x : {10, 20, 30, 40, 50, 25})
        root = insert(root, x);

    cout << "中序遍历(有序): ";
    inorder(root);
    cout << endl;

    cout << "根结点值: " << root->data
         << ", 平衡因子: " << bf(root) << endl;
    cout << "左子树高度: " << height(root->left)
         << ", 右子树高度: " << height(root->right) << endl;
    cout << "整棵树高度: " << height(root) << endl;
    return 0;
}

Java

public class AVL {

    static class Node {
        int data, height;
        Node left, right;
        Node(int d) { data = d; height = 1; }
    }

    private Node root;

    private int height(Node n) { return n == null ? 0 : n.height; }
    private int max(int a, int b) { return a > b ? a : b; }
    private int bf(Node n) { return n == null ? 0 : height(n.left) - height(n.right); }

    // 右旋(LL)
    private Node rotateRight(Node y) {
        Node x = y.left;
        Node T2 = x.right;
        x.right = y;
        y.left = T2;
        y.height = max(height(y.left), height(y.right)) + 1;
        x.height = max(height(x.left), height(x.right)) + 1;
        return x;
    }

    // 左旋(RR)
    private Node rotateLeft(Node x) {
        Node y = x.right;
        Node T2 = y.left;
        y.left = x;
        x.right = T2;
        x.height = max(height(x.left), height(x.right)) + 1;
        y.height = max(height(y.left), height(y.right)) + 1;
        return y;
    }

    public void insert(int data) { root = insert(root, data); }
    private Node insert(Node n, int data) {
        if (n == null) return new Node(data);
        if (data < n.data) n.left = insert(n.left, data);
        else if (data > n.data) n.right = insert(n.right, data);
        else return n;

        n.height = max(height(n.left), height(n.right)) + 1;
        int balance = bf(n);

        if (balance > 1 && data < n.left.data) return rotateRight(n);    // LL
        if (balance < -1 && data > n.right.data) return rotateLeft(n);   // RR
        if (balance > 1 && data > n.left.data) {                         // LR
            n.left = rotateLeft(n.left);
            return rotateRight(n);
        }
        if (balance < -1 && data < n.right.data) {                       // RL
            n.right = rotateRight(n.right);
            return rotateLeft(n);
        }
        return n;
    }

    public void inorder() { inorder(root); System.out.println(); }
    private void inorder(Node n) {
        if (n == null) return;
        inorder(n.left);
        System.out.print(n.data + " ");
        inorder(n.right);
    }

    public static void main(String[] args) {
        AVL tree = new AVL();
        int[] a = {10, 20, 30, 40, 50, 25};
        for (int x : a) tree.insert(x);

        System.out.print("中序遍历(有序): ");
        tree.inorder();

        System.out.println("根结点值: " + tree.root.data
                + ", 平衡因子: " + tree.bf(tree.root));
        System.out.println("左子树高度: " + tree.height(tree.root.left)
                + ", 右子树高度: " + tree.height(tree.root.right));
        System.out.println("整棵树高度: " + tree.height(tree.root));
    }
}

Python

class Node:
    """AVL 结点,含高度字段"""
    def __init__(self, data):
        self.data = data
        self.left = None
        self.right = None
        self.height = 1


def height(n):
    return 0 if n is None else n.height


def bf(n):
    """平衡因子 = 左高 - 右高"""
    return 0 if n is None else height(n.left) - height(n.right)


def rotate_right(y):
    """右旋(LL)"""
    x = y.left
    t2 = x.right
    x.right = y
    y.left = t2
    y.height = max(height(y.left), height(y.right)) + 1
    x.height = max(height(x.left), height(x.right)) + 1
    return x


def rotate_left(x):
    """左旋(RR)"""
    y = x.right
    t2 = y.left
    y.left = x
    x.right = t2
    x.height = max(height(x.left), height(x.right)) + 1
    y.height = max(height(y.left), height(y.right)) + 1
    return y


def insert(n, data):
    """AVL 插入并自动平衡"""
    if n is None:
        return Node(data)
    if data < n.data:
        n.left = insert(n.left, data)
    elif data > n.data:
        n.right = insert(n.right, data)
    else:
        return n

    n.height = max(height(n.left), height(n.right)) + 1
    balance = bf(n)

    if balance > 1 and data < n.left.data:            # LL
        return rotate_right(n)
    if balance < -1 and data > n.right.data:          # RR
        return rotate_left(n)
    if balance > 1 and data > n.left.data:            # LR
        n.left = rotate_left(n.left)
        return rotate_right(n)
    if balance < -1 and data < n.right.data:          # RL
        n.right = rotate_right(n.right)
        return rotate_left(n)
    return n


def inorder(n):
    if n is None:
        return []
    return inorder(n.left) + [n.data] + inorder(n.right)


if __name__ == "__main__":
    root = None
    for x in [10, 20, 30, 40, 50, 25]:
        root = insert(root, x)

    print("中序遍历(有序):", inorder(root))
    print("根结点值:", root.data, ", 平衡因子:", bf(root))
    print("左子树高度:", height(root.left), ", 右子树高度:", height(root.right))
    print("整棵树高度:", height(root))