#pragma GCC target("avx2,bmi,bmi2,popcnt,lzcnt")
#include <bits/stdc++.h>
#include <immintrin.h>
// V13: retain V12's exact low-16-bit sorting kernels and checked recovery.
// First scatter: 256 cache-line-sized (64-byte) software buffers, total 16 KiB.
// Each full buffer holds 21 three-byte records and one padding byte. Flush it
// with two aligned 256-bit non-temporal stores; SFENCE before reading output.
// Second scatter and its sampler explicitly skip each block's padding byte.
// All guesses are checked. Original input survives until the scatter is valid.
// No timers, threads, I/O, run-time tuning, or distinct/uniform-key assumption.
namespace duck_sort_v13_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 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; 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_v13_detail
void sort(unsigned* a,int n) {
static_assert(sizeof(unsigned)==4,"32-bit unsigned required");
if(n==100000000)duck_sort_v13_detail::sort_impl<100000000>(a,n);
else duck_sort_v13_detail::sort_impl<0>(a,n);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 464.897 ms | 676 MB + 144 KB | Accepted | Score: 100 | 显示更多 |