// Duck.ac 1001, C++17 / AVX2.
// v10 candidate, based on supplied v6.
// Original reference: https://duck.ac/submission/48070
// Sorting network: Bert Dobbelaere
// https://bertdobbelaere.github.io/sorting_networks.html#N12L39D9
#include <algorithm>
#include <cstdint>
#include <cstdlib>
#include <cstring>
#include <immintrin.h>
#include <utility>
#pragma GCC optimize("O3,unroll-loops,no-strict-aliasing")
#pragma GCC target("avx2,bmi,bmi2,popcnt,lzcnt")
static_assert(
sizeof(unsigned) == 4 && sizeof(uint16_t) == 2,
"32-bit unsigned required"
);
static_assert(
__BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__,
"little-endian required"
);
namespace fastsort {
template<class F, size_t... I>
[[gnu::always_inline]]
inline void each(F&& f, std::index_sequence<I...>) {
(f(std::integral_constant<size_t, I>{}), ...);
}
template<size_t N, class F>
[[gnu::always_inline]]
inline void repeat(F&& f) {
each(std::forward<F>(f), std::make_index_sequence<N>{});
}
} // namespace fastsort
namespace v2_tail {
using U = unsigned;
using B = uint8_t;
using H = uint16_t;
using V = __m256i;
using W = __m128i;
alignas(64) static B columns[256 * 8192];
inline U load16(const H* p) {
H v;
std::memcpy(&v, p, 2);
return v;
}
[[gnu::noinline]]
inline void dense(const H* src, U n, U* dst, U base) {
alignas(64) U counts[65536] = {};
for (U i = 0; i < n; ++i)
++counts[load16(src + i)];
for (U v = 0; v < 65536; ++v)
for (U c = counts[v]; c; --c)
*dst++ = base | v;
}
template<bool Room>
inline void write8(W x, U* dst, U c, U base, U* end) {
V y = _mm256_or_si256(
_mm256_cvtepu8_epi32(x),
_mm256_set1_epi32(base)
);
if (Room || end - dst >= 8) {
_mm256_storeu_si256((V*)dst, y);
} else {
V mask = _mm256_cmpgt_epi32(
_mm256_set1_epi32(c),
_mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7)
);
_mm256_maskstore_epi32((int*)dst, mask, y);
}
}
[[gnu::always_inline]]
inline void batch_sort(const B* src, B* out, B* tail) {
V x[12], a[8], b[8];
fastsort::repeat<12>(
[&](auto j) __attribute__((always_inline)) {
x[j] = _mm256_load_si256(
(const V*)(src + 256 * j)
);
}
);
#define CMP(A, B) do { \
V lo = _mm256_min_epu8(x[A], x[B]); \
x[B] = _mm256_max_epu8(x[A], x[B]); \
x[A] = lo; \
} while (0)
CMP(0,8); CMP(1,7); CMP(2,6);
CMP(3,11); CMP(4,10); CMP(5,9);
CMP(0,1); CMP(2,5); CMP(3,4);
CMP(6,9); CMP(7,8); CMP(10,11);
CMP(0,2); CMP(1,6); CMP(5,10); CMP(9,11);
CMP(0,3); CMP(1,2); CMP(4,6);
CMP(5,7); CMP(8,11); CMP(9,10);
CMP(1,4); CMP(3,5); CMP(6,8); CMP(7,10);
CMP(1,3); CMP(2,5); CMP(6,9); CMP(8,10);
CMP(2,3); CMP(4,5); CMP(6,7); CMP(8,9);
CMP(4,6); CMP(5,7);
CMP(3,4); CMP(5,6); CMP(7,8);
#undef CMP
#pragma GCC unroll 4
for (U i = 0; i < 4; ++i) {
a[2*i] = _mm256_unpacklo_epi8(x[2*i], x[2*i+1]);
a[2*i+1] = _mm256_unpackhi_epi8(x[2*i], x[2*i+1]);
}
#pragma GCC unroll 2
for (U i = 0; i < 2; ++i) {
#pragma GCC unroll 2
for (U j = 0; j < 2; ++j) {
b[4*i+2*j] = _mm256_unpacklo_epi16(
a[4*i+j], a[4*i+j+2]
);
b[4*i+2*j+1] = _mm256_unpackhi_epi16(
a[4*i+j], a[4*i+j+2]
);
}
}
// 保留 lane 内布局,省去跨 lane 重排。
#pragma GCC unroll 4
for (U i = 0; i < 4; ++i) {
V lo = _mm256_unpacklo_epi32(b[i], b[i+4]);
V hi = _mm256_unpackhi_epi32(b[i], b[i+4]);
_mm256_store_si256((V*)(out + 64*i), lo);
_mm256_store_si256((V*)(out + 64*i + 32), hi);
}
a[0] = _mm256_unpacklo_epi8(x[8], x[9]);
a[1] = _mm256_unpackhi_epi8(x[8], x[9]);
a[2] = _mm256_unpacklo_epi8(x[10], x[11]);
a[3] = _mm256_unpackhi_epi8(x[10], x[11]);
b[0] = _mm256_unpacklo_epi16(a[0], a[2]);
b[1] = _mm256_unpackhi_epi16(a[0], a[2]);
b[2] = _mm256_unpacklo_epi16(a[1], a[3]);
b[3] = _mm256_unpackhi_epi16(a[1], a[3]);
#pragma GCC unroll 4
for (U i = 0; i < 4; ++i)
_mm256_store_si256((V*)(tail + 32*i), b[i]);
}
// x 已排序,末尾至少有一个 255 哨兵。
inline W insert_byte(W x, U byte) {
return _mm_min_epu8(
x,
_mm_max_epu8(
_mm_slli_si128(x, 1),
_mm_set1_epi8(char(byte))
)
);
}
[[gnu::noinline]]
inline void rare_column(U h, U c, U* dst, U p) {
B copy[32];
for (U j = 0; j < c; ++j)
copy[j] = columns[256*j + h];
std::sort(copy, copy + c);
for (U j = 0; j < c; ++j)
dst[j] = p | copy[j];
}
template<bool Room>
[[gnu::always_inline]]
inline void emit_block(
const B* first,
const B* second,
const U* count,
U block,
U*& dst,
U base,
U* end
) {
V prefix = _mm256_set1_epi32(base | (block << 8));
#pragma GCC unroll 1
for (U half = 0; half < 2; ++half) {
#pragma GCC unroll 1
for (U group = 0; group < 4; ++group) {
fastsort::repeat<4>(
[&](auto Q) __attribute__((always_inline)) {
constexpr U q = decltype(Q)::value;
U h = block + 16*half + 4*group + q;
// count[h] = h + 256*c,已确认 c <= 32。
U c = reinterpret_cast<
const unsigned char*
>(count)[4*h + 1];
U p = base | (h << 8);
const B* f = first
+ 16*half + 64*group
+ 32*(q/2) + 8*(q%2);
const B* s = second
+ 16*half + 32*group + 4*q;
if (__builtin_expect(c <= 12, 1)) {
W x = _mm_loadl_epi64((const W*)f);
uint32_t four;
std::memcpy(&four, s, 4);
V wide = _mm256_or_si256(
_mm256_cvtepu8_epi32(x),
prefix
);
W y = _mm_or_si128(
_mm_cvtepu8_epi32(
_mm_cvtsi32_si128(four)
),
_mm256_castsi256_si128(prefix)
);
if (Room ||
__builtin_expect(end - dst >= 12, 1)) {
__asm__ volatile(
"vmovdqu {%1, %0|%0, %1}"
: "=m"(
*reinterpret_cast<__m256i_u*>(dst)
)
: "x"(wide)
);
_mm_storeu_si128((W*)(dst + 8), y);
} else {
alignas(32) U copy[12];
_mm256_store_si256((V*)copy, wide);
_mm_store_si128((W*)(copy + 8), y);
std::memcpy(dst, copy, c * 4);
}
} else if (c <= 16) {
W x = _mm_loadl_epi64((const W*)f);
uint32_t four;
std::memcpy(&four, s, 4);
x = _mm_unpacklo_epi64(
x,
_mm_cvtsi64_si128(
uint64_t(four) | 0xffffffff00000000ull
)
);
for (U j = 12; j < c; ++j)
x = insert_byte(x, columns[256*j + h]);
write8<Room>(x, dst, 8, p, end);
write8<Room>(
_mm_srli_si128(x, 8),
dst + 8,
c - 8,
p,
end
);
} else {
rare_column(h, c, dst, p);
}
dst += c;
prefix = _mm256_add_epi32(
prefix, _mm256_set1_epi32(256)
);
}
);
}
}
}
// 编码值 > 8447 等价于 c > 32。
[[gnu::noinline, gnu::cold]]
static bool exceeds32(const U* count) {
V maximum = _mm256_load_si256(
reinterpret_cast<const V*>(count)
);
#pragma GCC unroll 1
for (U k = 8; k < 256; k += 8) {
maximum = _mm256_max_epu32(
maximum,
_mm256_load_si256(
reinterpret_cast<const V*>(count + k)
)
);
}
V bad = _mm256_cmpgt_epi32(
maximum, _mm256_set1_epi32(8447)
);
return !_mm256_testz_si256(bad, bad);
}
[[gnu::always_inline]]
inline bool counts_fit32(const U* count) {
V b0 = _mm256_load_si256(
reinterpret_cast<const V*>(count)
);
V b1 = _mm256_load_si256(
reinterpret_cast<const V*>(count + 8)
);
V b2 = _mm256_load_si256(
reinterpret_cast<const V*>(count + 16)
);
V b3 = _mm256_load_si256(
reinterpret_cast<const V*>(count + 24)
);
#pragma GCC unroll 1
for (U k = 32; k < 256; k += 32) {
b0 = _mm256_or_si256(
b0, _mm256_load_si256(
reinterpret_cast<const V*>(count + k)
)
);
b1 = _mm256_or_si256(
b1, _mm256_load_si256(
reinterpret_cast<const V*>(count + k + 8)
)
);
b2 = _mm256_or_si256(
b2, _mm256_load_si256(
reinterpret_cast<const V*>(count + k + 16)
)
);
b3 = _mm256_or_si256(
b3, _mm256_load_si256(
reinterpret_cast<const V*>(count + k + 24)
)
);
}
V bits = _mm256_or_si256(
_mm256_or_si256(b0, b1),
_mm256_or_si256(b2, b3)
);
// 没有高于低 13 位的置位,则所有 c < 32。
if (__builtin_expect(
_mm256_testz_si256(bits, _mm256_set1_epi32(-8192)),
1
)) {
return true;
}
// 精确检查,继续允许 c == 32。
return !exceeds32(count);
}
// 只有独立源、正向输出路径可以传入扩展的 end。
inline void columns_merge(
const H* src, U n, U* dst, U base, U* end
) {
if (n < 64) {
U copy[64];
for (U i = 0; i < n; ++i)
copy[i] = base | load16(src + i);
std::sort(copy, copy + n);
std::memcpy(dst, copy, n * 4);
return;
}
if (n > 8192) {
dense(src, n, dst, base);
return;
}
alignas(64) U count[256];
std::memset(columns, 255, 3072);
for (U i = 0; i < 256; ++i)
count[i] = i;
for (U i = 0; i < n; ++i) {
U v = load16(src + i);
U h = v >> 8;
U at = count[h];
count[h] = at + 256;
columns[at] = B(v);
}
if (!counts_fit32(count)) {
dense(src, n, dst, base);
return;
}
const bool room = end - dst >= n + 12;
for (U block = 0; block < 256; block += 32) {
alignas(32) B first[256], second[128];
batch_sort(columns + block, first, second);
if (room) {
emit_block<true>(
first, second, count, block, dst, base, end
);
} else {
emit_block<false>(
first, second, count, block, dst, base, end
);
}
}
}
} // namespace v2_tail
namespace fastsort {
using U = unsigned;
using B = unsigned char;
using H = uint16_t;
using Z = uint64_t;
static B* mem;
inline U get(const B* p) {
U x;
std::memcpy(&x, p, 4);
return x & 0xffffff;
}
template<class F>
[[gnu::always_inline]]
inline void scan24(const B* __restrict__ s, U n, F&& f) {
U blocks = n / 84;
U rem = n % 84;
for (U b = 0; b < blocks; ++b, s += 256) {
_mm_prefetch((const char*)(s + 256), _MM_HINT_T0);
#pragma GCC unroll 1
for (U j = 0; j < 84; j += 12) {
repeat<12>(
[&](auto q) __attribute__((always_inline)) {
f(get(s + 3 * (j + q)));
}
);
}
}
for (U j = 0; j < rem; ++j)
f(get(s + 3*j));
}
inline bool fits(B** p, B** e) {
__m256i bad = _mm256_setzero_si256();
for (U k = 0; k < 256; k += 4) {
bad = _mm256_or_si256(
bad,
_mm256_cmpgt_epi64(
_mm256_load_si256((__m256i*)(p + k)),
_mm256_load_si256((__m256i*)(e + k))
)
);
}
return _mm256_testz_si256(bad, bad);
}
bool split(
U* __restrict__ a, U n, U* cap, B** start, U* cnt
) {
constexpr U guard = 25344;
alignas(64) B cache[65536] = {};
alignas(64) B* p[256];
alignas(64) B* e[256];
U pos[256];
size_t off = 0;
for (U k = 0; k < 256; ++k) {
start[k] = p[k] = mem + off;
e[k] = p[k] + ((size_t(cap[k]) + 83) / 84) * 256;
off = size_t(e[k] - mem) + guard;
pos[k] = k * 256;
}
auto push = [&](U x) __attribute__((always_inline)) {
U k = x >> 24;
U t = pos[k];
std::memcpy(cache + t, &x, 4);
t += 3;
if (__builtin_expect((t & 255) == 252, 0)) {
t -= 252;
repeat<8>(
[&](auto j) __attribute__((always_inline)) {
_mm256_stream_si256(
(__m256i*)(p[k] + 32*j),
_mm256_load_si256(
(const __m256i*)(cache + t + 32*j)
)
);
}
);
p[k] += 256;
}
pos[k] = t;
};
for (U i = 0; i < n;) {
U stop = std::min(n, i + 8192);
for (; i + 8 <= stop; i += 8) {
repeat<8>(
[&](auto j) __attribute__((always_inline)) {
push(a[i + j]);
}
);
}
for (; i < stop; ++i)
push(a[i]);
if (!fits(p, e)) {
_mm_sfence();
return false;
}
}
for (U k = 0; k < 256; ++k) {
U r = pos[k] & 255;
cnt[k] = U((p[k] - start[k]) / 256) * 84 + r / 3;
if (r) {
std::memcpy(p[k], cache + k*256, 256);
p[k] += 256;
}
}
_mm_sfence();
return fits(p, e);
}
void leaf(
const B* s, U n, U base, U* d, U* end = nullptr
) {
v2_tail::columns_merge(
(const H*)s, n, d, base, end ? end : d + n
);
}
// 调用时 src 与 dst 具有相同的 64 字节内偏移。
[[gnu::always_inline]]
inline void copy_nt(
U* __restrict__ dst,
const U* __restrict__ src,
U n
) {
while (n && (uintptr_t(dst) & 63)) {
*dst++ = *src++;
--n;
}
#pragma GCC unroll 1
for (; n >= 64; n -= 64, dst += 64, src += 64) {
repeat<8>(
[&](auto J) __attribute__((always_inline)) {
constexpr U j = decltype(J)::value;
__m256i x = _mm256_load_si256(
reinterpret_cast<const __m256i*>(src + 8*j)
);
_mm256_stream_si256(
reinterpret_cast<__m256i*>(dst + 8*j), x
);
}
);
}
#pragma GCC unroll 1
for (; n >= 16; n -= 16, dst += 16, src += 16) {
__m256i x0 = _mm256_load_si256(
reinterpret_cast<const __m256i*>(src)
);
__m256i x1 = _mm256_load_si256(
reinterpret_cast<const __m256i*>(src + 8)
);
_mm256_stream_si256(
reinterpret_cast<__m256i*>(dst), x0
);
_mm256_stream_si256(
reinterpret_cast<__m256i*>(dst + 8), x1
);
}
while (n--)
*dst++ = *src++;
}
inline void leaf_stream(
const B* s, U n, U base, U* d, U* end
) {
if (n < 64 || n > 8192) {
leaf(s, n, base, d, end);
return;
}
alignas(64) static U scratch[8192 + 32];
U* local = scratch + ((uintptr_t(d) & 63) >> 2);
v2_tail::columns_merge(
reinterpret_cast<const H*>(s),
n,
local,
base,
local + n + 12
);
copy_nt(d, local, n);
}
template<U Group = 84>
[[gnu::always_inline]]
inline void scatter24_bytes(
const B* __restrict__ s,
U n,
B** __restrict__ p
) {
static_assert(
Group > 0 && 84 % Group == 0,
"Group must divide 84"
);
const U blocks = n / 84;
const U remain = n - blocks * 84;
for (U block = 0; block < blocks; ++block, s += 256) {
_mm_prefetch(
reinterpret_cast<const char*>(s + 256),
_MM_HINT_T0
);
#pragma GCC unroll 1
for (U j = 0; j < 84; j += Group) {
repeat<Group>(
[&](auto Q) __attribute__((always_inline)) {
constexpr U q = decltype(Q)::value;
const B* src = s + 3 * (j + q);
const U h = src[2];
H lo;
std::memcpy(&lo, src, sizeof(lo));
B* const out = p[h];
p[h] = out + sizeof(lo);
std::memcpy(out, &lo, sizeof(lo));
}
);
}
}
for (U j = 0; j < remain; ++j) {
const B* src = s + 3*j;
const U h = src[2];
H lo;
std::memcpy(&lo, src, sizeof(lo));
B* const out = p[h];
p[h] = out + sizeof(lo);
std::memcpy(out, &lo, sizeof(lo));
}
}
void middle(
const B* __restrict__ s,
U n,
U base,
U* __restrict__ d
) {
if (n < 4096) {
U i = 0;
scan24(s, n, [&](U x) {
d[i++] = base | x;
});
std::sort(d, d + n);
return;
}
U cnt[256] = {};
U begin[257], pos[256];
if (n > 600000 || n < 16384) {
scan24(s, n, [&](U x) {
++cnt[x >> 16];
});
begin[0] = 0;
for (U k = 0; k < 256; ++k) {
pos[k] = begin[k];
begin[k+1] = begin[k] + cnt[k];
}
scan24(s, n, [&](U x) {
H lo = H(x);
std::memcpy(
(B*)d + 2 * pos[x >> 16]++,
&lo,
2
);
});
// 原地从 16 位展开为 32 位,必须倒序。
for (int k = 255; k >= 0; --k) {
if (cnt[k]) {
leaf(
(B*)d + 2 * begin[k],
cnt[k],
base | (U(k) << 16),
d + begin[k]
);
}
}
return;
}
B* tmp = mem + 380000000;
B* p[256];
U bytes = 0;
for (U k = 0; k < 256; ++k) {
begin[k] = bytes;
p[k] = tmp + bytes;
bytes += 64 * (
((((n + 255) / 256 * 5 / 4 + 32) * 2 + 63) / 64)
| 1
);
}
begin[256] = bytes;
for (;;) {
scatter24_bytes<84>(s, n, p);
bool bad = false;
for (U k = 0; k < 256; ++k) {
cnt[k] = U(p[k] - (tmp + begin[k])) / 2;
bad |= p[k] > tmp + begin[k+1];
}
if (!bad)
break;
// 根据精确计数重新分配,完整重跑。
bytes = 0;
for (U k = 0; k < 256; ++k) {
begin[k] = bytes;
p[k] = tmp + bytes;
bytes += 2 * cnt[k];
}
begin[256] = bytes;
}
U* end = d + n;
for (U k = 0; k < 256; ++k) {
if (cnt[k]) {
leaf_stream(
tmp + begin[k],
cnt[k],
base | (k << 16),
d,
end
);
}
d += cnt[k];
}
}
} // namespace fastsort
void sort(unsigned* a, int n) {
using namespace fastsort;
if (n < 2)
return;
if (n < 4096 || n > 100000000) {
std::sort(a, a + n);
return;
}
void* raw = std::malloc(384000063);
if (!raw) {
std::sort(a, a + n);
return;
}
mem = (B*)((uintptr_t(raw) + 63) & ~uintptr_t(63));
U cap[256], cnt[256];
B* start[256];
for (U k = 0; k < 256; ++k)
cap[k] = U((Z(n) + 255) / 256) * 6 / 5 + 32;
if (!split(a, n, cap, start, cnt)) {
std::memset(cap, 0, sizeof cap);
for (int i = 0; i < n; ++i)
++cap[a[i] >> 24];
split(a, n, cap, start, cnt);
}
U off = 0;
for (U k = 0; k < 256; ++k) {
if (cnt[k])
middle(start[k], cnt[k], k << 24, a + off);
off += cnt[k];
}
_mm_sfence();
std::free(raw);
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 380.41 ms | 673 MB + 756 KB | Accepted | Score: 100 | 显示更多 |