提交记录 61630


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_dsh_v41_0919 1001a. 测测你的排序2 Accepted 100 54.36 us 100 KB C++17 9.83 KB
提交时间 评测时间
2026-09-19 21:18:33 2026-09-19 21:18:35
// 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);
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #154.36 us100 KBAcceptedScore: 100


Judge Duck Online | 评测鸭在线
Server Time: 2026-09-25 05:42:20 | Loaded in 2 ms | Server Status
个人娱乐项目,仅供学习交流使用 | 捐赠