#include "/srv/home/gc-sdk/ipu-stack/device/static_runtime.S"
#include "arch/gc_tile_defines.h"

	.text
	.allow_optimizations

	.set SYNC_COMPUTE_SET, TEXCH_SYNCZONE_LOCAL
	.set REFERENCE_GELU_ALPHA_F16X2, 0x3a623a62
	.set REFERENCE_GELU_BETA_F16X2, 0x29b929b9
	.set ONE_F16X2, 0x3c003c00
	.set HALF_F16X2, 0x38003800
	// Callable supervisor ABI: m2 output, m3 input, m4 physical element
	// count, and m10 return. The contiguous entry accepts whole half pairs.

	.macro REFERENCE_GELU_PAIR value
	// tanh GeLU: x/2 * (1 + tanh(sqrt(2/pi) * (x + 0.044715*x^3))).
	f16v2mul $a2, \value, \value
	f16v2mul $a2, $a2, \value
	f16v2mul $a2, $a4, $a2
	f16v2add $a2, \value, $a2
	f16v2mul $a2, $a5, $a2
	f16v2tanh $a2, $a2
	f16v2add $a2, $a6, $a2
	f16v2mul $a2, $a7, $a2
	f16v2mul $a2, \value, $a2
	.endm

	.macro LOAD_REFERENCE_GELU_CONSTANTS
	ldconst $a4, REFERENCE_GELU_BETA_F16X2
	ldconst $a5, REFERENCE_GELU_ALPHA_F16X2
	ldconst $a6, ONE_F16X2
	ldconst $a7, HALF_F16X2
	.endm

	.macro REFERENCE_GELU_TWO_PAIRS
	f16v4mul $a2:3, $a0:1, $a0:1
	f16v4mul $a2:3, $a2:3, $a0:1
	f16v4mul $a2:3, $a4:BL, $a2:3
	f16v4add $a2:3, $a0:1, $a2:3
	f16v4mul $a2:3, $a5:BL, $a2:3
	f16v2tanh $a2, $a2
	f16v2tanh $a3, $a3
	f16v4add $a2:3, $a6:BL, $a2:3
	f16v4mul $a2:3, $a7:BL, $a2:3
	f16v4mul $a2:3, $a0:1, $a2:3
	.endm

	.macro REFERENCE_GELU_EIGHT_PAIRS
	ld32 $a0, $m3, $m15, 0
	ld32 $a1, $m3, $m15, 1
	REFERENCE_GELU_TWO_PAIRS
	st32 $a2, $m2, $m15, 0
	st32 $a3, $m2, $m15, 1
	ld32 $a0, $m3, $m15, 2
	ld32 $a1, $m3, $m15, 3
	REFERENCE_GELU_TWO_PAIRS
	st32 $a2, $m2, $m15, 2
	st32 $a3, $m2, $m15, 3
	ld32 $a0, $m3, $m15, 4
	ld32 $a1, $m3, $m15, 5
	REFERENCE_GELU_TWO_PAIRS
	st32 $a2, $m2, $m15, 4
	st32 $a3, $m2, $m15, 5
	ld32 $a0, $m3, $m15, 6
	ld32 $a1, $m3, $m15, 7
	REFERENCE_GELU_TWO_PAIRS
	st32 $a2, $m2, $m15, 6
	st32 $a3, $m2, $m15, 7
	.endm

	.macro REFERENCE_GELU_ENTRY symbol, worker
	.section .text.\symbol,"ax",@progbits
	.globl \symbol
	.p2align 2
	.type \symbol,@function
\symbol:
	.supervisor
	add $m11, $m11, -16
	st32 $m2, $m11, $m15, 0
	st32 $m3, $m11, $m15, 1
	st32 $m4, $m11, $m15, 2
	setzi $m0, \worker
	runall $m0, $m11, 0
	sync SYNC_COMPUTE_SET
	add $m11, $m11, 16
	br $m10
	.size \symbol, .-\symbol
	.endm

	REFERENCE_GELU_ENTRY reference_ipu_stack_gelu_tanh_approx_f16, .Lreference_gelu_contiguous_worker

	.worker
	.p2align 3
.Lreference_gelu_contiguous_worker:
	ld32 $m2, $mvertex_base, $m15, 0
	ld32 $m3, $mvertex_base, $m15, 1
	ld32 $m8, $mvertex_base, $m15, 2
	shr $m8, $m8, 1
	get $m4, $WSR
	and $m4, $m4, CSR_W_WSR__CTXTID_M1__MASK
	shl $m4, $m4, 3
	cmpult $m6, $m4, $m8
	brz $m6, .Lreference_gelu_contiguous_done
	LOAD_REFERENCE_GELU_CONSTANTS
	shl $m5, $m4, 2
	add $m2, $m2, $m5
	add $m3, $m3, $m5
	setzi $m5, 8
	add $m6, $m5, -1
	and $m6, $m8, $m6
	brnz $m6, .Lreference_gelu_contiguous_loop
.Lreference_gelu_contiguous_aligned_loop:
	REFERENCE_GELU_EIGHT_PAIRS
	add $m2, $m2, 192
	add $m3, $m3, 192
	add $m4, $m4, 48
	cmpult $m6, $m4, $m8
	brnz $m6, .Lreference_gelu_contiguous_aligned_loop
	bri .Lreference_gelu_contiguous_done
.Lreference_gelu_contiguous_loop:
	sub $m6, $m8, $m4
	cmpult $m6, $m6, $m5
	brnz $m6, .Lreference_gelu_contiguous_tail
	REFERENCE_GELU_EIGHT_PAIRS
	add $m2, $m2, 192
	add $m3, $m3, 192
	add $m4, $m4, 48
	cmpult $m6, $m4, $m8
	brnz $m6, .Lreference_gelu_contiguous_loop
	bri .Lreference_gelu_contiguous_done
.Lreference_gelu_contiguous_tail:
	ld32 $a0, $m3, $m15, 0
	REFERENCE_GELU_PAIR $a0
	st32 $a2, $m2, $m15, 0
	add $m2, $m2, 4
	add $m3, $m3, 4
	add $m4, $m4, 1
	cmpult $m6, $m4, $m8
	brnz $m6, .Lreference_gelu_contiguous_tail
.Lreference_gelu_contiguous_done:
	exitz $m15

#include "/srv/home/gc-sdk/ipu-stack/device/gelu_f16.S"
#include "arch/gc_tile_defines.h"

	.text
	.allow_optimizations
	.set SYNC_COMPUTE_SET, TEXCH_SYNCZONE_LOCAL
	.set RECIPROCAL_3_SHL17, 43691
	// Callable supervisor ABI: m2 output, m3 initial partial, m4 packed remote
	// partials, m5 remote-partial count, m6 physical half-element count, and
	// m10 return. Each remote partial occupies one contiguous element-count row.
	.section .text.reference_ipu_stack_reduce_sum_f16,"ax",@progbits
	.globl reference_ipu_stack_reduce_sum_f16
	.p2align 2
	.type reference_ipu_stack_reduce_sum_f16,@function
reference_ipu_stack_reduce_sum_f16:
	.supervisor
	add $m11, $m11, -24
	st32 $m2, $m11, $m15, 0
	st32 $m3, $m11, $m15, 1
	st32 $m4, $m11, $m15, 2
	st32 $m5, $m11, $m15, 3
	st32 $m6, $m11, $m15, 4
	setzi $m0, .Lreference_reduce_sum_f16_worker
	runall $m0, $m11, 0
	sync SYNC_COMPUTE_SET
	add $m11, $m11, 24
	br $m10
	.size reference_ipu_stack_reduce_sum_f16, .-reference_ipu_stack_reduce_sum_f16

	.worker
	.p2align 3
.Lreference_reduce_sum_f16_worker:
	ld32 $m2, $mvertex_base, $m15, 0
	ld32 $m3, $mvertex_base, $m15, 1
	ld32 $m4, $mvertex_base, $m15, 2
	ld32 $m0, $mvertex_base, $m15, 3
	ld32 $m9, $mvertex_base, $m15, 4
	shr $m9, $m9, 3
	shl $m1, $m9, 1
	get $m5, $WSR
	and $m5, $m5, CSR_W_WSR__CTXTID_M1__MASK
	cmpult $m6, $m5, $m9
	brz $m6, .Lreference_reduce_sum_done
	shl $m10, $m5, 4
	add $m2, $m2, $m10
	add $m3, $m3, $m10
	add $m4, $m4, $m10
	add $m6, $m9, 5
	sub $m6, $m6, $m5
	setzi $m10, RECIPROCAL_3_SHL17
	mul $m6, $m6, $m10
	shr $m6, $m6, 18
	brnzdec $m6, .Lreference_reduce_sum_loop
	bri .Lreference_reduce_sum_done
.Lreference_reduce_sum_loop:
	ld64step $a0:1, $mzero, $m3+=, 1
	ld64step $a2:3, $mzero, $m3+=, 11
	mov $m7, $m4
	add $m10, $m4, 8
	// Preload one partial, then consume the old registers while fetching the
	// next. A zero repeat count skips the body for a single remote partial.
	sub $m8, $m0, 1
	ld64step $a4:5, $mzero, $m4+=, $m1
	ld64step $a6:7, $mzero, $m10+=, $m1
	.p2align 3
	{ rpt $m8, 1
	  fnop }
	{ ld64step $a4:5, $mzero, $m4+=, $m1
	  f16v4add $a0:1, $a0:1, $a4:5 }
	{ ld64step $a6:7, $mzero, $m10+=, $m1
	  f16v4add $a2:3, $a2:3, $a6:7 }
	{ add $m4, $m7, 96
	  f16v4add $a0:1, $a0:1, $a4:5 }
	{ st64step $a0:1, $mzero, $m2+=, 1
	  f16v4add $a2:3, $a2:3, $a6:7 }
	st64step $a2:3, $mzero, $m2+=, 11
	brnzdec $m6, .Lreference_reduce_sum_loop
.Lreference_reduce_sum_done:
	exitz $m15

#include "/srv/home/gc-sdk/ipu-stack/device/reduce_add_f16.S"

.section .text.kernel_test_benign_fp,"ax",@progbits
.supervisor
.p2align 2
.globl kernel_test_benign_fp
kernel_test_benign_fp:
get $m0, $FP_ICTL
ldconst $m1, 0xfffffff8
and $m0, $m0, $m1
put $FP_ICTL, $m0
br $m10

.section .text.kernel_test_strict_fp,"ax",@progbits
.supervisor
.p2align 2
.globl kernel_test_strict_fp
kernel_test_strict_fp:
get $m0, $FP_ICTL
setzi $m1, 7
or $m0, $m0, $m1
put $FP_ICTL, $m0
br $m10
