Implementation of the Chtholly Tree Data Structure

The Chtholly Tree organizes consecutive array segments with identical values into nodes, stored in a sorted set. Each node represents a range [l, r] with a mutable value.

struct Segment {
    int left, right;
    mutable long long value;
    Segment(int l = 0, int r = 0, long long v = 0) : left(l), right(r), value(v) {}
    bool operator<(const Segment& other) const {
        return left < other.left;
    }
};
using Iterator = set<Segment>::iterator;
set<Segment> segments;

The split operation divides a node at a specified position.

Iterator split(int pos) {
    Iterator it = segments.lower_bound(Segment(pos));
    if (it != segments.end() && it->left == pos) return it;
    --it;
    int L = it->left, R = it->right;
    long long V = it->value;
    segments.erase(it);
    segments.insert(Segment(L, pos - 1, V));
    return segments.insert(Segment(pos, R, V)).first;
}

Range assignment merges nodes to maintain efficiency.

void assign(int l, int r, long long val) {
    Iterator right_it = split(r + 1);
    Iterator left_it = split(l);
    segments.erase(left_it, right_it);
    segments.insert(Segment(l, r, val));
}

Range addition modifies values across multiple nodes.

void add(int l, int r, long long delta) {
    Iterator right_it = split(r + 1);
    Iterator left_it = split(l);
    for (Iterator it = left_it; it != right_it; ++it)
        it->value += delta;
}

To find the k-th smallest value, collect and sort node values.

long long kth_smallest(int l, int r, int k) {
    Iterator right_it = split(r + 1);
    Iterator left_it = split(l);
    vector<pair<long long, int>> values;
    for (Iterator it = left_it; it != right_it; ++it)
        values.emplace_back(it->value, it->right - it->left + 1);
    sort(values.begin(), values.end());
    for (auto& [val, cnt] : values) {
        k -= cnt;
        if (k <= 0) return val;
    }
    return -1;
}

Computing the sum of powers uses modular exponentiation.

long long mod_pow(long long base, long long exp, long long mod) {
    long long result = 1;
    base %= mod;
    while (exp) {
        if (exp & 1) result = (result * base) % mod;
        base = (base * base) % mod;
        exp >>= 1;
    }
    return result;
}

long long power_sum(int l, int r, long long exp, long long mod) {
    Iterator right_it = split(r + 1);
    Iterator left_it = split(l);
    long long total = 0;
    for (Iterator it = left_it; it != right_it; ++it) {
        int length = it->right - it->left + 1;
        total = (total + length * mod_pow(it->value, exp, mod)) % mod;
    }
    return total;
}

For operations like copying ranges, store nodes temporarily.

Segment buffer[100010];
void copy_range(int src_l, int src_r, int dst_l, int dst_r) {
    int idx = 0;
    Iterator src_end = split(src_r + 1);
    Iterator src_start = split(src_l);
    for (Iterator it = src_start; it != src_end; ++it)
        buffer[++idx] = Segment(it->left, it->right, it->value);
    Iterator dst_end = split(dst_r + 1);
    Iterator dst_start = split(dst_l);
    segments.erase(dst_start, dst_end);
    for (int i = 1; i <= idx; ++i) {
        int offset = buffer[i].left - src_l;
        segments.insert(Segment(dst_l + offset, dst_l + offset + buffer[i].right - buffer[i].left, buffer[i].value));
    }
}

Swapping ranges involves handling both segments.

Segment temp1[100010], temp2[100010];
void swap_ranges(int l1, int r1, int l2, int r2) {
    if (l1 > l2) swap(l1, l2), swap(r1, r2);
    int cnt1 = 0, cnt2 = 0;
    Iterator end1 = split(r1 + 1), start1 = split(l1);
    for (Iterator it = start1; it != end1; ++it)
        temp1[++cnt1] = Segment(it->left, it->right, it->value);
    segments.erase(start1, end1);
    Iterator end2 = split(r2 + 1), start2 = split(l2);
    for (Iterator it = start2; it != end2; ++it)
        temp2[++cnt2] = Segment(it->left, it->right, it->value);
    segments.erase(start2, end2);
    for (int i = 1; i <= cnt1; ++i) {
        int shift = temp1[i].left - l1;
        segments.insert(Segment(l2 + shift, l2 + shift + temp1[i].right - temp1[i].left, temp1[i].value));
    }
    for (int i = 1; i <= cnt2; ++i) {
        int shift = temp2[i].left - l2;
        segments.insert(Segment(l1 + shift, l1 + shift + temp2[i].right - temp2[i].left, temp2[i].value));
    }
}

Reversing a range trensforms node boundaries.

void reverse_range(int l, int r) {
    int idx = 0;
    Iterator right_it = split(r + 1);
    Iterator left_it = split(l);
    for (Iterator it = left_it; it != right_it; ++it)
        buffer[++idx] = Segment(it->left, it->right, it->value);
    segments.erase(left_it, right_it);
    for (int i = 1; i <= idx; ++i) {
        int new_left = l + r - buffer[i].right;
        int new_right = l + r - buffer[i].left;
        segments.insert(Segment(new_left, new_right, buffer[i].value));
    }
}

For queries involving color constraints, use a two-pointer approach with frequency tracking.

int color_count[101], distinct_colors = 0;
void add_color(int col) {
    if (++color_count[col] == 1) ++distinct_colors;
}
void remove_color(int col) {
    if (--color_count[col] == 0) --distinct_colors;
}

long long min_sum_with_all_colors(int l, int r, int total_colors) {
    if (total_colors == 1) return query_min(l, r);
    Iterator right_it = split(r + 1), left_it = split(l), current = left_it;
    memset(color_count, 0, sizeof(color_count));
    distinct_colors = 0;
    long long answer = INF;
    while (current != right_it) {
        add_color(current->value);
        while (distinct_colors == total_colors) {
            answer = min(answer, query_sum(left_it->right, current->left));
            remove_color(left_it->value);
            ++left_it;
        }
        ++current;
    }
    return answer == INF ? -1 : answer;
}

Finding the maximum sum without duplicate colors uses a similar technique.

long long max_sum_no_duplicates(int l, int r) {
    memset(color_count, 0, sizeof(color_count));
    Iterator right_it = split(r + 1), left_it = split(l), current = left_it;
    long long answer = query_max(l, r);
    while (current != right_it) {
        ++color_count[current->value];
        while (left_it != current && color_count[current->value] > 1)
            --color_count[left_it->value], ++left_it;
        if (current != left_it)
            answer = max(answer, query_sum(left_it->right, current->left));
        if (current->left != current->right) {
            while (left_it != current)
                --color_count[left_it->value], ++left_it;
        }
        ++current;
    }
    return answer;
}

Tags: data-structures chtholly-tree competitive-programming cplusplus

Posted on Thu, 01 Oct 2026 16:29:25 +0000 by ryankentp