キャッシュ効率型基数ソートを使用する

キャッシュ効率型基数ソート (CRadix sort) は、MSD(最上位桁優先)基数ソートをキャッシュミスが少なくなるよう改めた文字列整列向けのアルゴリズムである。

通常の MSD 基数ソートでは、文字列そのものではなく「各文字列へのポインタ」の配列を並べ替えることが多い。ポインタ配列は連続に読めても、各ポインタの先にある文字列はメモリ上の別々の場所にある。桁を調べるたびにその先へ辿ると、アクセス先が飛び飛びになりキャッシュミスが増えやすい。

キャッシュ効率型基数ソートは、各キーに短いキーバッファを割り当て、先に使う数桁をそこへコピーしてから区分する。並べ替えではポインタ(本記事では整数キーそのもの)と対応するバッファを同じ順番で動かす。次の桁はばらばらの文字列を追い直さず、並んだバッファを順に読めばよいので、キャッシュに載りやすい。

手順は次のとおりである。

  1. キーバッファの確保: キーごとに長さ bs のバッファを用意する。理論上の目安はアルファベットサイズ m・件数 n に対しおよそ log n / log m だが、実装では小さな定数(本記事では bs = 2)で足りることが多い。
  2. バッファへの読込み: まだ見ていない桁のうち先頭 bs 個を各バッファへコピーする。キー本体へのアクセスはこのタイミングに寄せる。
  3. バッファ先頭桁での区分: MSD と同様、バッファの先頭文字(桁)0..r-1 で安定にバケット分けする。キーポインタ(本記事では整数キーそのもの)とバッファブロックを同じ順で入れ替える。
  4. 使用済み桁の廃棄: 調べた桁をバッファ先頭から捨て、残りを前へ詰める。次のパスでも常に先頭だけ見ればよい。
  5. 再充填と再帰: バッファが空になったら次の bs 桁を読み直す。要素が 2 個以上残る各バケットについて、下位桁で 3〜5 を繰り返す。
procedure cradix_sort(A)
  if length(A) = 0 then return
  width = digit_width(maximum(A))
  B[i] = fill_buffer(A[i], start=0, width) for each i
  cradix(A, B, digit_pos=0, width)

procedure cradix(A, B, digit_pos, width)
  if length(A) <= 1 or digit_pos >= width then return
  stable_partition A and B by B[i][0]
  for each non-empty bucket S
    discard_front_digit(B in S)
    next = digit_pos + 1
    if next >= width then continue
    if next mod bs = 0 then
      refill B[i] from A[i] at digit next
    cradix(S, B in S, next, width)

桁数を w、基数を r とすると時間計算量は通常の MSD と同様に O(w · (n + r)) 程度である。追加でキーバッファ O(bs · n) と区分用の作業領域を要する。各パスが安定なら全体も安定になる。

本記事と計測コードでは、サイト共通の整数配列向けに十進桁へ写した簡略版を用いる。文字列ポインタ版と同じく「キーとバッファを同じ順で動かす」点が中心で、素朴な LSD 基数ソートとは設計目標が異なる。

以下のデモでは 3 桁の整数を bs = 2 で扱う。棒の上の括弧内がキーバッファの中身である。

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

基数ソートの記事は LSD(下位桁から)の素朴なカウンティング繰り返しが中心である。

キャッシュ効率型基数ソートは MSD 側に立ち、キーバッファでキャッシュ線上の参照をまとめる点が異なる。

アメリカ国旗ソートも MSD だが、インプレース交換でバケットを作ることに主眼があり、キーバッファは用いない。

バーストソートはキャッシュ効率をトライの遅延展開で稼ぐ別系統である。

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

Size Average time Maximum time Average memory Maximum memory
256 0.000003 0.000062 66 72
512 0.000006 0.000059 86 92
1024 0.000018 0.000052 106 112
2048 0.000034 0.000144 122 128
4096 0.000063 0.000148 122 128
8192 0.000124 0.000256 181 188
16384 0.000354 0.000685 382 388
32768 0.000667 0.001810 372 412
65536 0.001359 0.006057 744 784
131072 0.003513 0.010571 2556 2596
262144 0.006693 0.014070 4076 4144
計測に使用したコードを表示する

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::{
    env,
    process::Command,
    time::{Duration, Instant},
};
const MIN_POWER: u32 = 8;
const MAX_POWER: u32 = 18;
const RUNS: usize = 8192;


const RADIX: usize = 10;
const BS: usize = 2;

fn digit_width(max: usize) -> usize {
    if max == 0 {
        return 1;
    }
    let mut w = 0usize;
    let mut v = max;
    while v > 0 {
        w += 1;
        v /= RADIX;
    }
    w
}

fn digit_at(value: usize, pos: usize, width: usize) -> u8 {
    let power = width - 1 - pos;
    let mut div = 1usize;
    for _ in 0..power {
        div = div.saturating_mul(RADIX);
    }
    ((value / div) % RADIX) as u8
}

fn fill_buffer(value: usize, start: usize, width: usize) -> [u8; BS] {
    let mut buf = [0u8; BS];
    for i in 0..BS {
        if start + i < width {
            buf[i] = digit_at(value, start + i, width);
        }
    }
    buf
}

fn cradix_rec(a: &mut [usize], buffers: &mut [[u8; BS]], digit_pos: usize, width: usize) {
    let n = a.len();
    if n <= 1 || digit_pos >= width {
        return;
    }

    let mut count = [0usize; RADIX];
    for b in buffers.iter() {
        count[b[0] as usize] += 1;
    }

    let mut offset = [0usize; RADIX];
    for i in 1..RADIX {
        offset[i] = offset[i - 1] + count[i - 1];
    }

    let mut out_a = vec![0usize; n];
    let mut out_b = vec![[0u8; BS]; n];
    let mut cursor = offset;
    for i in 0..n {
        let d = buffers[i][0] as usize;
        out_a[cursor[d]] = a[i];
        out_b[cursor[d]] = buffers[i];
        cursor[d] += 1;
    }
    a.copy_from_slice(&out_a);
    buffers.copy_from_slice(&out_b);

    for r in 0..RADIX {
        let start = offset[r];
        let len = count[r];
        if len <= 1 {
            continue;
        }
        let next_pos = digit_pos + 1;
        if next_pos >= width {
            continue;
        }

        let end = start + len;
        for i in start..end {
            for j in 0..BS - 1 {
                buffers[i][j] = buffers[i][j + 1];
            }
            buffers[i][BS - 1] = 0;
        }

        if next_pos % BS == 0 {
            for i in start..end {
                buffers[i] = fill_buffer(a[i], next_pos, width);
            }
        }

        cradix_rec(
            &mut a[start..end],
            &mut buffers[start..end],
            next_pos,
            width,
        );
    }
}

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

    let max = *a.iter().max().unwrap();
    let width = digit_width(max);
    let n = a.len();
    let mut buffers = vec![[0u8; BS]; n];
    for i in 0..n {
        buffers[i] = fill_buffer(a[i], 0, width);
    }
    cradix_rec(a, &mut buffers, 0, width);
}


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

    cradix_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],
    );
    for seed in 0..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 memory_usage_kb() -> usize {
    // VmHWM (peak RSS, KiB). Reported memory subtracts a per-size baseline that only
    // holds the input array, so the table reflects auxiliary space during sorting.
    let contents = std::fs::read_to_string("/proc/self/status")
        .unwrap_or_default();

    for line in contents.lines() {
        if let Some(rest) = line.strip_prefix("VmHWM:") {
            let kb = rest
                .split_whitespace()
                .next()
                .unwrap_or("0")
                .parse::<usize>()
                .unwrap_or(0);

            return kb;
        }
    }

    0
}

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

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

fn run_baseline(size: usize) -> usize {
    let _hold = input_array(size, 1);
    memory_usage_kb()
}

fn run_once(size: usize, seed: usize) -> (u128, usize) {
    let mut array = input_array(size, seed as u64);

    let start = Instant::now();

    benchmark_sort(&mut array);

    let elapsed = start.elapsed();
    let mem = memory_usage_kb();

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

    (micros(elapsed), mem)
}

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

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 == "--baseline-once") {
        run_baseline_child(&args);
        return;
    }
    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 baseline_output = Command::new(env::current_exe().expect("failed to find current executable"))
            .arg("--baseline-once")
            .arg(size.to_string())
            .output()
            .expect("failed to run benchmark baseline process");

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

        let baseline_stdout = String::from_utf8(baseline_output.stdout)
            .expect("baseline process returned non-UTF-8 output");
        let baseline_mem = baseline_stdout
            .split_whitespace()
            .next()
            .expect("missing baseline memory usage")
            .parse::<usize>()
            .expect("invalid baseline memory usage");

        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 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;
            }

            let aux_mem = mem.saturating_sub(baseline_mem);

            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