提交记录 48068


用户 题目 状态 得分 用时 内存 语言 代码长度
iMMIQ 1001. 测测你的排序 Accepted 100 444.792 ms 692368 KB C++17 31.73 KB
提交时间 评测时间
2026-09-15 12:57:27 2026-09-15 12:57:33
#pragma GCC target("avx2,bmi,bmi2,popcnt,lzcnt")
#include <bits/stdc++.h>
#include <immintrin.h>

// V14: keep V13's first two scatters, samplers, and checked retry paths.
// Low-16 kernel: a 12-input byte network in 32 parallel columns;
// SIMD insertion for entries 13..16; exact old fallbacks for larger buckets.
// 39-comparator schedule: Bert Dobbelaere's sorting-network list,
// https://bertdobbelaere.github.io/sorting_networks.html#N12L39D9
// No timers, autotuning, threads, I/O, or distribution correctness assumptions.
namespace duck_sort_v14_detail {
template <int STRIDE>
std::array<int, 256> get_population_upper_bounds(uint8_t* A, int N, int budget, int sample_size) {
    std::array<int, 256> results;
    results.fill(0);

    size_t n = (size_t)sample_size;
    
    // 1. Calculate Cache Line Alignment Info
    uintptr_t start_addr = (uintptr_t)A;
    // Align up to next 64-byte boundary
    uintptr_t aligned_start = (start_addr + 63) & ~63ULL;
    // Align down end address
    uintptr_t end_addr_exclusive = start_addr + (size_t)N * STRIDE;
    uintptr_t aligned_end = end_addr_exclusive & ~63ULL;

    if (aligned_end <= aligned_start) {
        // Not enough data for aligned sampling, fallback to full scan
         for (int i = 0; i < N; ++i) results[A[(size_t)i * STRIDE]]++;
         return results;
    }

    size_t num_lines = (aligned_end - aligned_start) / 64;
    
    // Safety check: if no cache lines available, fallback to full scan
    if (num_lines == 0) {
        for (int i = 0; i < N; ++i) results[A[(size_t)i * STRIDE]]++;
        return results;
    }
    
    // Determine the offset pattern for the FIRST aligned block
    // We need (aligned_start + off) % STRIDE == start_addr % STRIDE
    size_t diff = aligned_start - start_addr;
    int base_offset = (STRIDE - (diff % STRIDE)) % STRIDE;

    // 2. Adjust Sample Size to be in terms of cache lines
    // Average items per line
    int items_per_line_approx = 64 / STRIDE;
    size_t lines_to_sample = (n + items_per_line_approx - 1) / items_per_line_approx;
    
    // Cap at available lines
    if (lines_to_sample > num_lines) lines_to_sample = num_lines;
    
    // Recalculate actual n for statistics
    // (This is an approximation if stride=3 because different lines have different counts, 
    // but for large N it converges)
    // For Stride=4, count is always 16.
    // For Stride=3, count is 21 or 22 (avg 21.33).
    // optimizing: just counting actually sampled items is better, 
    // but user code expects 'n' to be passed to math formulas.
    // We will count exact sampled items in the loop.
    
    // 3. Sparse Sampling of Cache Lines
    std::array<int, 256> sample_counts;
    sample_counts.fill(0);
    size_t actual_sampled_count = 0;
    
    static std::mt19937 gen;
    const uint32_t mod_blocks = (uint32_t)num_lines;
    const uint64_t mu = ((unsigned __int128)1 << 64) / mod_blocks;

    for (size_t i = 0; i < lines_to_sample; ++i) {
        // Random Block Index
        uint32_t x = gen();
        uint64_t q = ((unsigned __int128)x * mu) >> 64;
        uint32_t blk_idx = x - q * mod_blocks;
        if (blk_idx >= mod_blocks) blk_idx -= mod_blocks;

        uint8_t* p_line = (uint8_t*)(aligned_start + (size_t)blk_idx * 64);

        // Calculate offset for this specific block
        // Block addr changes by 64. 64 % 3 = 1. 64 % 4 = 0.
        // offset_new = (offset_old - delta_addr) % STRIDE
        // delta_addr = blk_idx * 64
        int current_offset;
        if constexpr (STRIDE == 4) {
             current_offset = base_offset;
        } else {
             // STRIDE == 3
             // shift = (blk_idx) % 3
             // off = (base - shift) % 3
             int shift = blk_idx % 3;
             current_offset = base_offset - shift;
             if (current_offset < 0) current_offset += 3;
        }

        // Fetch fixed number of items per cache line
        // Safe max index check:
        // Stride 4: offset max 3. count 16. max idx = 3 + 15*4 = 63 < 64.
        // Stride 3: offset max 2. count 21. max idx = 2 + 20*3 = 62 < 64.
        const int ITEMS = 64 / STRIDE;
        
        #pragma GCC unroll 21
        for (int k = 0; k < ITEMS; ++k) {
            sample_counts[p_line[current_offset + k * STRIDE]]++;
        }
        actual_sampled_count += ITEMS;
    }
    
    n = actual_sampled_count;

    // 2. 二分查找最优 Z 值
    // 目标:找到最大的 Z,使得 Sum(UpperBounds(Z)) <= Budget
    double low_z = 0.0;
    double high_z = 10.0;
    double best_z = 0.0;
    
    double n_double = (double)n;
    double N_double = (double)N;
    double fpc = (double)(N - n) / (double)(N - 1);
    if (fpc < 0) fpc = 0; // Safety

    // 预计算 p_hat 以加速循环
    std::array<double, 256> p_hats;
    for(int i=0; i<256; ++i) p_hats[i] = sample_counts[i] / n_double;

    for (int iter = 0; iter < 20; ++iter) {
        double mid_z = (low_z + high_z) * 0.5;
        double z2 = mid_z * mid_z;
        double div_factor = 1.0 / (1.0 + z2 / n_double);
        
        long long current_sum = 0;
        
        for (int i = 0; i < 256; ++i) {
            double p_hat = p_hats[i];
            
            // Wilson Score Interval
            double term1 = p_hat + z2 / (2.0 * n_double);
            double variance_term = (p_hat * (1.0 - p_hat) / n_double) * fpc;
            if (variance_term < 0) variance_term = 0;
            double term2 = mid_z * std::sqrt(variance_term + z2 / (4.0 * n_double * n_double));
            double p_upper = (term1 + term2) * div_factor;
            
            int limit = (int)std::ceil(N_double * p_upper);
            current_sum += limit;
        }

        if (current_sum <= budget) {
            best_z = mid_z;
            low_z = mid_z;
        } else {
            high_z = mid_z;
        }
    }

    // 3. 使用最佳 Z 生成最终结果
    double z = best_z;
    double z2 = z * z;
    double div_factor = 1.0 / (1.0 + z2 / n_double);
    
    for (int i = 0; i < 256; ++i) {
        double p_hat = p_hats[i];
        double term1 = p_hat + z2 / (2.0 * n_double);
        double variance_term = (p_hat * (1.0 - p_hat) / n_double) * fpc;
        if (variance_term < 0) variance_term = 0;
        double term2 = z * std::sqrt(variance_term + z2 / (4.0 * n_double * n_double));
        double p_upper = (term1 + term2) * div_factor;
        
        int limit = (int)std::ceil(N_double * p_upper);
        if (limit > N) limit = N;
        results[i] = limit;
    }

    return results;
}

std::array<int,256> block_bounds(uint8_t* A,int N,int budget,int sample_size){
    std::array<int,256> results{};
    size_t n=0;
    std::array<int,256> sample_counts{};
    const unsigned num_blocks=N/21;
    if(!num_blocks){
        for(int i=0;i<N;++i)++results[A[3*i+2]];
        return results;
    }
    const unsigned lines_to_sample=std::min(unsigned((sample_size+20)/21),num_blocks);
    static std::mt19937 gen;
    const uint64_t mu=((unsigned __int128)1<<64)/num_blocks;
    for(unsigned i=0;i<lines_to_sample;++i){
        uint32_t x=gen();uint64_t q=((unsigned __int128)x*mu)>>64;
        uint32_t idx=x-q*num_blocks;
        if(idx>=num_blocks)idx-=num_blocks;
        const uint8_t* p=A+size_t(idx)*64;
        #pragma GCC unroll 21
        for(int j=0;j<21;++j)++sample_counts[p[3*j+2]];
    }
    const size_t actual_sampled_count=size_t(lines_to_sample)*21;
    n = actual_sampled_count;

    // 2. 二分查找最优 Z 值
    // 目标:找到最大的 Z,使得 Sum(UpperBounds(Z)) <= Budget
    double low_z = 0.0;
    double high_z = 10.0;
    double best_z = 0.0;
    
    double n_double = (double)n;
    double N_double = (double)N;
    double fpc = (double)(N - n) / (double)(N - 1);
    if (fpc < 0) fpc = 0; // Safety

    // 预计算 p_hat 以加速循环
    std::array<double, 256> p_hats;
    for(int i=0; i<256; ++i) p_hats[i] = sample_counts[i] / n_double;

    for (int iter = 0; iter < 20; ++iter) {
        double mid_z = (low_z + high_z) * 0.5;
        double z2 = mid_z * mid_z;
        double div_factor = 1.0 / (1.0 + z2 / n_double);
        
        long long current_sum = 0;
        
        for (int i = 0; i < 256; ++i) {
            double p_hat = p_hats[i];
            
            // Wilson Score Interval
            double term1 = p_hat + z2 / (2.0 * n_double);
            double variance_term = (p_hat * (1.0 - p_hat) / n_double) * fpc;
            if (variance_term < 0) variance_term = 0;
            double term2 = mid_z * std::sqrt(variance_term + z2 / (4.0 * n_double * n_double));
            double p_upper = (term1 + term2) * div_factor;
            
            int limit = (int)std::ceil(N_double * p_upper);
            current_sum += limit;
        }

        if (current_sum <= budget) {
            best_z = mid_z;
            low_z = mid_z;
        } else {
            high_z = mid_z;
        }
    }

    // 3. 使用最佳 Z 生成最终结果
    double z = best_z;
    double z2 = z * z;
    double div_factor = 1.0 / (1.0 + z2 / n_double);
    
    for (int i = 0; i < 256; ++i) {
        double p_hat = p_hats[i];
        double term1 = p_hat + z2 / (2.0 * n_double);
        double variance_term = (p_hat * (1.0 - p_hat) / n_double) * fpc;
        if (variance_term < 0) variance_term = 0;
        double term2 = z * std::sqrt(variance_term + z2 / (4.0 * n_double * n_double));
        double p_upper = (term1 + term2) * div_factor;
        
        int limit = (int)std::ceil(N_double * p_upper);
        if (limit > N) limit = N;
        results[i] = limit;
    }

    return results;
}


using namespace std;

const int PREFETCH_DIST = 64;

// 辅助函数:向地址 p 写入 3 字节 (利用 uint32 覆盖写,需保证 buffer 有 padding)
// Input val: [B0, B1, B2, X] (Little Endian) -> Writes B0, B1, B2
inline void store3(uint8_t* __restrict__ p, uint32_t val) {
    std::memcpy(p, &val, 4);
}

// 辅助函数:向地址 p 写入 2 字节
inline void store2(uint8_t* __restrict__ p, uint16_t val) {
    std::memcpy(p, &val, 2);
}


inline uint32_t load3(const uint8_t* p) {
    uint32_t v; std::memcpy(&v,p,4); return v;
}
inline uint16_t load2(const uint8_t* p) {
    uint16_t v; std::memcpy(&v,p,2); return v;
}
// Verify that a full tile can be scattered without leaving the allocation.
// Bucket overlap is detected after the pass; the original source stays intact.
inline bool top_tile_fits(uint8_t* const* p, const uint8_t* limit) {
    const __m256i bound=_mm256_set1_epi64x(reinterpret_cast<intptr_t>(limit));
    __m256i bad=_mm256_setzero_si256();
    for (int k=0;k<256;k+=4) {
        const __m256i q=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(p+k));
        bad=_mm256_or_si256(bad,_mm256_cmpgt_epi64(q,bound));
    }
    return _mm256_movemask_epi8(bad)==0;
}


struct Low16Workspace {
    // Only the first 12*256 bytes are initialized on the fast path.
    // column_storage owns the allocation; columns is 64-byte aligned.
    // The 2 MiB allocation also bounds a worst-case scatter: with c<=8192,
    // max written offset is 255+(8192-1)*256 = 2097151.
    uint8_t* columns=nullptr;
    void* column_storage=nullptr;
    unsigned* dense=nullptr;
    ~Low16Workspace(){std::free(dense);std::free(column_storage);}
};
inline void low16_dense(const uint8_t* src,int c,unsigned* dst,
                        unsigned prefix,Low16Workspace& ws){
    if(!ws.dense)ws.dense=static_cast<unsigned*>(std::malloc(65536*sizeof(unsigned)));
    if(!ws.dense){
        for(int i=0;i<c;++i)dst[i]=prefix|load2(src+2*size_t(i));
        std::sort(dst,dst+c);return;
    }
    std::memset(ws.dense,0,65536*sizeof(unsigned));
    for(int i=0;i<c;++i)++ws.dense[load2(src+2*size_t(i))];
    for(unsigned v=0;v<65536;++v){
        unsigned count=ws.dense[v];
        for(unsigned j=0;j<count;++j)*dst++=prefix|v;
    }
}
__attribute__((always_inline)) inline void compare_byte_vectors(__m256i& a,__m256i& b){
    __m256i t=_mm256_min_epu8(a,b);
    b=_mm256_max_epu8(a,b);a=t;
}
__attribute__((always_inline)) inline void column_batch(const uint8_t* src,uint8_t* dst0,uint8_t* dst1){
    __m256i r0=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+0));
    __m256i r1=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+256));
    __m256i r2=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+512));
    __m256i r3=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+768));
    __m256i r4=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+1024));
    __m256i r5=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+1280));
    __m256i r6=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+1536));
    __m256i r7=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+1792));
    __m256i r8=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+2048));
    __m256i r9=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+2304));
    __m256i r10=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+2560));
    __m256i r11=_mm256_loadu_si256(reinterpret_cast<const __m256i*>(src+2816));
    compare_byte_vectors(r0,r8);
    compare_byte_vectors(r1,r7);
    compare_byte_vectors(r2,r6);
    compare_byte_vectors(r3,r11);
    compare_byte_vectors(r4,r10);
    compare_byte_vectors(r5,r9);
    compare_byte_vectors(r0,r1);
    compare_byte_vectors(r2,r5);
    compare_byte_vectors(r3,r4);
    compare_byte_vectors(r6,r9);
    compare_byte_vectors(r7,r8);
    compare_byte_vectors(r10,r11);
    compare_byte_vectors(r0,r2);
    compare_byte_vectors(r1,r6);
    compare_byte_vectors(r5,r10);
    compare_byte_vectors(r9,r11);
    compare_byte_vectors(r0,r3);
    compare_byte_vectors(r1,r2);
    compare_byte_vectors(r4,r6);
    compare_byte_vectors(r5,r7);
    compare_byte_vectors(r8,r11);
    compare_byte_vectors(r9,r10);
    compare_byte_vectors(r1,r4);
    compare_byte_vectors(r3,r5);
    compare_byte_vectors(r6,r8);
    compare_byte_vectors(r7,r10);
    compare_byte_vectors(r1,r3);
    compare_byte_vectors(r2,r5);
    compare_byte_vectors(r6,r9);
    compare_byte_vectors(r8,r10);
    compare_byte_vectors(r2,r3);
    compare_byte_vectors(r4,r5);
    compare_byte_vectors(r6,r7);
    compare_byte_vectors(r8,r9);
    compare_byte_vectors(r4,r6);
    compare_byte_vectors(r5,r7);
    compare_byte_vectors(r3,r4);
    compare_byte_vectors(r5,r6);
    compare_byte_vectors(r7,r8);
    {uint8_t* dst=dst0;
    __m256i t0=_mm256_unpacklo_epi8(r0,r1);
    __m256i t1=_mm256_unpackhi_epi8(r0,r1);
    __m256i t2=_mm256_unpacklo_epi8(r2,r3);
    __m256i t3=_mm256_unpackhi_epi8(r2,r3);
    __m256i t4=_mm256_unpacklo_epi8(r4,r5);
    __m256i t5=_mm256_unpackhi_epi8(r4,r5);
    __m256i t6=_mm256_unpacklo_epi8(r6,r7);
    __m256i t7=_mm256_unpackhi_epi8(r6,r7);
    __m256i u0=_mm256_unpacklo_epi16(t0,t2);
    __m256i u1=_mm256_unpackhi_epi16(t0,t2);
    __m256i u2=_mm256_unpacklo_epi16(t1,t3);
    __m256i u3=_mm256_unpackhi_epi16(t1,t3);
    __m256i u4=_mm256_unpacklo_epi16(t4,t6);
    __m256i u5=_mm256_unpackhi_epi16(t4,t6);
    __m256i u6=_mm256_unpacklo_epi16(t5,t7);
    __m256i u7=_mm256_unpackhi_epi16(t5,t7);
    __m256i v0=_mm256_unpacklo_epi32(u0,u4);
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+0),_mm256_castsi256_si128(v0));
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+128),_mm256_extracti128_si256(v0,1));
    __m256i v1=_mm256_unpackhi_epi32(u0,u4);
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+16),_mm256_castsi256_si128(v1));
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+144),_mm256_extracti128_si256(v1,1));
    __m256i v2=_mm256_unpacklo_epi32(u1,u5);
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+32),_mm256_castsi256_si128(v2));
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+160),_mm256_extracti128_si256(v2,1));
    __m256i v3=_mm256_unpackhi_epi32(u1,u5);
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+48),_mm256_castsi256_si128(v3));
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+176),_mm256_extracti128_si256(v3,1));
    __m256i v4=_mm256_unpacklo_epi32(u2,u6);
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+64),_mm256_castsi256_si128(v4));
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+192),_mm256_extracti128_si256(v4,1));
    __m256i v5=_mm256_unpackhi_epi32(u2,u6);
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+80),_mm256_castsi256_si128(v5));
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+208),_mm256_extracti128_si256(v5,1));
    __m256i v6=_mm256_unpacklo_epi32(u3,u7);
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+96),_mm256_castsi256_si128(v6));
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+224),_mm256_extracti128_si256(v6,1));
    __m256i v7=_mm256_unpackhi_epi32(u3,u7);
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+112),_mm256_castsi256_si128(v7));
    _mm_store_si128(reinterpret_cast<__m128i*>(dst+240),_mm256_extracti128_si256(v7,1));
    }
    {
        __m256i t0=_mm256_unpacklo_epi8(r8,r9);
        __m256i t1=_mm256_unpackhi_epi8(r8,r9);
        __m256i t2=_mm256_unpacklo_epi8(r10,r11);
        __m256i t3=_mm256_unpackhi_epi8(r10,r11);
        __m256i v0=_mm256_unpacklo_epi16(t0,t2);
        __m256i v1=_mm256_unpackhi_epi16(t0,t2);
        __m256i v2=_mm256_unpacklo_epi16(t1,t3);
        __m256i v3=_mm256_unpackhi_epi16(t1,t3);
        _mm_store_si128(reinterpret_cast<__m128i*>(dst1+0),_mm256_castsi256_si128(v0));
        _mm_store_si128(reinterpret_cast<__m128i*>(dst1+64),_mm256_extracti128_si256(v0,1));
        _mm_store_si128(reinterpret_cast<__m128i*>(dst1+16),_mm256_castsi256_si128(v1));
        _mm_store_si128(reinterpret_cast<__m128i*>(dst1+80),_mm256_extracti128_si256(v1,1));
        _mm_store_si128(reinterpret_cast<__m128i*>(dst1+32),_mm256_castsi256_si128(v2));
        _mm_store_si128(reinterpret_cast<__m128i*>(dst1+96),_mm256_extracti128_si256(v2,1));
        _mm_store_si128(reinterpret_cast<__m128i*>(dst1+48),_mm256_castsi256_si128(v3));
        _mm_store_si128(reinterpret_cast<__m128i*>(dst1+112),_mm256_extracti128_si256(v3,1));
    }

}

alignas(16) static const uint8_t insert_masks[16][16]={
 {128,0,1,2,3,4,5,6,7,8,9,10,11,12,13,14},
 {0,128,1,2,3,4,5,6,7,8,9,10,11,12,13,14},
 {0,1,128,2,3,4,5,6,7,8,9,10,11,12,13,14},
 {0,1,2,128,3,4,5,6,7,8,9,10,11,12,13,14},
 {0,1,2,3,128,4,5,6,7,8,9,10,11,12,13,14},
 {0,1,2,3,4,128,5,6,7,8,9,10,11,12,13,14},
 {0,1,2,3,4,5,128,6,7,8,9,10,11,12,13,14},
 {0,1,2,3,4,5,6,128,7,8,9,10,11,12,13,14},
 {0,1,2,3,4,5,6,7,128,8,9,10,11,12,13,14},
 {0,1,2,3,4,5,6,7,8,128,9,10,11,12,13,14},
 {0,1,2,3,4,5,6,7,8,9,128,10,11,12,13,14},
 {0,1,2,3,4,5,6,7,8,9,10,128,11,12,13,14},
 {0,1,2,3,4,5,6,7,8,9,10,11,128,12,13,14},
 {0,1,2,3,4,5,6,7,8,9,10,11,12,128,13,14},
 {0,1,2,3,4,5,6,7,8,9,10,11,12,13,128,14},
 {0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,128},
};
__attribute__((always_inline)) inline __m128i insert_byte(__m128i sorted,unsigned x){
    __m128i v=_mm_set1_epi8(char(x));
    unsigned mask=unsigned(_mm_movemask_epi8(_mm_cmpeq_epi8(sorted,_mm_max_epu8(sorted,v))));
    // There is at least one 255 sentinel until the last insertion, so mask
    // cannot be zero (even when x==255). Equal values keep exact multiplicity.
    unsigned idx=__builtin_ctz(mask);
    __m128i order=_mm_load_si128(reinterpret_cast<const __m128i*>(insert_masks[idx]));
    return _mm_blendv_epi8(_mm_shuffle_epi8(sorted,order),v,order);
}
__attribute__((noinline))
void low16_network(const uint8_t* __restrict__ src,int c,unsigned* __restrict__ dst,
                  unsigned prefix,Low16Workspace& ws){
    if(c<64){
        for(int i=0;i<c;++i)dst[i]=prefix|load2(src+2*size_t(i));
        std::sort(dst,dst+c);return;
    }
    if(c>8192){low16_dense(src,c,dst,prefix,ws);return;}
    if(!ws.columns){
        ws.column_storage=std::malloc(256*8192+63);
        if(ws.column_storage)ws.columns=reinterpret_cast<uint8_t*>((reinterpret_cast<uintptr_t>(ws.column_storage)+63)&~uintptr_t(63));
    }
    if(!ws.columns){low16_dense(src,c,dst,prefix,ws);return;}
    uint8_t* const col=ws.columns;
    // Unsigned maximum is a valid padding value even when real data equals
    // 255: each column emits exactly its recovered element count.
    std::memset(col,255,256*12);
    alignas(64) unsigned pos[256];
    __m256i starts=_mm256_setr_epi32(0,1,2,3,4,5,6,7);
    const __m256i step=_mm256_set1_epi32(8);
    for(unsigned k=0;k<256;k+=8){
        _mm256_store_si256(reinterpret_cast<__m256i*>(pos+k),starts);
        starts=_mm256_add_epi32(starts,step);
    }
    unsigned i=0;
    for(;i+16<=unsigned(c);i+=16){
        #pragma GCC unroll 16
        for(unsigned j=0;j<16;++j){
            unsigned v=load2(src+size_t(i+j)*2),key=v>>8;
            col[pos[key]]=uint8_t(v);pos[key]+=256;
        }
    }
    for(;i<unsigned(c);++i){unsigned v=load2(src+size_t(i)*2),key=v>>8;col[pos[key]]=uint8_t(v);pos[key]+=256;}
    starts=_mm256_setr_epi32(0,1,2,3,4,5,6,7);
    __m256i overflow=_mm256_setzero_si256();
    for(unsigned k=0;k<256;k+=8){
        __m256i used=_mm256_srli_epi32(_mm256_sub_epi32(
            _mm256_load_si256(reinterpret_cast<const __m256i*>(pos+k)),starts),8);
        _mm256_store_si256(reinterpret_cast<__m256i*>(pos+k),used);
        overflow=_mm256_or_si256(overflow,
            _mm256_cmpgt_epi32(used,_mm256_set1_epi32(32)));
        starts=_mm256_add_epi32(starts,step);
    }
    if(__builtin_expect(_mm256_movemask_epi8(overflow)!=0,0)){low16_dense(src,c,dst,prefix,ws);return;}
    unsigned* const output_end=dst+c;
    alignas(64) uint8_t first[256],second[128];
    for(unsigned block=0;block<256;block+=32){
        // Sort 12 ranks in each of 32 independent byte lanes.
        // Transpose the first 8 and last 4 ranks into contiguous per-bucket rows.
        column_batch(col+block,first,second);
        for(unsigned t=0;t<32;++t){
            const unsigned hi=block+t,size=pos[hi];
            const unsigned base=prefix|(hi<<8);
            if(__builtin_expect(size<=12,1)){
                const __m256i bases=_mm256_set1_epi32(base);
                __m256i v0=_mm256_or_si256(_mm256_cvtepu8_epi32(_mm_loadl_epi64(reinterpret_cast<const __m128i*>(first+8*t))),bases);
                uint32_t four;std::memcpy(&four,second+4*t,4);
                __m128i v1=_mm_or_si128(_mm_cvtepu8_epi32(_mm_cvtsi32_si128(int(four))),_mm256_castsi256_si128(bases));
                // Extra lanes only touch this function's unfinalized suffix.
                // Later columns overwrite them. No store crosses output_end.
                if(__builtin_expect(output_end-dst>=12,1)){
                    _mm256_storeu_si256(reinterpret_cast<__m256i*>(dst),v0);
                    _mm_storeu_si128(reinterpret_cast<__m128i*>(dst+8),v1);
                }else {
                    alignas(32) unsigned values[12];
                    _mm256_store_si256(reinterpret_cast<__m256i*>(values),v0);
                    _mm_store_si128(reinterpret_cast<__m128i*>(values+8),v1);
                    std::memcpy(dst,values,size*4);
                }
            }else if(size<=16){
                uint32_t four;std::memcpy(&four,second+4*t,4);
                __m128i v=_mm_unpacklo_epi64(_mm_loadl_epi64(reinterpret_cast<const __m128i*>(first+8*t)),_mm_cvtsi32_si128(int(four)));
                v=_mm_or_si128(v,_mm_setr_epi32(0,0,0,-1));
                for(unsigned j=12;j<size;++j)v=insert_byte(v,col[hi+256*j]);
                __m256i bases=_mm256_set1_epi32(base);
                __m256i v0=_mm256_or_si256(_mm256_cvtepu8_epi32(v),bases);
                __m256i v1=_mm256_or_si256(_mm256_cvtepu8_epi32(_mm_srli_si128(v,8)),bases);
                _mm256_storeu_si256(reinterpret_cast<__m256i*>(dst),v0);
                if(output_end-dst>=16)_mm256_storeu_si256(reinterpret_cast<__m256i*>(dst+8),v1);
                else {
                    __m256i indices=_mm256_setr_epi32(0,1,2,3,4,5,6,7);
                    __m256i mask=_mm256_cmpgt_epi32(_mm256_set1_epi32(size-8),indices);
                    _mm256_maskstore_epi32(reinterpret_cast<int*>(dst+8),mask,v1);
                }
            }else{
                uint8_t values[32];
                for(unsigned j=0;j<size;++j)values[j]=col[hi+256*j];
                std::sort(values,values+size);
                for(unsigned j=0;j<size;++j)dst[j]=base|values[j];
            }
            dst+=size;
        }
    }
}

template<int FixedN>
void sort_impl(uint* a, int __n) {
    const int n=FixedN ? FixedN : __n;
    if(n<=1)return;
    if(n<4096){std::sort(a,a+n);return;}


    // ---------------------------------------------------------
    // Pass 1: Global MSD (Partition by B3)
    // Read a; write blocks of 21 x [B0,B1,B2] + one padding byte.
    // ---------------------------------------------------------
    
    uint cnt_global[256];
    
    // 1.1 统计 B3 (Sampling & Upper Bounds)
    // Budget set to n * 1.47 (47% over-provisioning)
    int budget = (int)(n * 1.47);
    int sample_size = 20000;
    // A 是 (uint8_t*)a + 3 (B3 byte), stride = 4
    auto bounds = get_population_upper_bounds<4>((uint8_t*)a + 3, n, budget, sample_size);
    
    // 1.2 计算 B3 Offset (Bytes in b)
    // Each high-byte bucket starts on a complete 64-byte block.
    uint ptr_global[256];
    uint32_t offset_b3 = 0;
    for (int i = 0; i < 256; i++) {
        ptr_global[i] = offset_b3;
        offset_b3 += ((unsigned(bounds[i])+20)/21)*64; // Rounded whole-block upper bound.
    }
    
    // 申请 b 数组 
    constexpr int TILE=16384;
    const size_t main_bytes=((size_t(budget)+20)/21)*64+256*64;
    uint8_t* b_alloc=static_cast<uint8_t*>(std::malloc(main_bytes+4*TILE+4096));
    if(!b_alloc){std::sort(a,a+n);return;}
    uint8_t* b=reinterpret_cast<uint8_t*>((reinterpret_cast<uintptr_t>(b_alloc)+63)&~uintptr_t(63));
    bool incomplete=false;
    // 1.3 执行 Pass 1 分发
    {
        uint* __restrict__ src = a;
        uint8_t* pp[256];
        for(int z=0;z<256;++z)pp[z]=b+ptr_global[z];

        // Initializing all bytes also makes full-vector tail writes defined.
        alignas(64) uint8_t buffers[256][64] = {};
        uint8_t* write[256];
        for(unsigned k=0;k<256;++k)write[k]=buffers[k];
        int i=0;
        for(;i<n;) {
            if (__builtin_expect(!top_tile_fits(pp,b+main_bytes),0)) {
                incomplete=true;break;
            }
            const int end=std::min(n,i+TILE);
            for(;i+8<=end;i+=8) {
                _mm_prefetch(reinterpret_cast<const char*>(
                    reinterpret_cast<uintptr_t>(src)+size_t(i+PREFETCH_DIST)*4),_MM_HINT_NTA);
                #pragma GCC unroll 8
                for(int j=0;j<8;++j) {
                    const unsigned v=src[i+j], k=v>>24;
                    uint8_t* q=write[k];
                    store3(q,v);q+=3;
                    if(__builtin_expect((reinterpret_cast<uintptr_t>(q)&63)==63,0)){
                        q-=63;
                        _mm256_stream_si256(reinterpret_cast<__m256i*>(pp[k]),_mm256_load_si256(reinterpret_cast<const __m256i*>(q)));
                        _mm256_stream_si256(reinterpret_cast<__m256i*>(pp[k]+32),_mm256_load_si256(reinterpret_cast<const __m256i*>(q+32)));
                        pp[k]+=64;
                    }
                    write[k]=q;
                }
            }
            for(;i<end;++i){
                const unsigned v=src[i],k=v>>24;
                uint8_t* q=write[k];store3(q,v);q+=3;
                if((reinterpret_cast<uintptr_t>(q)&63)==63){
                    q-=63;
                    _mm256_stream_si256(reinterpret_cast<__m256i*>(pp[k]),_mm256_load_si256(reinterpret_cast<const __m256i*>(q)));
                    _mm256_stream_si256(reinterpret_cast<__m256i*>(pp[k]+32),_mm256_load_si256(reinterpret_cast<const __m256i*>(q+32)));
                    pp[k]+=64;
                }
                write[k]=q;
            }
        }
        for(unsigned k=0;k<256;++k){
            cnt_global[k]=unsigned((pp[k]-(b+ptr_global[k]))/64)*21+unsigned((write[k]-buffers[k])/3);
            if(write[k]!=buffers[k] && !incomplete){
                _mm256_stream_si256(reinterpret_cast<__m256i*>(pp[k]),_mm256_load_si256(reinterpret_cast<const __m256i*>(buffers[k])));
                _mm256_stream_si256(reinterpret_cast<__m256i*>(pp[k]+32),_mm256_load_si256(reinterpret_cast<const __m256i*>(buffers[k]+32)));
            }
        }
        // Order all streaming writes before validation, retry, or reading b.
        _mm_sfence();
    }
    bool retry_global=incomplete;
    if(incomplete) {
        std::memset(cnt_global,0,sizeof(cnt_global));
        for(int i=0;i<n;++i)++cnt_global[a[i]>>24];
    }
    for(int k=0;k<256;++k)
        retry_global |= cnt_global[k] && cnt_global[k]>=static_cast<unsigned>(bounds[k]);
    if(retry_global) {
        unsigned off=0;
        uint8_t* pp[256];
        for(int k=0;k<256;++k) {
            ptr_global[k]=off;pp[k]=b+off;
            off+=((cnt_global[k]+20)/21)*64; // Exact, whole-block capacity.
        }
        int i=0;
        for(;i+16<=n;i+=16) {
            #pragma GCC unroll 16
            for(int j=0;j<16;++j) {
                unsigned v=a[i+j],k=v>>24;store3(pp[k],v);pp[k]+=3;if((reinterpret_cast<uintptr_t>(pp[k])&63)==63)++pp[k];
            }
        }
        for(;i<n;++i){unsigned v=a[i],k=v>>24;store3(pp[k],v);pp[k]+=3;if((reinterpret_cast<uintptr_t>(pp[k])&63)==63)++pp[k];}
    }

    uint8_t* scratch=nullptr;
    size_t scratch_capacity=0;
    Low16Workspace ws;
    unsigned emitted=0;
    for(unsigned h=0;h<256;++h) {
        unsigned count=cnt_global[h];
        if(!count)continue;
        uint8_t* src=b+ptr_global[h];
        unsigned* dst=a+emitted;
        if(count<1024){
            for(unsigned j=0;j<count;++j)dst[j]=(h<<24)|(load3(src+size_t(j/21)*64+size_t(j%21)*3)&0xffffffu);
            std::sort(dst,dst+count);emitted+=count;continue;
        }
        auto cap=block_bounds(src,int(count),int(count*2),5000);
        unsigned start[256],cnt[256];
        size_t bytes=0;
        for(unsigned k=0;k<256;++k){start[k]=unsigned(bytes);bytes+=size_t(cap[k])*2;}
        size_t need=std::max(bytes+size_t(count)*2+8,size_t(count)*2+1032);
        uint8_t* temp;
        // The whole current output range and temp are disjoint. Unlike the
        // old four-pass variant, output is now final directly after Pass 2.
        size_t free_bytes=size_t(n-emitted-count)*4;
        if(need+64<=free_bytes){
            uintptr_t end_addr=reinterpret_cast<uintptr_t>(a+n);
            temp=reinterpret_cast<uint8_t*>((end_addr-need)&~uintptr_t(63));
        }else{
            if(need>scratch_capacity){
                uint8_t* q=static_cast<uint8_t*>(std::malloc(need));
                if(!q){
                    unsigned t=emitted;
                    for(unsigned k=h;k<256;++k)
                        for(unsigned j=0;j<cnt_global[k];++j)
                            a[t++]=(k<<24)|(load3(b+ptr_global[k]+size_t(j/21)*64+size_t(j%21)*3)&0xffffffu);
                    std::sort(a+emitted,a+n);
                    std::free(scratch);std::free(b_alloc);return;
                }
                std::free(scratch);scratch=q;scratch_capacity=need;
            }
            temp=scratch;
        }
        uint8_t* p[256];
        for(unsigned k=0;k<256;++k)p[k]=temp+start[k];
        unsigned i=0;
        const uint8_t* read=src;
        for(;i+21<=count;i+=21,read+=64){
            _mm_prefetch(reinterpret_cast<const char*>(reinterpret_cast<uintptr_t>(read)+256),_MM_HINT_T0);
            #pragma GCC unroll 21
            for(unsigned j=0;j<21;++j){
                unsigned v=load3(read+size_t(j)*3), key=(v>>16)&255;
                store2(p[key],uint16_t(v));p[key]+=2;
            }
        }
        for(unsigned j=0;i<count;++i,++j){unsigned v=load3(read+size_t(j)*3),key=(v>>16)&255;store2(p[key],uint16_t(v));p[key]+=2;}
        bool retry=false;
        for(unsigned k=0;k<256;++k){cnt[k]=unsigned(p[k]-(temp+start[k]))/2;retry|=cnt[k]>unsigned(cap[k]);}
        if(retry){
            unsigned off=0;
            for(unsigned k=0;k<256;++k){start[k]=off;p[k]=temp+off;off+=cnt[k]*2;}
            for(unsigned i=0;i<count;++i){unsigned v=load3(src+size_t(i/21)*64+size_t(i%21)*3),key=(v>>16)&255;store2(p[key],uint16_t(v));p[key]+=2;}
        }
        for(unsigned k=0;k<256;++k){
            if(cnt[k])low16_network(temp+start[k],int(cnt[k]),dst,(h<<24)|(k<<16),ws);
            dst+=cnt[k];
        }
        emitted+=count;
    }
    std::free(scratch);std::free(b_alloc);
}
} // namespace duck_sort_v14_detail
void sort(unsigned* a,int n) {
    static_assert(sizeof(unsigned)==4,"32-bit unsigned required");
    if(n==100000000)duck_sort_v14_detail::sort_impl<100000000>(a,n);
    else duck_sort_v14_detail::sort_impl<0>(a,n);
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #1444.792 ms676 MB + 144 KBAcceptedScore: 100


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