blob: d1895848fb7741db575d8e56bedad2fcc319c9c4 [file]
// Copyright 2020 Google LLC
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// https://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// -----------------------------------------------------------------------------
//
// ANS DSP
//
// Author: Skal (pascal.massimino@gmail.com)
#include <cassert>
#include <cstdint>
#include "src/dsp/dsp.h"
#include "src/dsp/math.h"
namespace WP2 {
namespace {
//------------------------------------------------------------------------------
// C version
void ANSUpdateCDF_C(uint32_t n, const uint16_t cdf_base[],
const uint16_t cdf_var[], uint32_t mult, uint16_t cumul[]) {
if (mult == 65535) { // special case 'copy'
for (uint32_t i = 1; i < n; ++i) {
const int32_t v = (int32_t)cdf_var[i] + cdf_base[i] - cumul[i];
cumul[i] += v;
}
} else {
for (uint32_t i = 1; i < n; ++i) {
const int32_t v = (int32_t)cdf_var[i] + cdf_base[i] - cumul[i];
cumul[i] += RightShift(v * mult, 16);
}
}
for (uint32_t i = 1; i < n; ++i) assert(cumul[i] >= cumul[i - 1]);
}
uint32_t ANSGetSymbol_C(uint32_t max_symbol, const uint16_t cumul[],
uint32_t proba) {
(void)max_symbol;
uint32_t s = 0;
#if 0
// loop version
for (; s < max_symbol - 1; ++s) {
if (proba < cumul_[s + 1]) break;
}
#else
// binary search
#if (APROBA_MAX_SYMBOL > 16)
if (proba >= cumul[16]) s += 16;
#endif
#if (APROBA_MAX_SYMBOL > 8)
if (proba >= cumul[s + 8]) s += 8;
#endif
#if (APROBA_MAX_SYMBOL > 4)
if (proba >= cumul[s + 4]) s += 4;
#endif
if (proba >= cumul[s + 2]) s += 2;
if (proba >= cumul[s + 1]) s += 1;
#endif
return s;
}
//------------------------------------------------------------------------------
// SSE version
#if defined(WP2_USE_SSE)
void ANSUpdateCDF_SSE(uint32_t n, const uint16_t cdf_base[],
const uint16_t cdf_var[], uint32_t mult,
uint16_t cumul[]) {
if (mult == 65535) { // special case 'copy'
for (uint32_t i = 0; i < n; i += 8) {
const __m128i A = _mm_loadu_si128((const __m128i*)&cdf_var[i]);
const __m128i B = _mm_loadu_si128((const __m128i*)&cdf_base[i]);
const __m128i D = _mm_add_epi16(A, B);
_mm_storeu_si128((__m128i*)&cumul[i], D);
}
} else {
assert(mult < 32768);
const __m128i M = _mm_set1_epi16(mult);
for (uint32_t i = 0; i < n; i += 8) {
const __m128i A = _mm_loadu_si128((const __m128i*)&cdf_var[i]);
const __m128i B = _mm_loadu_si128((const __m128i*)&cdf_base[i]);
const __m128i C = _mm_loadu_si128((const __m128i*)&cumul[i]);
const __m128i D = _mm_add_epi16(A, B);
const __m128i E = _mm_sub_epi16(D, C);
const __m128i F = _mm_mulhi_epi16(E, M);
const __m128i I = _mm_add_epi16(C, F);
_mm_storeu_si128((__m128i*)&cumul[i], I);
}
}
}
uint32_t ANSGetSymbol_SSE(uint32_t max_symbol, const uint16_t cumul[],
uint32_t proba) {
(void)max_symbol;
// TODO(yguyon): SSE other than APROBA_MAX_SYMBOL=16
#if (APROBA_MAX_SYMBOL == 16)
const __m128i A0 = _mm_loadu_si128((const __m128i*)&cumul[0]);
const __m128i A1 = _mm_loadu_si128((const __m128i*)&cumul[8]);
const __m128i B = _mm_set1_epi16(proba + 1);
const __m128i C0 = _mm_cmplt_epi16(A0, B); // we'd need mm_cmple_epi16 !!
const __m128i C1 = _mm_cmplt_epi16(A1, B);
const __m128i D = _mm_packs_epi16(C0, C1);
const uint32_t bits = _mm_movemask_epi8(D);
return WP2Log2Floor(bits);
#elif (APROBA_MAX_SYMBOL == 32)
const __m128i A0 = _mm_loadu_si128((const __m128i*)&cumul[0]);
const __m128i A1 = _mm_loadu_si128((const __m128i*)&cumul[8]);
const __m128i A2 = _mm_loadu_si128((const __m128i*)&cumul[16]);
const __m128i A3 = _mm_loadu_si128((const __m128i*)&cumul[24]);
const __m128i B = _mm_set1_epi16(proba + 1);
const __m128i C0 = _mm_cmplt_epi16(A0, B); // we'd need mm_cmple_epi16 !!
const __m128i C1 = _mm_cmplt_epi16(A1, B);
const __m128i C2 = _mm_cmplt_epi16(A2, B);
const __m128i C3 = _mm_cmplt_epi16(A3, B);
const __m128i D0 = _mm_packs_epi16(C0, C1);
const __m128i D1 = _mm_packs_epi16(C2, C3);
const uint32_t E0 = _mm_movemask_epi8(D0);
const uint32_t E1 = _mm_movemask_epi8(D1);
const uint32_t bits = E0 | (E1 << 16);
return WP2Log2Floor(bits);
#else
return ANSGetSymbol_C(max_symbol, cumul, proba);
#endif
}
#endif // WP2_USE_SSE
} // namespace
//------------------------------------------------------------------------------
ANSUpdateCDFFunc ANSUpdateCDF = nullptr;
ANSGetSymbolFunc ANSGetSymbol = nullptr;
static volatile WP2CPUInfo ans_last_cpuinfo_used =
(WP2CPUInfo)&ans_last_cpuinfo_used;
WP2_TSAN_IGNORE_FUNCTION void ANSInit() {
if (ans_last_cpuinfo_used == WP2GetCPUInfo) return;
ANSUpdateCDF = ANSUpdateCDF_C;
ANSGetSymbol = ANSGetSymbol_C;
if (WP2GetCPUInfo != nullptr) {
#if defined(WP2_USE_SSE)
if (WP2GetCPUInfo(kSSE)) {
ANSUpdateCDF = ANSUpdateCDF_SSE;
ANSGetSymbol = ANSGetSymbol_SSE;
}
#endif
}
ans_last_cpuinfo_used = WP2GetCPUInfo;
}
//------------------------------------------------------------------------------
} // namespace WP2