Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 30 additions & 19 deletions libvmaf/src/feature/x86/adm_avx2.c
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,17 @@
#define MIN(x, y) (((x) < (y)) ? (x) : (y))
#define MAX(x, y) (((x) > (y)) ? (x) : (y))

static inline int64_t mm_hadd_epi64(__m128i v)
{
#if defined(__x86_64__) || defined(_M_X64)
return (int64_t)_mm_extract_epi64(v, 0) + (int64_t)_mm_extract_epi64(v, 1);
#else
int64_t lanes[2];
_mm_storeu_si128((__m128i *)lanes, v);
return lanes[0] + lanes[1];
#endif
}

#define shift15_64b_signExt_256(a, r) \
{ \
r = _mm256_add_epi64( _mm256_srli_epi64(a, 15) , _mm256_and_si256(a, _mm256_set1_epi64x(0xFFFE000000000000))); \
Expand Down Expand Up @@ -822,12 +833,12 @@ void adm_decouple_avx2(AdmBuffer *buf, int w, int h, int stride,
__m256 od_inv_64 = _mm256_mul_ps(inv_64, _mm256_cvtepi32_ps(od));
__m256 rst_d_f = _mm256_mul_ps(kd_inv_32768, od_inv_64);

__m256i gt0_rst_h_f = (__m256i)(_mm256_cmp_ps(rst_h_f, _mm256_setzero_ps(), 14));
__m256i lt0_rst_h_f = (__m256i)(_mm256_cmp_ps(rst_h_f, _mm256_setzero_ps(), 1));
__m256i gt0_rst_v_f = (__m256i)(_mm256_cmp_ps(rst_v_f, _mm256_setzero_ps(), 14));
__m256i lt0_rst_v_f = (__m256i)(_mm256_cmp_ps(rst_v_f, _mm256_setzero_ps(), 1));
__m256i gt0_rst_d_f = (__m256i)(_mm256_cmp_ps(rst_d_f, _mm256_setzero_ps(), 14));
__m256i lt0_rst_d_f = (__m256i)(_mm256_cmp_ps(rst_d_f, _mm256_setzero_ps(), 1));
__m256i gt0_rst_h_f = _mm256_castps_si256(_mm256_cmp_ps(rst_h_f, _mm256_setzero_ps(), 14));
__m256i lt0_rst_h_f = _mm256_castps_si256(_mm256_cmp_ps(rst_h_f, _mm256_setzero_ps(), 1));
__m256i gt0_rst_v_f = _mm256_castps_si256(_mm256_cmp_ps(rst_v_f, _mm256_setzero_ps(), 14));
__m256i lt0_rst_v_f = _mm256_castps_si256(_mm256_cmp_ps(rst_v_f, _mm256_setzero_ps(), 1));
__m256i gt0_rst_d_f = _mm256_castps_si256(_mm256_cmp_ps(rst_d_f, _mm256_setzero_ps(), 14));
__m256i lt0_rst_d_f = _mm256_castps_si256(_mm256_cmp_ps(rst_d_f, _mm256_setzero_ps(), 1));

__m256i mask_min_max_h = _mm256_or_si256(gt0_rst_h_f, lt0_rst_h_f);
__m256i mask_min_max_v = _mm256_or_si256(gt0_rst_v_f, lt0_rst_v_f);
Expand All @@ -837,7 +848,7 @@ void adm_decouple_avx2(AdmBuffer *buf, int w, int h, int stride,
__m256i mask_rst_v = _mm256_and_si256(mask_min_max_v, angle_flag);
__m256i mask_rst_d = _mm256_and_si256(mask_min_max_d, angle_flag);

__m256d adm_gain_d = _mm256_set1_pd(adm_enhn_gain_limit);
__m256d adm_gain_d = _mm256_set1_pd(adm_enhn_gain_limit);
__m256d rst_h_gainlo_d = _mm256_mul_pd(_mm256_cvtepi32_pd(_mm256_extractf128_si256(rst_h, 0)), adm_gain_d);
__m256d rst_h_gainhi_d = _mm256_mul_pd(_mm256_cvtepi32_pd(_mm256_extractf128_si256(rst_h, 1)), adm_gain_d);
__m256i rst_h_gain = _mm256_insertf128_si256(_mm256_castsi128_si256(_mm256_cvtpd_epi32(rst_h_gainlo_d)), _mm256_cvtpd_epi32(rst_h_gainhi_d),1);
Expand Down Expand Up @@ -2152,15 +2163,15 @@ float adm_cm_avx2(AdmBuffer *buf, int w, int h, int src_stride, int csf_a_stride
}
accum_inner_h_lo_256 = _mm256_add_epi64(accum_inner_h_lo_256, accum_inner_h_hi_256);
__m128i r2_h = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_h_lo_256), _mm256_extracti128_si256(accum_inner_h_lo_256, 1));
int64_t res_h = r2_h[0] + r2_h[1];
int64_t res_h = mm_hadd_epi64(r2_h);

accum_inner_v_lo_256 = _mm256_add_epi64(accum_inner_v_lo_256, accum_inner_v_hi_256);
__m128i r2_v = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_v_lo_256), _mm256_extracti128_si256(accum_inner_v_lo_256, 1));
int64_t res_v = r2_v[0] + r2_v[1];
int64_t res_v = mm_hadd_epi64(r2_v);

accum_inner_d_lo_256 = _mm256_add_epi64(accum_inner_d_lo_256, accum_inner_d_hi_256);
__m128i r2_d = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_d_lo_256), _mm256_extracti128_si256(accum_inner_d_lo_256, 1));
int64_t res_d = r2_d[0] + r2_d[1];
int64_t res_d = mm_hadd_epi64(r2_d);

for (j = end_col_mod6; j < end_col; ++j) {
xh = src->band_h[i * src_stride + j] * i_rfactor[0];
Expand Down Expand Up @@ -2610,13 +2621,13 @@ float i4_adm_cm_avx2(AdmBuffer *buf, int w, int h, int src_stride, int csf_a_str
}

__m128i r2_h = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_h_256), _mm256_extracti128_si256(accum_inner_h_256, 1));
int64_t res_h = r2_h[0] + r2_h[1];
int64_t res_h = mm_hadd_epi64(r2_h);

__m128i r2_v = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_v_256), _mm256_extracti128_si256(accum_inner_v_256, 1));
int64_t res_v = r2_v[0] + r2_v[1];
int64_t res_v = mm_hadd_epi64(r2_v);

__m128i r2_d = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_d_256), _mm256_extracti128_si256(accum_inner_d_256, 1));
int64_t res_d = r2_d[0] + r2_d[1];
int64_t res_d = mm_hadd_epi64(r2_d);

for (j = end_col_mod2; j < end_col; ++j)
{
Expand Down Expand Up @@ -3689,15 +3700,15 @@ float adm_csf_den_scale_avx2(const adm_dwt_band_t *src, int w, int h,

accum_inner_h_lo = _mm256_add_epi64(accum_inner_h_lo, accum_inner_h_hi);
__m128i h_r2 = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_h_lo), _mm256_extracti128_si256(accum_inner_h_lo, 1));
uint64_t h_r1 = h_r2[0] + h_r2[1];
uint64_t h_r1 = mm_hadd_epi64(h_r2);

accum_inner_v_lo = _mm256_add_epi64(accum_inner_v_lo, accum_inner_v_hi);
__m128i v_r2 = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_v_lo), _mm256_extracti128_si256(accum_inner_v_lo, 1));
uint64_t v_r1 = v_r2[0] + v_r2[1];
uint64_t v_r1 = mm_hadd_epi64(v_r2);

accum_inner_d_lo = _mm256_add_epi64(accum_inner_d_lo, accum_inner_d_hi);
__m128i d_r2 = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_d_lo), _mm256_extracti128_si256(accum_inner_d_lo, 1));
uint64_t d_r1 = d_r2[0] + d_r2[1];
uint64_t d_r1 = mm_hadd_epi64(d_r2);

for (int j = right_mod_8; j < right; ++j) {
uint16_t h = (uint16_t)abs(src_h[j]);
Expand Down Expand Up @@ -4157,13 +4168,13 @@ float adm_csf_den_s123_avx2(const i4_adm_dwt_band_t *src, int scale, int w, int
accum_inner_d_256 = _mm256_add_epi64(accum_inner_d_256, d_cu);
}
__m128i h_r2 = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_h_256), _mm256_extracti128_si256(accum_inner_h_256, 1));
uint64_t h_r1 = h_r2[0] + h_r2[1];
uint64_t h_r1 = mm_hadd_epi64(h_r2);

__m128i d_r2 = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_d_256), _mm256_extracti128_si256(accum_inner_d_256, 1));
uint64_t d_r1 = d_r2[0] + d_r2[1];
uint64_t d_r1 = mm_hadd_epi64(d_r2);

__m128i v_r2 = _mm_add_epi64(_mm256_castsi256_si128(accum_inner_v_256), _mm256_extracti128_si256(accum_inner_v_256, 1));
uint64_t v_r1 = v_r2[0] + v_r2[1];
uint64_t v_r1 = mm_hadd_epi64(v_r2);

for (int j = right_mod_4; j < right; ++j)
{
Expand Down
35 changes: 23 additions & 12 deletions libvmaf/src/feature/x86/adm_avx512.c
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,17 @@
#define MIN(x, y) (((x) < (y)) ? (x) : (y))
#define MAX(x, y) (((x) > (y)) ? (x) : (y))

static inline int64_t mm_hadd_epi64(__m128i v)
{
#if defined(__x86_64__) || defined(_M_X64)
return (int64_t)_mm_extract_epi64(v, 0) + (int64_t)_mm_extract_epi64(v, 1);
#else
int64_t lanes[2];
_mm_storeu_si128((__m128i *)lanes, v);
return lanes[0] + lanes[1];
#endif
}

#define COS_1DEG_SQ cos(1.0 * M_PI / 180.0) * cos(1.0 * M_PI / 180.0)

#if 1
Expand Down Expand Up @@ -1807,17 +1818,17 @@ float adm_cm_avx512(AdmBuffer *buf, int w, int h, int src_stride, int csf_a_stri
accum_inner_h_lo_512 = _mm512_add_epi64(accum_inner_h_lo_512, accum_inner_h_hi_512);
__m256i r4_h = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_h_lo_512), _mm512_extracti64x4_epi64(accum_inner_h_lo_512, 1));
__m128i r2_h = _mm_add_epi64(_mm256_castsi256_si128(r4_h), _mm256_extracti128_si256(r4_h, 1));
int64_t res_h = r2_h[0] + r2_h[1];
int64_t res_h = mm_hadd_epi64(r2_h);

accum_inner_v_lo_512 = _mm512_add_epi64(accum_inner_v_lo_512, accum_inner_v_hi_512);
__m256i r4_v = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_v_lo_512), _mm512_extracti64x4_epi64(accum_inner_v_lo_512, 1));
__m128i r2_v = _mm_add_epi64(_mm256_castsi256_si128(r4_v), _mm256_extracti128_si256(r4_v, 1));
int64_t res_v = r2_v[0] + r2_v[1];
int64_t res_v = mm_hadd_epi64(r2_v);

accum_inner_d_lo_512 = _mm512_add_epi64(accum_inner_d_lo_512, accum_inner_d_hi_512);
__m256i r4_d = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_d_lo_512), _mm512_extracti64x4_epi64(accum_inner_d_lo_512, 1));
__m128i r2_d = _mm_add_epi64(_mm256_castsi256_si128(r4_d), _mm256_extracti128_si256(r4_d, 1));
int64_t res_d = r2_d[0] + r2_d[1];
int64_t res_d = mm_hadd_epi64(r2_d);

for (j = end_col_mod14; j < end_col; ++j) {
xh = src->band_h[i * src_stride + j] * i_rfactor[0];
Expand Down Expand Up @@ -2254,15 +2265,15 @@ float i4_adm_cm_avx512(AdmBuffer *buf, int w, int h, int src_stride, int csf_a_s

__m256i r4_h = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_h_512), _mm512_extracti64x4_epi64(accum_inner_h_512, 1));
__m128i r2_h = _mm_add_epi64(_mm256_castsi256_si128(r4_h), _mm256_extracti128_si256(r4_h, 1));
int64_t res_h = r2_h[0] + r2_h[1];
int64_t res_h = mm_hadd_epi64(r2_h);

__m256i r4_v = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_v_512), _mm512_extracti64x4_epi64(accum_inner_v_512, 1));
__m128i r2_v = _mm_add_epi64(_mm256_castsi256_si128(r4_v), _mm256_extracti128_si256(r4_v, 1));
int64_t res_v = r2_v[0] + r2_v[1];
int64_t res_v = mm_hadd_epi64(r2_v);

__m256i r4_d = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_d_512), _mm512_extracti64x4_epi64(accum_inner_d_512, 1));
__m128i r2_d = _mm_add_epi64(_mm256_castsi256_si128(r4_d), _mm256_extracti128_si256(r4_d, 1));
int64_t res_d = r2_d[0] + r2_d[1];
int64_t res_d = mm_hadd_epi64(r2_d);

for (j = end_col_mod6; j < end_col; ++j)
{
Expand Down Expand Up @@ -4006,17 +4017,17 @@ float adm_csf_den_scale_avx512(const adm_dwt_band_t *src, int w, int h,
accum_inner_h_lo = _mm512_add_epi64(accum_inner_h_lo, accum_inner_h_hi);
__m256i h_r4 = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_h_lo), _mm512_extracti64x4_epi64(accum_inner_h_lo, 1));
__m128i h_r2 = _mm_add_epi64(_mm256_castsi256_si128(h_r4), _mm256_extracti64x2_epi64(h_r4, 1));
uint64_t h_r1 = h_r2[0] + h_r2[1];
uint64_t h_r1 = mm_hadd_epi64(h_r2);

accum_inner_v_lo = _mm512_add_epi64(accum_inner_v_lo, accum_inner_v_hi);
__m256i v_r4 = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_v_lo), _mm512_extracti64x4_epi64(accum_inner_v_lo, 1));
__m128i v_r2 = _mm_add_epi64(_mm256_castsi256_si128(v_r4), _mm256_extracti64x2_epi64(v_r4, 1));
uint64_t v_r1 = v_r2[0] + v_r2[1];
uint64_t v_r1 = mm_hadd_epi64(v_r2);

accum_inner_d_lo = _mm512_add_epi64(accum_inner_d_lo, accum_inner_d_hi);
__m256i d_r4 = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_d_lo), _mm512_extracti64x4_epi64(accum_inner_d_lo, 1));
__m128i d_r2 = _mm_add_epi64(_mm256_castsi256_si128(d_r4), _mm256_extracti64x2_epi64(d_r4, 1));
uint64_t d_r1 = d_r2[0] + d_r2[1];
uint64_t d_r1 = mm_hadd_epi64(d_r2);

for (int j = right_mod_16; j < right; ++j) {
uint16_t h = (uint16_t)abs(src_h[j]);
Expand Down Expand Up @@ -4147,15 +4158,15 @@ float adm_csf_den_s123_avx512(const i4_adm_dwt_band_t *src, int scale, int w, in
}
__m256i h_r4 = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_h_512), _mm512_extracti64x4_epi64(accum_inner_h_512, 1));
__m128i h_r2 = _mm_add_epi64(_mm256_castsi256_si128(h_r4), _mm256_extracti64x2_epi64(h_r4, 1));
uint64_t h_r1 = h_r2[0] + h_r2[1];
uint64_t h_r1 = mm_hadd_epi64(h_r2);

__m256i d_r4 = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_d_512), _mm512_extracti64x4_epi64(accum_inner_d_512, 1));
__m128i d_r2 = _mm_add_epi64(_mm256_castsi256_si128(d_r4), _mm256_extracti64x2_epi64(d_r4, 1));
uint64_t d_r1 = d_r2[0] + d_r2[1];
uint64_t d_r1 = mm_hadd_epi64(d_r2);

__m256i v_r4 = _mm256_add_epi64(_mm512_castsi512_si256(accum_inner_v_512), _mm512_extracti64x4_epi64(accum_inner_v_512, 1));
__m128i v_r2 = _mm_add_epi64(_mm256_castsi256_si128(v_r4), _mm256_extracti64x2_epi64(v_r4, 1));
uint64_t v_r1 = v_r2[0] + v_r2[1];
uint64_t v_r1 = mm_hadd_epi64(v_r2);

for (int j = right_mod_8; j < right; ++j)
{
Expand Down
Loading