提交记录 48063


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

// V12: keep V11's two checked high-byte partitions. In each low-16-bit
// bucket, scatter the low bytes into columns, and sort 32 independent columns
// at a time with an AVX2 8-input, 19-comparator network. The first and second
// eight values are sorted separately; 9..16-value columns then use a merge.
// Output vector stores may overlap FUTURE output within the SAME low16 bucket,
// but never write beyond its end. Tail stores use exact masks.
// Capacity guesses only affect speed, not correctness. The intact input is
// retained until capacity checks complete. Oversized/skewed buckets use counting.
// No timers, threads, runtime tuning, I/O, or assumptions of distinct keys.
namespace duck_sort_v12_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;
}

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 16*256 bytes are initialized on the fast path.
    // 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;
    unsigned* dense=nullptr;
    ~Low16Workspace(){std::free(dense);std::free(columns);}
};
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;
    }
}
template<int J,int Mask>
__attribute__((always_inline)) inline __m256i step_words(__m256i x){
    const __m256i perm=_mm256_setr_epi8(
        ((0^J)*2),((0^J)*2+1),((1^J)*2),((1^J)*2+1),
        ((2^J)*2),((2^J)*2+1),((3^J)*2),((3^J)*2+1),
        ((4^J)*2),((4^J)*2+1),((5^J)*2),((5^J)*2+1),
        ((6^J)*2),((6^J)*2+1),((7^J)*2),((7^J)*2+1),
        ((0^J)*2),((0^J)*2+1),((1^J)*2),((1^J)*2+1),
        ((2^J)*2),((2^J)*2+1),((3^J)*2),((3^J)*2+1),
        ((4^J)*2),((4^J)*2+1),((5^J)*2),((5^J)*2+1),
        ((6^J)*2),((6^J)*2+1),((7^J)*2),((7^J)*2+1));
    __m256i y=_mm256_shuffle_epi8(x,perm);
    return _mm256_blend_epi16(_mm256_min_epu16(x,y),_mm256_max_epu16(x,y),Mask);
}
__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;
}
// Sort 32 independent 8-byte sequences, then transpose to consecutive rows.
__attribute__((always_inline)) inline void columns8(const uint8_t* src,uint8_t* dst){
    __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));
    compare_byte_vectors(r0,r1);
    compare_byte_vectors(r2,r3);
    compare_byte_vectors(r0,r2);
    compare_byte_vectors(r1,r3);
    compare_byte_vectors(r1,r2);
    compare_byte_vectors(r4,r5);
    compare_byte_vectors(r6,r7);
    compare_byte_vectors(r4,r6);
    compare_byte_vectors(r5,r7);
    compare_byte_vectors(r5,r6);
    compare_byte_vectors(r0,r4);
    compare_byte_vectors(r2,r6);
    compare_byte_vectors(r2,r4);
    compare_byte_vectors(r1,r5);
    compare_byte_vectors(r3,r7);
    compare_byte_vectors(r3,r5);
    compare_byte_vectors(r1,r2);
    compare_byte_vectors(r3,r4);
    compare_byte_vectors(r5,r6);
    __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));
}

__attribute__((always_inline)) inline __m256i merge8_pairs(const uint8_t* a,const uint8_t* b){
    __m128i x=_mm_loadl_epi64(reinterpret_cast<const __m128i*>(a));
    __m128i y=_mm_loadl_epi64(reinterpret_cast<const __m128i*>(b));
    y=_mm_shuffle_epi8(y,_mm_setr_epi8(7,6,5,4,3,2,1,0,-1,-1,-1,-1,-1,-1,-1,-1));
    __m256i v=_mm256_cvtepu8_epi16(_mm_unpacklo_epi64(x,y));
    __m256i w=_mm256_permute2x128_si256(v,v,1);
    v=_mm256_blend_epi32(_mm256_min_epu16(v,w),_mm256_max_epu16(v,w),0xf0);
    v=step_words<4,0xf0>(v);v=step_words<2,0xcc>(v);v=step_words<1,0xaa>(v);
    return v;
}

__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.columns=static_cast<uint8_t*>(std::malloc(256*8192));
    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*16);
    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;
    const __m256i indices=_mm256_setr_epi32(0,1,2,3,4,5,6,7);
    alignas(64) uint8_t first[256],second[256];
    for(unsigned block=0;block<256;block+=32){
        // The network compares corresponding bytes across 32 independent
        // columns. Transpose only after sorting, to make each column contiguous.
        columns8(col+block,first);
        columns8(col+2048+block,second);
        for(unsigned t=0;t<32;++t){
            const unsigned hi=block+t,size=pos[hi];
            if(!size)continue;
            const unsigned base=prefix|(hi<<8);
            const __m256i bases=_mm256_set1_epi32(base);
            if(size<=8){
                __m256i v=_mm256_cvtepu8_epi32(_mm_loadl_epi64(reinterpret_cast<const __m128i*>(first+8*t)));
                v=_mm256_or_si256(v,bases);
                // Any extra lanes are in not-yet-finalized output, and will be
                // overwritten by later columns before this function returns.
                if(output_end-dst>=8)_mm256_storeu_si256(reinterpret_cast<__m256i*>(dst),v);
                else {
                    __m256i mask=_mm256_cmpgt_epi32(_mm256_set1_epi32(size),indices);
                    _mm256_maskstore_epi32(reinterpret_cast<int*>(dst),mask,v);
                }
            }else if(size<=16){
                __m256i v=merge8_pairs(first+8*t,second+8*t);
                __m256i a=_mm256_or_si256(_mm256_cvtepu16_epi32(_mm256_castsi256_si128(v)),bases);
                __m256i b=_mm256_or_si256(_mm256_cvtepu16_epi32(_mm256_extracti128_si256(v,1)),bases);
                _mm256_storeu_si256(reinterpret_cast<__m256i*>(dst),a);
                if(output_end-dst>=16)_mm256_storeu_si256(reinterpret_cast<__m256i*>(dst+8),b);
                else {
                    __m256i mask=_mm256_cmpgt_epi32(_mm256_set1_epi32(size-8),indices);
                    _mm256_maskstore_epi32(reinterpret_cast<int*>(dst+8),mask,b);
                }
            }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 (4 bytes) -> Write: b (3 bytes: [B0, B1, B2])
    // ---------------------------------------------------------
    
    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)
    // 增加 4 字节 Padding 以安全使用 store3
    uint ptr_global[256];
    uint32_t offset_b3 = 0;
    for (int i = 0; i < 256; i++) {
        ptr_global[i] = offset_b3;
        offset_b3 += bounds[i] * 3; // Use Upper Bound
    }
    
    // 申请 b 数组 
    constexpr int TILE=16384;
    const size_t main_bytes=size_t(budget)*3;
    uint8_t* b=static_cast<uint8_t*>(std::malloc(main_bytes+3*TILE+4096));
    if(!b){std::sort(a,a+n);return;}
    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];

        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+16<=end;i+=16) {
                _mm_prefetch(reinterpret_cast<const char*>(
                    reinterpret_cast<uintptr_t>(src)+size_t(i+PREFETCH_DIST)*4),_MM_HINT_NTA);
                #pragma GCC unroll 16
                for(int j=0;j<16;++j) {
                    const unsigned v=src[i+j], k=v>>24;
                    uint8_t* q=pp[k];
                    store3(q,v); pp[k]=q+3;
                }
            }
            for(;i<end;++i){const unsigned v=src[i],k=v>>24;store3(pp[k],v);pp[k]+=3;}
        }
        
        // Reconstruct exact counts from pointer progress
        for(int k=0; k<256; ++k) {
            cnt_global[k] = (pp[k] - (b+ptr_global[k])) / 3;
        }
    }
    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+=3*cnt_global[k]+4; // Allow the fourth byte of store3.
        }
        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;
            }
        }
        for(;i<n;++i){unsigned v=a[i],k=v>>24;store3(pp[k],v);pp[k]+=3;}
    }

    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)*3)&0xffffffu);
            std::sort(dst,dst+count);emitted+=count;continue;
        }
        auto cap=get_population_upper_bounds<3>(src+2,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)*3)&0xffffffu);
                    std::sort(a+emitted,a+n);
                    std::free(scratch);std::free(b);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;
        for(;i+16<=count;i+=16){
            _mm_prefetch(reinterpret_cast<const char*>(reinterpret_cast<uintptr_t>(src)+size_t(i+64)*3),_MM_HINT_T0);
            #pragma GCC unroll 16
            for(unsigned j=0;j<16;++j){
                unsigned v=load3(src+size_t(i+j)*3), key=(v>>16)&255;
                store2(p[key],uint16_t(v));p[key]+=2;
            }
        }
        for(;i<count;++i){unsigned v=load3(src+size_t(i)*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)*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);
}
} // namespace duck_sort_v12_detail
void sort(unsigned* a,int n) {
    static_assert(sizeof(unsigned)==4,"32-bit unsigned required");
    if(n==100000000)duck_sort_v12_detail::sort_impl<100000000>(a,n);
    else duck_sort_v12_detail::sort_impl<0>(a,n);
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #1465.36 ms671 MB + 596 KBAcceptedScore: 100


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