プログラミング
あなたのコードは速い – ただし運が良ければ
Your code is fast – if you're lucky (tiki.li)
要約
この記事では、現代のコンパイラがループを最適化する際に、特定のプログラミングスタイル(例えばソートネットワーク)を使用すると、分岐のない高速な命令を利用できることを解説しています。これによりコードは高速化されますが、コンパイラの最適化能力に依存するため、必ずしも常に最速になるとは限らないという「運」の要素についても触れています。コード例として、ソートネットワークを用いたQuicksortの実装が示されています。
全文翻訳
あなたのコードは速い – ただし運が良ければ
最近、最適化されたQuicksortの実装に取り組んでいるときに、非常に興味深い癖に出くわしました。現代のコンパイラ(特にClang)は、適切なプログラミングスタイルを使用している場合、高速な分岐なし命令を使用してループを最適化します。
sort.h - ソートネットワークによるクイックソート
// SPDX-License-Identifier: MIT
// sort.h - 分岐なしクイックソート
// (c) christof.kaser@gmail.com
#ifndef SORT_H
#define SORT_H
#ifndef BLQS_CMP
#define BLQS_CMP(a, b) ((a) < (b))
#endif
#include <stddef.h>
#include <string.h>
#define min(a, b) (((a) < (b)) ? (a) : (b))
#define SMALLPART 1024
#define SWSZ 512
#define UNROLL 16
#define sort2(a, b) do { \
unsigned m = BLQS_CMP(a, b); \
BLQS_TYPE x = a; \
a = m ? a : b; \
b = m ? b : x; \
} while(0)
#define sort3(a, b, c) do { \
sort2(a, b); sort2(b, c); sort2(a, b); \
} while(0)
#define sort4(a, b, c, d) do { \
sort2(a, b); sort2(c, d); sort2(a, c); \
sort2(b, d); sort2(b, c); \
} while(0)
#define sort5(a, b, c, d, e) do { \
sort2(b, c); sort2(d, e); sort2(b, d); \
sort2(a, c); sort2(a, d); sort2(c, e); \
sort2(a, b); sort2(c, d); sort2(b, c); \
} while(0)
#define sort6(a, b, c, d, e, f) do { \
sort2(a, b); sort2(c, d); sort2(e, f); \
sort2(a, c); sort2(b, d); sort2(e, f); \
sort2(a, e); sort2(b, f); sort2(c, e); \
sort2(d, f); sort2(b, c); sort2(d, e); \
sort2(c, d); \
} while(0)
#define sort7(a, b, c, d, e, f, g) do { \
sort2(a, b); sort2(c, d); sort2(a, c); \
sort2(b, d); sort2(b, c); sort2(e, f); \
sort2(e, g); sort2(f, g); sort2(a, e); \
sort2(b, f); sort2(c, g); sort2(b, e); \
sort2(d, g); sort2(c, e); sort2(b, c); \
sort2(d, f); sort2(d, e); \
} while(0)
#define sort8(a,b,c,d,e,f,g,h) do { \
sort2(a,b); sort2(c,d); sort2(e,f); sort2(g,h); \
sort2(a,c); sort2(b,d); sort2(e,g); sort2(f,h); \
sort2(b,c); sort2(f,g); \
sort2(a,e); sort2(b,f); sort2(c,g); sort2(d,h); \
sort2(c,e); sort2(d,f); \
sort2(b,c); sort2(d,e); sort2(f,g); \
} while (0)
#define sort9(a,b,c,d,e,f,g,h,i) do { \
sort2(a,d); sort2(b,h); sort2(c,f); sort2(e,i); \
sort2(a,h); sort2(c,e); sort2(d,i); sort2(f,g); \
sort2(a,c); sort2(b,d); sort2(e,f); sort2(h,i); \
sort2(b,e); sort2(d,g); sort2(f,h); \
sort2(a,b); sort2(c,e); sort2(d,f); sort2(g,i); \
sort2(c,d); sort2(e,f); sort2(g,h); \
sort2(b,c); sort2(d,e); sort2(f,g); \
} while (0)
#define sort10(a,b,c,d,e,f,g,h,i,j) do { \
sort2(a,i); sort2(b,j); sort2(c,h); sort2(d,f); sort2(e,g); \
sort2(a,c); sort2(b,e); sort2(f,i); sort2(h,j); \
sort2(a,d); sort2(c,e); sort2(f,h); sort2(g,j); \
sort2(a,b); sort2(d,g); sort2(i,j); \
sort2(b,f); sort2(c,d); sort2(e,i); sort2(g,h); \
sort2(b,c); sort2(d,f); sort2(e,g); sort2(h,i); \
sort2(c,d); sort2(e,f); sort2(g,h); \
sort2(d,e); sort2(f,g); \
} while (0)
#define sort11(a,b,c,d,e,f,g,h,i,j,k) do { \
sort2(a,j); sort2(b,g); sort2(c,e); sort2(d,h); sort2(f,i); \
sort2(a,b); sort2(d,f); sort2(e,k); sort2(g,j); sort2(h,i); \
sort2(b,d); sort2(c,f); sort2(e,h); sort2(i,k); \
sort2(a,e); sort2(b,c); sort2(d,h); sort2(f,j); sort2(g,i); \
sort2(a,b); sort2(c,g); sort2(e,f); sort2(h,i); sort2(j,k); \
sort2(c,e); sort2(d,g); sort2(f,h); sort2(i,j); \
sort2(b,c); sort2(d,e); sort2(f,g); sort2(h,i); \
sort2(c,d); sort2(e,f); sort2(g,h); \
} while (0)
#define sort12(a,b,c,d,e,f,g,h,i,j,k,l) do { \
sort2(a,i); sort2(b,h); sort2(c,g); sort2(d,l); sort2(e,k); sort2(f,j); \
sort2(a,c); sort2(b,e); sort2(d,f); sort2(g,i); sort2(h,k); sort2(j,l); \
sort2(a,b); sort2(c,j); sort2(e,h); sort2(f,g); sort2(k,l); \
sort2(b,d); sort2(c,h); sort2(e,j); sort2(i,k); \
sort2(a,b); sort2(c,d); sort2(e,f); sort2(g,h); sort2(i,j); sort2(k,l); \
sort2(b,c); sort2(d,f); sort2(g,i); sort2(j,k); \
sort2(c,e); sort2(d,g); sort2(f,i); sort2(h,j); \
sort2(b,c); sort2(d,e); sort2(f,g); sort2(h,i); sort2(j,k); \
} while (0)
static void sorting_network(BLQS_TYPE* l, int partszm1_min_1) {
switch (partszm1_min_1) {
case 0: break;
case 1: sort2(l[0],l[1]); break;
case 2: sort3(l[0],l[1],l[2]); break;
case 3: sort4(l[0],l[1],l[2],l[3]); break;
case 4: sort5(l[0],l[1],l[2],l[3],l[4]); break;
case 5: sort6(l[0],l[1],l[2],l[3],l[4],l[5]); break;
case 6: sort7(l[0],l[1],l[2],l[3],l[4],l[5],l[6]); break;
case 7: sort8(l[0],l[1],l[2],l[3],l[4],l[5],l[6],l[7]); break;
case 8: sort9(l[0],l[1],l[2],l[3],l[4],l[5],l[6],l[7],l[8]); break;
case 9: sort10(l[0],l[1],l[2],l[3],l[4],l[5],l[6],l[7],l[8],l[9]); break;
case 10: sort11(l[0],l[1],l[2],l[3],l[4],l[5],l[6],l[7],l[8],l[9],l[10]); break;
case 11: sort12(l[0],l[1],l[2],l[3],l[4],l[5],l[6],l[7],l[8],l[9],l[10],l[11]); break;
}
}
#define med5(a,b,c,d,e) do { \
sort2(a,b); sort2(c,d); sort2(a,c); \
sort2(b,d); sort2(b,c); sort2(c,e); \
sort2(b,c); \
} while(0)
static BLQS_TYPE* partition_small(BLQS_TYPE* left, BLQS_TYPE* right) {
BLQS_TYPE* outerleft = left;
BLQS_TYPE* pivp = left + 6;
BLQS_TYPE piv = *pivp;
BLQS_TYPE l1 = left[1],l2 = left[2];
BLQS_TYPE r1 = right[-1], r0 = *right;
med5(l1, l2, piv, r1, r0);
left[1] = l1;
left[2] = l2;
right[-1] = r1;
*right = r0;
left += 3;
right -= 2;
*pivp = *outerleft;
BLQS_TYPE swbuf[SMALLPART];
BLQS_TYPE* sw = swbuf;
BLQS_TYPE* lwr = left;
while (right - left >= UNROLL)
for (int i = UNROLL; i--;) {
BLQS_TYPE x = *left++;
if (BLQS_CMP(x, piv)) {
*lwr = x;
lwr++;
} else {
*sw = x;
sw++;
}
}
while (left <= right) {
BLQS_TYPE x = *left++;
if (BLQS_CMP(x, piv)) {
*lwr = x;
lwr++;
} else {
*sw = x;
sw++;
}
}
memcpy(lwr, swbuf, (sw - swbuf) * sizeof(BLQS_TYPE));
lwr -= 1;
*outerleft = *lwr;
*lwr = piv;
return lwr;
}
static BLQS_TYPE* partition(BLQS_TYPE* left, BLQS_TYPE* right) {
BLQS_TYPE* outerleft = left;
BLQS_TYPE* pivp = left + (right - left) / 2;
BLQS_TYPE piv = *pivp;
med5(left[1],left[2],left[3],left[4],left[5]);
med5(left[11],left[12],left[13],left[14],left[15]);
med5(pivp[-2], pivp[-1], piv, pivp[1], pivp[2]);
med5(right[-14], right[-13], right[-12], right[-11], right[-10]);
med5(right[-4], right[-3], right[-2], right[-1], right[0]);
med5(left[3], left[13], piv, right[-12], right[-2]);
left += 1;
*pivp = *outerleft;
BLQS_TYPE swbuf[SWSZ];
BLQS_TYPE *rwr = right, *sw = swbuf;
BLQS_TYPE *lwr = left;
while (UNROLL < SWSZ - (sw - swbuf) && left < right - UNROLL) {
ptrdiff_t avail = min(right - left, SWSZ - (sw - swbuf));
BLQS_TYPE* endp = right - avail;
while (right > endp + UNROLL) {
for (int i = UNROLL; i--;) {
BLQS_TYPE x = *right--;
if (BLQS_CMP(x, piv)) {
*sw = x;
sw++;
} else {
*rwr = x;
rwr--;
}
}
}
}
while (right - left >= UNROLL && (rwr - right > UNROLL || left - lwr > UNROLL)) {
while (rwr - right > UNROLL && right - left >= UNROLL) {
for (int i = UNROLL; i--;) {
BLQS_TYPE x = *left++;
if (BLQS_CMP(x, piv)) {
*lwr = x;
lwr++;
} else {
*rwr = x;
rwr--;
}
}
}
while (left - lwr > UNROLL && right - left >= UNROLL) {
for (int i = UNROLL; i--;) {
BLQS_TYPE x = *right--;
if (BLQS_CMP(x, piv)) {
*lwr = x;
lwr++;
} else {
*rwr = x;
rwr--;
}
}
}
}
do {
while (rwr > right && left <= right) {
BLQS_TYPE x = *left++;
if (BLQS_CMP(x, piv)) {
*lwr = x;
lwr++;
} else {
*rwr = x;
rwr--;
}
}
while (lwr < left && left <= right) {
BLQS_TYPE x = *right--;
if (BLQS_CMP(x, piv)) {
*lwr = x;
lwr++;
} else {
*rwr = x;
rwr--;
}
}
} while ((lwr < left||rwr > right) && left <= right);
while (left <= right && !BLQS_CMP(*right, piv)) {
right--;
rwr--;
}
memcpy(lwr, swbuf, (sw - swbuf) * sizeof(BLQS_TYPE));
*outerleft = *rwr;
*rwr = piv;
return rwr;
}
static void smallsort(BLQS_TYPE* left, BLQS_TYPE* right) {
while (right - left > 11) {
BLQS_TYPE* mid = partition_small(left, right);
smallsort(left, mid - 1);
left = mid + 1;
}
sorting_network(left, right - left);
}
static void sortr(BLQS_TYPE* left, BLQS_TYPE* right) {
while (1) {
ptrdiff_t partszm1 = right - left;
if (partszm1 <= SMALLPART) break;
BLQS_TYPE* mid = partition(left, right);
if (mid - left < partszm1 / 16) {
if (mid > left) sortr(left, mid - 1);
BLQS_TYPE piv = *mid;
mid += 1;
// collect duplicates
for (BLQS_TYPE* p = mid; p <= right; p++) {
if (!BLQS_CMP(piv, *p)) {
BLQS_TYPE h = *mid;
*mid = *p;
*p = h;
mid++;
}
}
left = mid;
if (right - left <