バーストソートを使用する

バーストソート (burstsort) は、キャッシュ効率を意識した文字列整列向けのアルゴリズムである。共通接頭辞をトライ木で管理し、まだ細分化していない末尾部分をバケットへ集める。バケットが閾値を超えたら「バースト」して下位桁へトライを伸ばし、小さなバケットだけ挿入ソートなどで仕上げる。

本記事では整数キーを十進桁ごとに扱う簡略版を示す。トライソートと同様に最上位桁優先の区分付けだが、最初から木全体を構築するのではなく、バケットが膨らんだときだけ子ノードへ展開する点が異なる。

  1. バケットへの挿入: 各要素を根のバケットへ追加する。ノードがすでにトライ化していれば、現行桁 d = (x / exp) % 10 に従い子へ降り、exp を 10 で割って繰り返す。
  2. バースト: 葉バケットの要素数が閾値 B(ここでは 16、デモでは 4)を超え、かつ exp > 0 なら、バケット内の全要素を現行桁で 10 分割し子ノードへ再配分する。
  3. 葉の整列: これ以上桁を分けられないノード(exp = 0)や、閾値以下に収まったバケットは挿入ソートで昇順にする。
  4. 収集: 子 0..9 の順に深さ優先走査し、整列済みバケットを左から連結して出力する。
procedure burst_sort(A)
  if length(A) = 0 then return
  maxVal = maximum(A)
  exp = highest_digit_weight(maxVal)
  root = burst node with empty bucket
  for each x in A
    burst_insert(root, x, exp)
  out = empty list
  burst_collect(root, out)   // children 0..9, then sort each bucket
  A = out

桁数を w とすると挿入と収集はともに O(n · w) であり、補助記憶域は O(n) となる。葉バケットを挿入ソートで仕上げるため安定ソートである。

以下のデモでは 2 桁の整数 20 個の乱数を使い、閾値を 4 に固定している。「シャッフル」で別の並びに差し替えられます。

類似アルゴリズムとの相違点

トライソートは全キーを最初から桁ごとの木へ挿入する。基数ソートの最上位桁優先版と区分は同型で、バーストソートはバケットが閾値を超えたときだけ下位桁へ展開する点が異なる。

計算時間量および空間計算量を計測する

Size Average time Maximum time Average memory Maximum memory
256 0.000009 0.000054 8 8
512 0.000017 0.000069 16 16
1024 0.000034 0.000072 32 32
2048 0.000066 0.000131 64 64
4096 0.000133 0.000205 129 129
8192 0.000280 0.000505 259 259
16384 0.000588 0.001185 517 517
32768 0.001183 0.002515 1035 1035
65536 0.002480 0.004254 2070 2070
131072 0.005392 0.011395 4141 4141
262144 0.011281 0.025826 8283 8283
計測に使用したコードを表示する

set -euo pipefail

WORKDIR="$(mktemp -d)"
trap 'rm -rf "$WORKDIR"' EXIT

cat > "$WORKDIR/Dockerfile" <<'EOF'
FROM rust:1.95.0

WORKDIR /app

RUN mkdir -p src

RUN cat > Cargo.toml <<'CARGO'
[package]
name = "rust-benchmark"
version = "0.1.0"
edition = "2021"

[profile.release]
lto = true
codegen-units = 1
panic = "abort"
CARGO

RUN cat > src/main.rs <<'RUST'
use std::{
    alloc::{GlobalAlloc, Layout, System},
    env,
    process::Command,
    sync::atomic::{AtomicUsize, Ordering},
    time::{Duration, Instant},
};

/// Counts live heap bytes and the high-water mark so auxiliary sort buffers
/// (swap Vecs, etc.) are measured as explicit heap growth during the sort.
struct TrackingAllocator;

static LIVE_BYTES: AtomicUsize = AtomicUsize::new(0);
static PEAK_BYTES: AtomicUsize = AtomicUsize::new(0);

fn record_alloc(size: usize) {
    let live = LIVE_BYTES.fetch_add(size, Ordering::Relaxed) + size;
    PEAK_BYTES.fetch_max(live, Ordering::Relaxed);
}

unsafe impl GlobalAlloc for TrackingAllocator {
    unsafe fn alloc(&self, layout: Layout) -> *mut u8 {
        let ptr = System.alloc(layout);
        if !ptr.is_null() {
            record_alloc(layout.size());
        }
        ptr
    }

    unsafe fn dealloc(&self, ptr: *mut u8, layout: Layout) {
        LIVE_BYTES.fetch_sub(layout.size(), Ordering::Relaxed);
        System.dealloc(ptr, layout);
    }

    unsafe fn alloc_zeroed(&self, layout: Layout) -> *mut u8 {
        let ptr = System.alloc_zeroed(layout);
        if !ptr.is_null() {
            record_alloc(layout.size());
        }
        ptr
    }

    unsafe fn realloc(&self, ptr: *mut u8, layout: Layout, new_size: usize) -> *mut u8 {
        let new_ptr = System.realloc(ptr, layout, new_size);
        if !new_ptr.is_null() {
            LIVE_BYTES.fetch_sub(layout.size(), Ordering::Relaxed);
            record_alloc(new_size);
        }
        new_ptr
    }
}

#[global_allocator]
static GLOBAL: TrackingAllocator = TrackingAllocator;
const MIN_POWER: u32 = 8;
const MAX_POWER: u32 = 18;
const RUNS: usize = 8192;
fn insertion_sort(a: &mut [usize]) {
    for i in 1..a.len() {
        let mut j = i;
        while j > 0 && a[j - 1] > a[j] {
            a.swap(j - 1, j);
            j -= 1;
        }
    }
}



const BURST_THRESHOLD: usize = 16;

#[derive(Default)]
struct BurstNode {
    children: [Option<Box<BurstNode>>; 10],
    bucket: Vec<usize>,
}

impl BurstNode {
    fn is_trie(&self) -> bool {
        self.children.iter().any(|c| c.is_some())
    }
}

fn max_digit_exp(max: usize) -> usize {
    let mut exp = 1usize;
    while exp.saturating_mul(10) <= max && exp <= usize::MAX / 10 {
        exp *= 10;
    }
    exp
}

fn burst_insert(node: &mut BurstNode, value: usize, exp: usize) {
    if node.is_trie() {
        if exp == 0 {
            node.bucket.push(value);
            return;
        }
        let digit = (value / exp) % 10;
        let child = node
            .children[digit]
            .get_or_insert_with(|| Box::new(BurstNode::default()));
        burst_insert(child, value, exp / 10);
        return;
    }

    node.bucket.push(value);
    if node.bucket.len() > BURST_THRESHOLD && exp > 0 {
        let exp_now = exp;
        let items = std::mem::take(&mut node.bucket);
        for v in items {
            let digit = (v / exp_now) % 10;
            let child = node
                .children[digit]
                .get_or_insert_with(|| Box::new(BurstNode::default()));
            burst_insert(child, v, exp_now / 10);
        }
    }
}

fn burst_collect(node: &BurstNode, out: &mut Vec<usize>) {
    if node.is_trie() {
        for child in node.children.iter().flatten() {
            burst_collect(child, out);
        }
    }
    if !node.bucket.is_empty() {
        let mut bucket = node.bucket.clone();
        insertion_sort(&mut bucket);
        out.extend(bucket);
    }
}

fn burst_sort(a: &mut [usize]) {
    if a.is_empty() {
        return;
    }

    let max = *a.iter().max().unwrap();
    let exp = max_digit_exp(max);
    let mut root = BurstNode::default();

    for &x in a.iter() {
        burst_insert(&mut root, x, exp);
    }

    let mut out = Vec::with_capacity(a.len());
    burst_collect(&root, &mut out);
    a.copy_from_slice(&out);
}


fn benchmark_sort(array: &mut [usize]) {

    burst_sort(array);

}

fn is_non_decreasing(a: &[usize]) -> bool {
    a.windows(2).all(|w| w[0] <= w[1])
}

fn same_multiset(a: &[usize], b: &[usize]) -> bool {
    if a.len() != b.len() {
        return false;
    }

    let mut left = a.to_vec();
    let mut right = b.to_vec();
    left.sort_unstable();
    right.sort_unstable();
    left == right
}

fn check_correctness_case(label: &str, mut input: Vec<usize>) {
    let original = input.clone();

    benchmark_sort(&mut input);

    if !is_non_decreasing(&input) {
        panic!("correctness case {}: output is not sorted", label);
    }

    if !same_multiset(&input, &original) {
        panic!("correctness case {}: elements were lost or added", label);
    }
}

fn few_unique_values(size: usize, unique: usize, seed: u64) -> Vec<usize> {
    let mut state = seed;

    (0..size)
        .map(|_| {
            state ^= state << 13;
            state ^= state >> 7;
            state ^= state << 17;
            (state as usize % unique) + 1
        })
        .collect()
}

fn run_correctness_checks() {
    check_correctness_case("empty", vec![]);
    check_correctness_case("single", vec![42]);
    check_correctness_case("duplicates", vec![3, 1, 3, 2, 1, 2]);
    check_correctness_case("sorted", vec![1, 2, 3, 4, 5]);
    check_correctness_case("reverse", vec![5, 4, 3, 2, 1]);
    check_correctness_case("all_equal", vec![7, 7, 7, 7]);
    check_correctness_case("skewed_range", vec![1_000_000, 2, 1_000_001, 1, 999_999]);
    // Static-buffer Grail skips the in-buffer build when key collection is sparse
    // (ideal_buffer = false). Exercising that path catches regressions in buffer gating.
    check_correctness_case(
        "few_keys_len16",
        vec![2, 2, 2, 2, 2, 2, 2, 2, 4, 3, 1, 2, 3, 4, 1, 4],
    );
    // Seed 0 is a fixed point of the xorshift below, so it would degenerate into
    // yet another all-equal case instead of a 4-value mix. Start at 1.
    for seed in 1..=32 {
        check_correctness_case(
            &format!("few_keys_len32_seed_{seed}"),
            few_unique_values(32, 4, seed),
        );
    }
}


fn shuffled(size: usize, seed: u64) -> Vec<usize> {
    let mut v: Vec<usize> = (1..=size).collect();

    let mut state = seed;

    for i in (1..size).rev() {
        state ^= state << 13;
        state ^= state >> 7;
        state ^= state << 17;

        let j = (state as usize) % (i + 1);

        v.swap(i, j);
    }

    v
}

fn micros(d: Duration) -> u128 {
    d.as_micros()
}

fn input_array(size: usize, seed: u64) -> Vec<usize> {
    shuffled(size, seed)
}

/// Peak heap growth during `benchmark_sort`, in KiB (explicit buffers such as swap).
fn run_once(size: usize, seed: usize) -> (u128, usize) {
    let mut array = input_array(size, seed as u64);

    let base_bytes = LIVE_BYTES.load(Ordering::Relaxed);
    PEAK_BYTES.store(base_bytes, Ordering::Relaxed);

    let start = Instant::now();

    benchmark_sort(&mut array);

    let elapsed = start.elapsed();
    let peak_bytes = PEAK_BYTES.load(Ordering::Relaxed);
    let aux_kb = peak_bytes.saturating_sub(base_bytes) / 1024;

    let expected: Vec<usize> = (1..=size).collect();
    if array != expected {
        panic!(
            "sort failed with seed {} for size {}",
            seed,
            size
        );
    }

    (micros(elapsed), aux_kb)
}

fn run_child(args: &[String]) {
    let size = args[2].parse::<usize>().expect("invalid size");
    let seed = args[3].parse::<usize>().expect("invalid seed");
    let (elapsed_us, mem) = run_once(size, seed);
    println!("{} {}", elapsed_us, mem);
}

fn main() {
    let args: Vec<String> = env::args().collect();
    if args.get(1).is_some_and(|arg| arg == "--run-once") {
        run_child(&args);
        return;
    }

    run_correctness_checks();

    println!(
        "| {:>10} | {:>15} | {:>15} | {:>15} | {:>15} |",
        "Size",
        "Average time",
        "Maximum time",
        "Average memory",
        "Maximum memory"
    );

    println!(
        "|{:-<11}:|{:-<16}:|{:-<16}:|{:-<16}:|{:-<16}:|",
        "",
        "",
        "",
        "",
        ""
    );

    for power in MIN_POWER..=MAX_POWER {
        let size = 1usize << power;

        let mut total_time: u128 = 0;
        let mut max_time: u128 = 0;

        let mut total_mem: usize = 0;
        let mut max_mem: usize = 0;

        for seed in 1..=RUNS {
            let output = Command::new(env::current_exe().expect("failed to find current executable"))
                .arg("--run-once")
                .arg(size.to_string())
                .arg(seed.to_string())
                .output()
                .expect("failed to run benchmark child process");

            if !output.status.success() {
                panic!(
                    "benchmark child process failed: {}",
                    String::from_utf8_lossy(&output.stderr)
                );
            }

            let stdout = String::from_utf8(output.stdout)
                .expect("child process returned non-UTF-8 output");
            let mut fields = stdout.split_whitespace();
            let elapsed_us = fields
                .next()
                .expect("missing elapsed time")
                .parse::<u128>()
                .expect("invalid elapsed time");
            let aux_mem = fields
                .next()
                .expect("missing memory usage")
                .parse::<usize>()
                .expect("invalid memory usage");

            total_time += elapsed_us;

            if elapsed_us > max_time {
                max_time = elapsed_us;
            }

            total_mem += aux_mem;

            if aux_mem > max_mem {
                max_mem = aux_mem;
            }
        }

        let avg_time = total_time / RUNS as u128;
        let avg_mem = total_mem / RUNS;

        println!(
            "| {:>10} | {:>15} | {:>15} | {:>15} | {:>15} |",
            size,
            format!("{}.{:06}", avg_time / 1_000_000, avg_time % 1_000_000),
            format!("{}.{:06}", max_time / 1_000_000, max_time % 1_000_000),
            avg_mem,
            max_mem
        );
    }
}
RUST

RUN cargo build --release

CMD ["./target/release/rust-benchmark"]
EOF

docker build -t rust-benchmark "$WORKDIR"
docker run --rm --init rust-benchmark