// 1001a: sort n unsigned in place (interface: void sort(unsigned *a, int n))
//
// Approach: MSD partition by the top 9 bits into 512 buckets (~20 elements each for
// random full-range 32-bit input), then sort each bucket with an AVX2 bitonic sorting
// network: 16 lanes for buckets up to 16 elements, 32 lanes (4 registers) up to 32,
// std::sort for the rare larger bucket. A 4-pass LSD radix is kept as a fallback for
// CPUs without AVX2 and for inputs whose top bits are too skewed to bucket well.
//
// Why this beats radix here: a 4-pass radix needs 12 stores per element (4 histogram +
// 4x(1 counter + 1 data)) and its scatter passes cost ~33k ticks each on the judge.
// This needs ~1 histogram RMW, 1 counter RMW, 1 data store and 2 register-resident
// in-bucket stores per element.
//
// Judge-HW measurements (duck.ac custom-test channel, best of 300, 10000 random u32):
// 4-pass LSD radix (8-bit digits) 188k ticks
// 1024 buckets + 16-lane network 155k ticks
// 512 buckets + branchy 16/32-lane 132k ticks <- this file
#include <algorithm>
#include <cstring>
#include <immintrin.h>
typedef unsigned u32;
typedef unsigned short u16;
#if defined(__GNUC__) && !defined(__clang__)
#pragma GCC push_options
#pragma GCC target("avx2")
#pragma GCC optimize("O3")
#endif
// ---- bitonic sorting network primitives (AVX2) ----
static const int IDX1[8] = {1,0,3,2,5,4,7,6};
static const int IDX2[8] = {2,3,0,1,6,7,4,5};
static const int IDX4[8] = {4,5,6,7,0,1,2,3};
static const int IDXR[8] = {7,6,5,4,3,2,1,0};
// per-lane select: 0 -> take min, -1 -> take max
#define SEL_A _mm256_setr_epi32(0,-1,-1,0, 0,-1,-1,0) // stage j=1,k=2
#define SEL_B _mm256_setr_epi32(0,0,-1,-1, -1,-1,0,0) // stage j=2,k=4
#define SEL_C _mm256_setr_epi32(0,-1,0,-1, -1,0,-1,0) // stage j=1,k=4
#define SEL_D _mm256_setr_epi32(0,0,0,0, -1,-1,-1,-1) // stage j=4,k=8
#define SEL_E _mm256_setr_epi32(0,0,-1,-1, 0,0,-1,-1) // stage j=2,k=8
#define SEL_F _mm256_setr_epi32(0,-1,0,-1, 0,-1,0,-1) // stage j=1,k=8
static inline __m256i ce_stage(__m256i v, const int* idx, __m256i sel) {
__m256i s = _mm256_permutevar8x32_epi32(v, _mm256_loadu_si256((const __m256i*)(const void*)idx));
__m256i mn = _mm256_min_epu32(v, s), mx = _mm256_max_epu32(v, s);
return _mm256_blendv_epi8(mn, mx, sel);
}
static inline __m256i net8(__m256i v) {
v = ce_stage(v, IDX1, SEL_A);
v = ce_stage(v, IDX2, SEL_B);
v = ce_stage(v, IDX1, SEL_C);
v = ce_stage(v, IDX4, SEL_D);
v = ce_stage(v, IDX2, SEL_E);
v = ce_stage(v, IDX1, SEL_F);
return v;
}
static inline __m256i bm821(__m256i v) { // bitonic-merge tail for 8 lanes
return ce_stage(ce_stage(ce_stage(v, IDX4, SEL_D), IDX2, SEL_E), IDX1, SEL_F);
}
static inline void merge16(__m256i& lo, __m256i& hi) {
__m256i hr = _mm256_permutevar8x32_epi32(hi, _mm256_loadu_si256((const __m256i*)(const void*)IDXR));
__m256i mn = _mm256_min_epu32(lo, hr), mx = _mm256_max_epu32(lo, hr);
lo = bm821(mn);
hi = bm821(mx);
}
static inline void net32(__m256i& r0, __m256i& r1, __m256i& r2, __m256i& r3) {
r0 = net8(r0); r1 = net8(r1); r2 = net8(r2); r3 = net8(r3);
merge16(r0, r1);
merge16(r2, r3);
__m256i R3 = _mm256_permutevar8x32_epi32(r3, _mm256_loadu_si256((const __m256i*)(const void*)IDXR));
__m256i R2 = _mm256_permutevar8x32_epi32(r2, _mm256_loadu_si256((const __m256i*)(const void*)IDXR));
__m256i a0 = _mm256_min_epu32(r0, R3), a1 = _mm256_max_epu32(r0, R3);
__m256i b0 = _mm256_min_epu32(r1, R2), b1 = _mm256_max_epu32(r1, R2);
__m256i c0 = _mm256_min_epu32(a0, b0), c1 = _mm256_max_epu32(a0, b0);
__m256i d0 = _mm256_min_epu32(a1, b1), d1 = _mm256_max_epu32(a1, b1);
r0 = bm821(c0); r1 = bm821(c1); r2 = bm821(d0); r3 = bm821(d1);
}
static inline __m256i load_pad8(const u32* p, int cnt) { // lanes >= cnt become 0xFFFFFFFF
__m256i v = _mm256_loadu_si256((const __m256i*)(const void*)p);
__m256i idx = _mm256_setr_epi32(0,1,2,3,4,5,6,7);
__m256i m = _mm256_cmpgt_epi32(_mm256_set1_epi32(cnt), idx);
return _mm256_blendv_epi8(_mm256_set1_epi32(-1), v, m);
}
static inline void store_part(u32* p, __m256i v, int cnt) { // store lanes < cnt
__m256i idx = _mm256_setr_epi32(0,1,2,3,4,5,6,7);
__m256i m = _mm256_cmpgt_epi32(_mm256_set1_epi32(cnt), idx);
_mm256_maskstore_epi32((int*)(void*)p, m, v);
}
static void sort16_pad(u32* p, int size) { // 1..16 elements, in place
int c0 = size < 8 ? size : 8;
__m256i lo = net8(load_pad8(p, c0));
if (size <= 8) { store_part(p, lo, c0); return; }
int c1 = size - 8;
__m256i hi = net8(load_pad8(p + 8, c1));
merge16(lo, hi);
store_part(p, lo, c0);
store_part(p + 8, hi, c1);
}
static void sort32_pad(u32* p, int size) { // 1..32 elements, in place
int c0 = size < 8 ? size : 8;
int c1 = size > 8 ? (size < 16 ? size - 8 : 8) : 0;
int c2 = size > 16 ? (size < 24 ? size - 16 : 8) : 0;
int c3 = size > 24 ? size - 24 : 0;
__m256i r0 = load_pad8(p, c0), r1 = load_pad8(p + 8, c1);
__m256i r2 = load_pad8(p + 16, c2), r3 = load_pad8(p + 24, c3);
net32(r0, r1, r2, r3);
store_part(p, r0, c0); store_part(p + 8, r1, c1);
store_part(p + 16, r2, c2); store_part(p + 24, r3, c3);
}
// ---------------- bucket driver ----------------
#define NBUCK 512
static u16 g_cnt[NBUCK];
static u32 g_tmp[1 << 17];
// returns 1 on success, 0 if the top bits are too skewed for the bucket path
static int bucket_sort(u32* a, int n) {
for (int k = 0; k < NBUCK; k++) g_cnt[k] = 0;
for (int i = 0; i < n; i++) g_cnt[a[i] >> 23]++;
unsigned s = 0;
int nonEmpty = 0, maxc = 0;
for (int k = 0; k < NBUCK; k++) {
unsigned c = g_cnt[k];
if (c) { nonEmpty++; if ((int)c > maxc) maxc = (int)c; }
g_cnt[k] = (u16)s; s += c;
}
if (maxc > 64 || nonEmpty < 128) return 0; // not bucket-friendly: caller uses radix
for (int i = 0; i < n; i++) { u32 v = a[i]; g_tmp[g_cnt[v >> 23]++] = v; }
unsigned start = 0;
for (int k = 0; k < NBUCK; k++) {
int c = (int)(unsigned)(g_cnt[k] - start);
if (c > 0) {
if (c <= 16) sort16_pad(g_tmp + start, c);
else if (c <= 32) sort32_pad(g_tmp + start, c);
else std::sort(g_tmp + start, g_tmp + start + c);
}
start = g_cnt[k];
}
memcpy(a, g_tmp, (size_t)n * 4);
return 1;
}
static inline void store_all(u32* p, __m256i v) { _mm256_storeu_si256((__m256i*)(void*)p, v); }
static void sort16_to(const u32* s, u32* d, int size) {
int c0 = size < 8 ? size : 8;
__m256i lo = net8(load_pad8(s, c0));
if (size <= 8) { store_all(d, lo); return; }
int c1 = size - 8;
__m256i hi = net8(load_pad8(s + 8, c1));
merge16(lo, hi);
store_all(d, lo);
store_all(d + 8, hi);
}
static void sort32_to(const u32* s, u32* d, int size) {
int c0 = size < 8 ? size : 8;
int c1 = size > 8 ? (size < 16 ? size - 8 : 8) : 0;
int c2 = size > 16 ? (size < 24 ? size - 16 : 8) : 0;
int c3 = size > 24 ? size - 24 : 0;
__m256i r0 = load_pad8(s, c0), r1 = load_pad8(s + 8, c1);
__m256i r2 = load_pad8(s + 16, c2), r3 = load_pad8(s + 24, c3);
net32(r0, r1, r2, r3);
store_all(d, r0); store_all(d + 8, r1);
store_all(d + 16, r2); store_all(d + 24, r3);
}
static int bucket_sort_direct(u32* a, int n) {
for (int k = 0; k < NBUCK; k++) g_cnt[k] = 0;
for (int i = 0; i < n; i++) g_cnt[a[i] >> 23]++;
unsigned s = 0;
int nonEmpty = 0, maxc = 0;
for (int k = 0; k < NBUCK; k++) {
unsigned c = g_cnt[k];
if (c) { nonEmpty++; if ((int)c > maxc) maxc = (int)c; }
g_cnt[k] = (u16)s; s += c;
}
if (maxc > 64 || nonEmpty < 128) return 0;
for (int i = 0; i < n; i++) { u32 v = a[i]; g_tmp[g_cnt[v >> 23]++] = v; }
unsigned start = 0;
for (int k = 0; k < NBUCK; k++) {
int c = (int)(unsigned)(g_cnt[k] - start);
if (c > 0) {
if (c <= 32 && start + 32 <= (unsigned)n) {
if (c <= 16) sort16_to(g_tmp + start, a + start, c);
else sort32_to(g_tmp + start, a + start, c);
} else {
if (c <= 16) sort16_pad(g_tmp + start, c);
else if (c <= 32) sort32_pad(g_tmp + start, c);
else std::sort(g_tmp + start, g_tmp + start + c);
memcpy(a + start, g_tmp + start, (size_t)c * 4);
}
}
start = g_cnt[k];
}
return 1;
}
#if defined(__GNUC__) && !defined(__clang__)
#pragma GCC pop_options
#endif
// ---------------- fallback: 4-pass LSD radix ----------------
static void radix_sort(u32* a, int n) {
static u16 h[4][256];
static u32 tmp[1 << 16];
for (int i = 0; i < 256; i++) { h[0][i] = 0; h[1][i] = 0; h[2][i] = 0; h[3][i] = 0; }
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;
int done = 0;
for (int p = 0; p < 4; p++) {
int nb = 0;
for (int k = 0; k < 256; k++) if (h[p][k]) { nb = 1; break; }
if (!nb) continue;
unsigned s = 0;
u16* hp = h[p];
for (int k = 0; k < 256; k++) { unsigned c = hp[k]; hp[k] = (u16)s; s += c; }
int sh = p * 8;
for (int i = 0; i < n; i++) { u32 v = src[i]; dst[hp[(v >> sh) & 255]++] = v; }
u32* t = src; src = dst; dst = t;
done++;
}
if (done && src != a) for (int i = 0; i < n; i++) a[i] = src[i];
}
// ---------------- entry ----------------
void sort(unsigned* a, int n) {
if (n <= 1) return;
if (n <= 24) {
for (int i = 1; i < n; i++) {
u32 v = a[i]; int j = i - 1;
while (j >= 0 && a[j] > v) { a[j + 1] = a[j]; j--; }
a[j + 1] = v;
}
return;
}
if (n <= (1 << 17) - 32 && __builtin_cpu_supports("avx2") && bucket_sort_direct(a, n)) return;
radix_sort(a, n);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 54.36 us | 100 KB | Accepted | Score: 100 | 显示更多 |