/* 1001 / 1001c adaptive MSD radix sort, 3-byte packed intermediate.
*
* Level 0: fixed-stride partition of `a` by the top byte into `tmp`. Because
* the bucket index IS the top byte, only the low 24 bits are stored:
* each record occupies 3 bytes and is written with a single unaligned
* 32-bit store (its 4th byte is the next record's first byte, written
* immediately afterwards; the last record's 4th byte lands in the
* bucket's own slack). This cuts the intermediate write+read traffic
* by 25% (1001c: 403MB instead of 537MB).
* Level 1: per bucket, partition by byte2 with a fixed-stride 4-byte layout in
* a small L3-resident scratch (rad16). Sub-buckets that fit in L1
* (~4K elements) are finished directly by the low-16-bit kernel
* (leaf16: count byte0, scatter byte0 fusing the byte1 count, scatter
* byte1 straight into `a` with the top byte OR'd back in). A
* sub-bucket that is still huge means the value range is narrow:
* descend to level 2 instead of degenerating.
* Level 2: exact counting partition by byte1 (rad8) then per run a counting
* sort by byte0 (or insertion sort for tiny runs).
*
* Every fixed stride is padded by +17 words: a stride that is a multiple of a
* page makes all concurrent bucket write streams land in the same L1 sets /
* DRAM banks; with balanced data (1001c) that collapses the scatter ~3.8x. */
#include <stdlib.h>
#include <string.h>
#include <immintrin.h>
#pragma GCC target("avx2")
typedef unsigned u32;
#define L16MAX 4096u /* largest range sorted by the L1 low-16-bit kernel */
#define TINY 24u /* largest range handled by insertion sort */
#define PADW 17u /* stride padding, see note above */
static u32 g_c1[256], g_c2[256], g_fill[256];
static u32 g_out_fill[256], g_out_off[257];
static u32 g_leaf[L16MAX];
static u32 *g_s2; /* fixed-stride byte2 layout for the largest bucket */
static u32 *g_bufA, *g_bufB;
static unsigned g_cap2;
static u32 g_top; /* top byte of the bucket currently being sorted */
static inline u32 ld24(const unsigned char *p) {
u32 v;
memcpy(&v, p, 4);
return v & 0xffffffu;
}
static void tiny_sort3(const unsigned char *s, u32 *d, unsigned m) {
u32 b[TINY];
for (unsigned i = 0; i < m; i++) b[i] = ld24(s + (size_t)i * 3u);
for (unsigned i = 1; i < m; i++) {
u32 v = b[i]; int j = (int)i - 1;
while (j >= 0 && b[j] > v) { b[j + 1] = b[j]; j--; }
b[j + 1] = v;
}
for (unsigned i = 0; i < m; i++) d[i] = b[i] | g_top;
}
static void tiny_sort(const u32 *s, u32 *d, unsigned m) {
u32 b[TINY];
for (unsigned i = 0; i < m; i++) b[i] = s[i];
for (unsigned i = 1; i < m; i++) {
u32 v = b[i]; int j = (int)i - 1;
while (j >= 0 && b[j] > v) { b[j + 1] = b[j]; j--; }
b[j + 1] = v;
}
for (unsigned i = 0; i < m; i++) d[i] = b[i] | g_top;
}
/* counting sort by byte0 */
static void leaf8(const u32 *s, u32 *d, unsigned m) {
memset(g_c1, 0, sizeof(g_c1));
for (unsigned i = 0; i < m; i++) g_c1[s[i] & 255u]++;
u32 t = 0;
for (int q = 0; q < 256; q++) { u32 c = g_c1[q]; g_c1[q] = t; t += c; }
for (unsigned i = 0; i < m; i++) { u32 v = s[i]; d[g_c1[v & 255u]++] = v | g_top; }
}
/* sort by the low 16 bits (range must be L1 resident) */
static u32 g_out[L16MAX];
static void leaf16(const u32 *s, u32 *d, unsigned m) {
memset(g_c1, 0, sizeof(g_c1));
for (unsigned i = 0; i < m; i++) g_c1[s[i] & 255u]++;
u32 t = 0;
for (int q = 0; q < 256; q++) { u32 c = g_c1[q]; g_c1[q] = t; t += c; }
memset(g_c2, 0, sizeof(g_c2));
for (unsigned i = 0; i < m; i++) { u32 v = s[i]; g_leaf[g_c1[v & 255u]++] = v; g_c2[(v >> 8) & 255u]++; }
t = 0;
for (int q = 0; q < 256; q++) { u32 c = g_c2[q]; g_c2[q] = t; t += c; }
u32 top = g_top;
u32 *o = g_out;
for (unsigned i = 0; i < m; i++) { u32 v = g_leaf[i]; o[g_c2[(v >> 8) & 255u]++] = v | top; }
/* sequential (cache-line friendly) copy out */
unsigned i = 0;
while (i < m && ((unsigned long)(d + i) & 31u)) { d[i] = o[i]; i++; }
for (; i + 8 <= m; i += 8) {
__m256i x; memcpy(&x, o + i, 32);
_mm256_stream_si256((__m256i *)(d + i), x);
}
for (; i < m; i++) d[i] = o[i];
}
/* byte1 level: exact counting partition into `out`, then byte0 per run */
static void rad8(const u32 *src, u32 *dst, unsigned m, u32 *out) {
if (m <= TINY) { tiny_sort(src, dst, m); return; }
u32 off[257];
memset(g_c1, 0, sizeof(g_c1));
for (unsigned i = 0; i < m; i++) g_c1[(src[i] >> 8) & 255u]++;
u32 t = 0;
for (int q = 0; q < 256; q++) { u32 c = g_c1[q]; off[q] = t; g_c1[q] = t; t += c; }
off[256] = t;
for (unsigned i = 0; i < m; i++) { u32 v = src[i]; out[g_c1[(v >> 8) & 255u]++] = v; }
for (int q = 0; q < 256; q++) {
unsigned c = off[q + 1] - off[q];
if (!c) continue;
if (c <= TINY) tiny_sort(out + off[q], dst + off[q], c);
else leaf8(out + off[q], dst + off[q], c);
}
}
/* byte2 level over 3-byte packed source records; falls back to an exact
* counting partition when a byte2 sub-bucket would overflow its stride */
static void rad16(const unsigned char *src, unsigned m, u32 *dst, u32 *out) {
if (m <= TINY) { tiny_sort3(src, dst, m); return; }
u32 off[257];
unsigned cap = m / 256 + (m >> 9) + 64 + PADW;
int of = 0;
if (cap > g_cap2) { cap = g_cap2; of = 2; }
if (!of) {
memset(g_fill, 0, sizeof(g_fill));
for (unsigned i = 0; i < m; i++) {
u32 v = ld24(src + (size_t)i * 3u);
u32 b = (v >> 16) & 255u; u32 f = g_fill[b];
if (f >= cap) { of = 1; break; }
g_s2[(size_t)b * cap + f] = v; g_fill[b] = f + 1;
}
}
if (!of) {
unsigned base = 0;
for (int b = 0; b < 256; b++) {
unsigned c = g_fill[b];
if (!c) continue;
u32 *sub = g_s2 + (size_t)b * cap;
if (c <= L16MAX) leaf16(sub, dst + base, c);
else rad8(sub, dst + base, c, out);
base += c;
}
return;
}
memset(g_c1, 0, sizeof(g_c1));
for (unsigned i = 0; i < m; i++) g_c1[(ld24(src + (size_t)i * 3u) >> 16) & 255u]++;
u32 t = 0;
for (int q = 0; q < 256; q++) { u32 c = g_c1[q]; off[q] = t; g_c1[q] = t; t += c; }
off[256] = t;
for (unsigned i = 0; i < m; i++) { u32 v = ld24(src + (size_t)i * 3u); out[g_c1[(v >> 16) & 255u]++] = v; }
for (int q = 0; q < 256; q++) {
unsigned c = off[q + 1] - off[q];
if (!c) continue;
if (c <= L16MAX) leaf16(out + off[q], dst + off[q], c);
else rad8(out + off[q], dst + off[q], c, g_bufB);
}
}
/* plain 4-pass LSD; only used if the top-byte distribution is degenerate */
static void lsd_fallback(u32 *a, int n, u32 *tmp) {
static u32 h[4][256];
memset(h, 0, sizeof(h));
for (int i = 0; i < n; i++) {
u32 v = a[i];
h[0][v & 255]++; h[1][(v >> 8) & 255]++; h[2][(v >> 16) & 255]++; h[3][v >> 24]++;
}
u32 *src = a, *dst = tmp;
for (int pass = 0; pass < 4; pass++) {
u32 s = 0;
for (int q = 0; q < 256; q++) { u32 c = h[pass][q]; h[pass][q] = s; s += c; }
for (int i = 0; i < n; i++) { u32 v = src[i]; dst[h[pass][(v >> (pass * 8)) & 255]++] = v; }
u32 *sw = src; src = dst; dst = sw;
}
if (src != a) memcpy(a, src, (size_t)n * 4);
}
void sort(u32 *a, int n) {
if (n <= 1) return;
unsigned stride = (unsigned)((n >> 8) + (n >> 9) + 8192 + PADW);
size_t need3 = (size_t)stride * 768u + 16u; /* 256 buckets * 3 bytes */
size_t need4 = (size_t)n * 4u + 64u; /* lsd fallback scratch */
u32 *tmp = (u32 *)malloc(need3 > need4 ? need3 : need4);
if (!tmp) return;
unsigned char *t3 = (unsigned char *)tmp;
int of = 0;
memset(g_out_fill, 0, sizeof(g_out_fill));
for (int i = 0; i < n; i++) {
u32 v = a[i];
u32 b = v >> 24;
u32 f = g_out_fill[b];
if (f + 1u >= stride) { of = 1; break; }
memcpy(t3 + ((size_t)b * stride + f) * 3u, &v, 4); /* low 24 bits (+1 overrun byte) */
g_out_fill[b] = f + 1;
}
if (of) { lsd_fallback(a, n, tmp); free(tmp); return; }
{
u32 t = 0;
for (int q = 0; q < 256; q++) { g_out_off[q] = t; t += g_out_fill[q]; }
g_out_off[256] = t;
}
unsigned maxb = 0;
for (int q = 0; q < 256; q++) if (g_out_fill[q] > maxb) maxb = g_out_fill[q];
if (!maxb) { free(tmp); return; }
if ((size_t)maxb * 4u > (size_t)n * 4u / 4u + (16u << 20)) { /* too skewed for per-bucket buffers */
lsd_fallback(a, n, tmp);
free(tmp);
return;
}
g_cap2 = maxb / 256 + (maxb >> 9) + 64 + PADW;
g_s2 = (u32 *)malloc((size_t)g_cap2 * 256u * 4u);
g_bufA = (u32 *)malloc((size_t)maxb * 4u);
g_bufB = (u32 *)malloc((size_t)maxb * 4u);
if (!g_s2 || !g_bufA || !g_bufB) { lsd_fallback(a, n, tmp); free(tmp); return; }
for (int b = 0; b < 256; b++) {
unsigned m = g_out_fill[b];
if (!m) continue;
const unsigned char *src = t3 + (size_t)b * stride * 3u;
u32 *dst = a + g_out_off[b];
g_top = (u32)b << 24;
if (m == 1) { dst[0] = ld24(src) | g_top; continue; }
rad16(src, m, dst, g_bufA);
}
free(g_s2); free(g_bufA); free(g_bufB); free(tmp);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 790.809 ms | 904 MB + 52 KB | Accepted | Score: 100 | 显示更多 |