クワッドソートを使用する

クワッドソート (quadsort) は、安定・適応的なボトムアップ型マージソートである。小区間を 4 要素単位で整え、隣接する整列済みブロックを 4 本まとめて併合(ピンポン・クワッドマージ)する点が特徴で、整列済み区間の併合を省略できる。

本記事の実装は説明用に簡略化しており、実際の実装で用いられる無分岐のパリティマージやクロスマージの代わりに、通常の安定二分マージを用いる。ブロック長も 8 ではなく 4 から始める。

  1. 整列・逆整列の早期終了: 全体が昇順なら何もしない。厳密な降順なら反転して終了する(比較・移動とも O(n))。
  2. クワッドスワップ: 4 要素ごとに比較交換ネットワークで整列し、長さ 4 の整列済みブロック列にする。余りは挿入ソートで整える。
  3. ピンポン・クワッドマージ: 隣接 4 ブロックを、まず 2 組ずつ補助配列へマージし、続けて本配列へ戻す 1 回のマージで 4 倍長の整列済み列にする(往復コピーを減らす)。
  4. スキップ: 4 ブロックの境界がすべて昇順なら併合を省略する。これにより完全整列済み入力は追加の O(n log n) マージを避けられる。
  5. 余り: 4 ブロックに満たない残りは、通常の二分ボトムアップマージで吸収する。ブロック長を 4 倍しながら繰り返す。
procedure quad_swap4(A, i0, i1, i2, i3)
  compare-exchange pairs and cross pairs until four keys are sorted

procedure quad_merge_four(A, swap, start, block)
  if A[start+block-1] ≤ A[start+block]
     and A[start+2*block-1] ≤ A[start+2*block]
     and A[start+3*block-1] ≤ A[start+3*block] then
    return
  merge A[start .. start+2*block) into swap via two halves
  merge A[start+2*block .. start+4*block) into swap
  merge the two halves of swap back into A[start .. start+4*block)

procedure quadsort(A)
  if A is sorted then return
  if A is reverse-sorted then reverse(A); return
  for each complete group of 4 elements
    quad_swap4(group)
  insertion_sort(tail shorter than 4)
  block := 4
  while block < length(A)
    for each aligned span of 4*block elements
      quad_merge_four(A, swap, span, block)
    binary bottom-up merge any leftover runs of size block
    block := block * 4

最良は整列済み検出により O(n)、平均・最悪は O(n log n) である。補助配列に最大 O(n) を使う安定ソートである。

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

マージソートは 2 分割再帰が基本である。クワッドソートはボトムアップで 4 ブロック単位のピンポン併合と、境界昇順時のスキップを前面に出す。

ティムソートは自然ランの検出とギャロッピング併合が中心である。クワッドソートは固定長ブロックのクワッドスワップから始め、整列度に応じて併合を省略する。

フラックスソートはクイック型の安定分割を主とし、小区間でクワッドソート系の仕上げを使う。本記事のクワッドソートは分割を行わず、マージ側だけで完結する。

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

Size Average time Maximum time Average memory Maximum memory
256 0.000005 0.000053 78 84
512 0.000013 0.000065 69 76
1024 0.000028 0.000277 66 72
2048 0.000062 0.000251 78 84
4096 0.000134 0.000310 90 96
8192 0.000295 0.002911 122 128
16384 0.000639 0.000919 74 80
32768 0.001367 0.002252 104 144
65536 0.002976 0.004847 360 400
131072 0.006242 0.011857 879 920
262144 0.013037 0.027442 1912 2060
計測に使用したコードを表示する

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;


/// Sort four elements with a small sorting network (equals keep order via `>`).
fn quad_swap4(a: &mut [usize], i0: usize, i1: usize, i2: usize, i3: usize) {
    if a[i0] > a[i1] {
        a.swap(i0, i1);
    }
    if a[i2] > a[i3] {
        a.swap(i2, i3);
    }
    if a[i0] > a[i2] {
        a.swap(i0, i2);
    }
    if a[i1] > a[i3] {
        a.swap(i1, i3);
    }
    if a[i1] > a[i2] {
        a.swap(i1, i2);
    }
}

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

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

fn quad_reverse(a: &mut [usize]) {
    let mut lo = 0;
    let mut hi = a.len();
    while lo + 1 < hi {
        hi -= 1;
        a.swap(lo, hi);
        lo += 1;
    }
}

/// Stable two-way merge from `src[lo..mid)` and `src[mid..hi)` into `dst[lo..hi)`.
fn quad_merge_two(src: &[usize], dst: &mut [usize], lo: usize, mid: usize, hi: usize) {
    let mut i = lo;
    let mut j = mid;
    let mut k = lo;
    while i < mid && j < hi {
        if src[i] <= src[j] {
            dst[k] = src[i];
            i += 1;
        } else {
            dst[k] = src[j];
            j += 1;
        }
        k += 1;
    }
    while i < mid {
        dst[k] = src[i];
        i += 1;
        k += 1;
    }
    while j < hi {
        dst[k] = src[j];
        j += 1;
        k += 1;
    }
}

/// True when four consecutive sorted blocks of length `block` are already ordered
/// across boundaries (skipping the merge is safe).
fn quad_blocks_ordered(a: &[usize], start: usize, block: usize) -> bool {
    a[start + block - 1] <= a[start + block]
        && a[start + block * 2 - 1] <= a[start + block * 2]
        && a[start + block * 3 - 1] <= a[start + block * 3]
}

/// Ping-pong quad merge: two pairwise merges into swap, then one merge back into `a`.
fn quad_merge_four(a: &mut [usize], swap: &mut [usize], start: usize, block: usize) {
    let mid1 = start + block;
    let mid2 = start + block * 2;
    let mid3 = start + block * 3;
    let end = start + block * 4;
    if quad_blocks_ordered(a, start, block) {
        return;
    }
    quad_merge_two(a, swap, start, mid1, mid2);
    quad_merge_two(a, swap, mid2, mid3, end);
    quad_merge_two(swap, a, start, mid2, end);
}

/// Binary bottom-up merge for a partial span that is not a full group of four blocks.
fn quad_merge_remainder(a: &mut [usize], swap: &mut [usize], start: usize, n: usize, block: usize) {
    let mut width = block;
    while start + width < n {
        let mut lo = start;
        while lo + width < n {
            let mid = lo + width;
            let hi = (lo + width * 2).min(n);
            if a[mid - 1] > a[mid] {
                quad_merge_two(a, swap, lo, mid, hi);
                a[lo..hi].copy_from_slice(&swap[lo..hi]);
            }
            lo = hi;
        }
        width *= 2;
    }
}

fn quad_sort(a: &mut [usize]) {
    let n = a.len();
    if n <= 1 {
        return;
    }
    if quad_is_sorted(a) {
        return;
    }
    if quad_is_reverse_sorted(a) {
        quad_reverse(a);
        return;
    }

    // Analyzer / quad-swap: leave sorted blocks of 4 (educational stand-in for 8).
    let mut i = 0;
    while i + 4 <= n {
        quad_swap4(a, i, i + 1, i + 2, i + 3);
        i += 4;
    }
    if i < n {
        for j in (i + 1)..n {
            let key = a[j];
            let mut k = j;
            while k > i && a[k - 1] > key {
                a[k] = a[k - 1];
                k -= 1;
            }
            a[k] = key;
        }
    }

    let mut swap = vec![0usize; n];
    let mut block = 4usize;
    while block < n {
        let stride = block * 4;
        let mut start = 0usize;
        while start < n {
            let rem = n - start;
            if rem <= block {
                break;
            }
            if rem >= stride {
                quad_merge_four(a, &mut swap, start, block);
                start += stride;
            } else {
                quad_merge_remainder(a, &mut swap, start, n, block);
                break;
            }
        }
        block *= 4;
    }
}


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

    quad_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 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