// Average value of an area for ARM NEON.
// Written by Nils Liaaen Corneliusen 2024.
// Based on code from 2014: https://www.ignorantus.com/source/average_neon_2014.c
// More information about this implementation: https://www.ignorantus.com/news/2024/#17-April-2024
// License: CC0 1.0 Universal (CC0 1.0) Public Domain Dedication license
#include <stdio.h>
#include <stdint.h>
#include <arm_neon.h>

#define MIN(a,b) ((a)<(b)?(a):(b))
#define MAX(a,b) ((a)>(b)?(a):(b))
#define LO8( X )  vget_low_u8( X )
#define HI8( X )  vget_high_u8( X )
#define LO64( X ) vget_low_u64( X )
#define HI64( X ) vget_high_u64( X )
#define qu64u8( X ) vreinterpretq_u64_u8( X )
#define qu8u64( X ) vreinterpretq_u8_u64( X )
#define vandq_u64x( X, V ) qu8u64( vandq_u64( qu64u8( X ), V ) )

float avg_neon( const uint8_t *src, int srcstride, int srcx, int srcy, int srcw, int srch )
{
    uint64x2_t ff = { 0xffffffffffffffff, 0xffffffffffffffff };
    uint64x2_t lmask, rmask;

    // Alignment mask left
    int leftal = srcx&0x0f;
    int64x2_t lsh = { MIN( leftal, 8 )*8,  MAX( leftal-8, 0 )*8 };
    lmask = vshlq_u64( ff, lsh );

    int w = srcw - MIN( 16-leftal, srcw ); // Width excluding alignment left

    // Overrun mask right
    if( w&~0x0f && !(w&0x0f) ) {
        // Remaining width>=16 and multiple of 16, move 1 block right
        rmask = ff;
        w -= 16;
    } else {
        // Normal overrun
        int rightal = (srcx+srcw)&0x0f;
        int64x2_t rsh = { -MAX( 8-rightal, 0 )*8, -MIN( 16-rightal, 8 )*8 };
        rmask = vshlq_u64( ff, rsh );
        w &= ~0x0f;
    }

    if( leftal+srcw < 16 ) {
        // Everything fits left, merge
        lmask = vandq_u64( lmask, rmask );
        rmask = vdupq_n_u64( 0 );
    }

    const uint8_t *data = (uint8_t *)(((size_t)src+srcx)&~0x0f) + srcy*srcstride;
    uint64x1_t sum = vcreate_u64( 0 );

    // 4 rows at a time
    int y = 0;

    for( ; y + 3 < srch; y += 4 ) {
        const uint8_t * restrict row0 = (uint8_t *)__builtin_assume_aligned( data,             16 );
        const uint8_t * restrict row1 = (uint8_t *)__builtin_assume_aligned( row0 + srcstride, 16 );
        const uint8_t * restrict row2 = (uint8_t *)__builtin_assume_aligned( row1 + srcstride, 16 );
        const uint8_t * restrict row3 = (uint8_t *)__builtin_assume_aligned( row2 + srcstride, 16 );

        // Left align
        uint8x16_t d00 = vld1q_u8( row0 ); row0 += 16; d00 = vandq_u64x( d00, lmask );
        uint8x16_t d01 = vld1q_u8( row1 ); row1 += 16; d01 = vandq_u64x( d01, lmask );
        uint8x16_t d02 = vld1q_u8( row2 ); row2 += 16; d02 = vandq_u64x( d02, lmask );
        uint8x16_t d03 = vld1q_u8( row3 ); row3 += 16; d03 = vandq_u64x( d03, lmask );
        uint16x8_t r00 = vaddl_u8( HI8( d00 ), LO8( d00 ) ); d00 = vld1q_u8( row0 ); row0 += 16;
        uint16x8_t r01 = vaddl_u8( HI8( d01 ), LO8( d01 ) ); d01 = vld1q_u8( row1 ); row1 += 16;
        uint16x8_t r02 = vaddl_u8( HI8( d02 ), LO8( d02 ) ); d02 = vld1q_u8( row2 ); row2 += 16;
        uint16x8_t r03 = vaddl_u8( HI8( d03 ), LO8( d03 ) ); d03 = vld1q_u8( row3 ); row3 += 16;

        // Middle 16-blocks
        for( int x = 0; x < w; x += 16 ) {
            r00 = vaddq_u16( r00, vaddl_u8( HI8( d00 ), LO8( d00 ) ) ); d00 = vld1q_u8( row0 ); row0 += 16;
            r01 = vaddq_u16( r01, vaddl_u8( HI8( d01 ), LO8( d01 ) ) ); d01 = vld1q_u8( row1 ); row1 += 16;
            r02 = vaddq_u16( r02, vaddl_u8( HI8( d02 ), LO8( d02 ) ) ); d02 = vld1q_u8( row2 ); row2 += 16;
            r03 = vaddq_u16( r03, vaddl_u8( HI8( d03 ), LO8( d03 ) ) ); d03 = vld1q_u8( row3 ); row3 += 16;
        }

        // Right overrun
        d00 = vandq_u64x( d00, rmask );
        d01 = vandq_u64x( d01, rmask );
        d02 = vandq_u64x( d02, rmask );
        d03 = vandq_u64x( d03, rmask );
        r00 = vaddq_u16( r00, vaddl_u8( HI8( d00 ), LO8( d00 ) ) );
        r01 = vaddq_u16( r01, vaddl_u8( HI8( d01 ), LO8( d01 ) ) );
        r02 = vaddq_u16( r02, vaddl_u8( HI8( d02 ), LO8( d02 ) ) );
        r03 = vaddq_u16( r03, vaddl_u8( HI8( d03 ), LO8( d03 ) ) );

        // If 1920x1080 all pixels 0xff support not necessary, replace inner paddles with vaddqs
        uint64x2_t rv = vpaddlq_u32( vaddq_u32( vaddq_u32( vpaddlq_u16( r00 ), vpaddlq_u16( r01 ) ),
                                                vaddq_u32( vpaddlq_u16( r02 ), vpaddlq_u16( r03 ) ) ) );

        sum = vadd_u64( sum, vadd_u64( HI64(rv), LO64(rv) ) );

        data += srcstride*4;
    }

    // Remaining rows
    for( ; y < srch; y++ ) {
        const uint8_t * restrict row0 = (uint8_t *)__builtin_assume_aligned( data, 16 );
        // Left align
        uint8x16_t d00 = vld1q_u8( row0 ); row0 += 16; d00 = vandq_u64x( d00, lmask );
        uint16x8_t r00 = vaddl_u8( HI8( d00 ), LO8( d00 ) ); d00 = vld1q_u8( row0 ); row0 += 16;
        // Middle 16-blocks
        for( int x = 0; x < w; x += 16 ) {
            r00 = vaddq_u16( r00, vaddl_u8( HI8( d00 ), LO8( d00 ) ) ); d00 = vld1q_u8( row0 ); row0 += 16;
        }
        // Right overrun
        d00 = vandq_u64x( d00, rmask );
        r00 = vaddq_u16( r00, vaddl_u8( HI8( d00 ), LO8( d00 ) ) );
        // Sum it all up
        uint64x2_t rv = vpaddlq_u32( vpaddlq_u16( r00 ) );
        sum = vadd_u64( sum, vadd_u64( HI64(rv), LO64(rv) ) );

        data += srcstride;
    }

    return (int)(uint64_t)sum / (float)(srcw*srch);
}
