diff --git a/field/koalabear/poseidon2/poseidon2.go b/field/koalabear/poseidon2/poseidon2.go index abaa52da5..b1604e367 100644 --- a/field/koalabear/poseidon2/poseidon2.go +++ b/field/koalabear/poseidon2/poseidon2.go @@ -510,14 +510,18 @@ func (h *Permutation) compressx16Columns( result [][8]fr.Element, ) { - if !h.params.hasFast16_6_21 || (runtime.GOARCH == "arm64" && state != nil) { + if !h.params.hasFast16_6_21 { h.compressx16ColumnsGeneric(state, matrix, colSize, result) return } nbSteps := uint64(colSize / 8) if runtime.GOARCH == "arm64" { - permutation16x16xN_columns_arm64(&matrix[0], h.params.RoundKeys, &result[0][0], nbSteps) + var statePtr *fr.Element + if state != nil { + statePtr = &state[0] + } + permutation16x16xN_columns_arm64(&matrix[0], h.params.RoundKeys, &result[0][0], nbSteps, statePtr) return } var statePtr *fr.Element diff --git a/field/koalabear/poseidon2/poseidon2_amd64.go b/field/koalabear/poseidon2/poseidon2_amd64.go index b30860b64..fa35099ee 100644 --- a/field/koalabear/poseidon2/poseidon2_amd64.go +++ b/field/koalabear/poseidon2/poseidon2_amd64.go @@ -51,6 +51,6 @@ func permutation16x16x512_arm64(matrix *fr.Element, roundKeys [][]fr.Element, re panic("permutation16x16x512_arm64 is not implemented") } -func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64) { +func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64, state *fr.Element) { panic("permutation16x16xN_columns_arm64 is not implemented") } diff --git a/field/koalabear/poseidon2/poseidon2_arm64.go b/field/koalabear/poseidon2/poseidon2_arm64.go index 54f368d25..a0ccc9dfa 100644 --- a/field/koalabear/poseidon2/poseidon2_arm64.go +++ b/field/koalabear/poseidon2/poseidon2_arm64.go @@ -39,4 +39,4 @@ func permutation16x16xN_columns_avx512(matrix *fr.Element, roundKeys [][]fr.Elem func permutation16x16x512_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element) //go:noescape -func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64) +func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64, state *fr.Element) diff --git a/field/koalabear/poseidon2/poseidon2_arm64.s b/field/koalabear/poseidon2/poseidon2_arm64.s index 6ed37a075..7f26ed38a 100644 --- a/field/koalabear/poseidon2/poseidon2_arm64.s +++ b/field/koalabear/poseidon2/poseidon2_arm64.s @@ -810,7 +810,7 @@ step_loop: BNE batch_loop RET -TEXT ·permutation16x16xN_columns_arm64(SB), $128-48 +TEXT ·permutation16x16xN_columns_arm64(SB), $128-56 MOVD matrix+0(FP), R0 MOVD roundKeys+8(FP), R1 MOVD result+32(FP), R2 @@ -820,11 +820,36 @@ TEXT ·permutation16x16xN_columns_arm64(SB), $128-48 VDUP R4, V1.S4 MOVD $1, R5 VDUP R5, V28.S4 - MOVD nbSteps+40(FP), R14 + MOVD nbSteps+40(FP), R15 + MOVD state+48(FP), R14 MOVD $0, R7 batch_loop: + CBZ R14, state_is_zero1 + LSL $4, R7, R13 + ADD R14, R13, R13 + VLD1.P 16(R13), [V2.S4] + ADD $0x30, R13, R13 + VLD1.P 16(R13), [V3.S4] + ADD $0x30, R13, R13 + VLD1.P 16(R13), [V4.S4] + ADD $0x30, R13, R13 + VLD1.P 16(R13), [V5.S4] + ADD $0x30, R13, R13 + VLD1.P 16(R13), [V6.S4] + ADD $0x30, R13, R13 + VLD1.P 16(R13), [V7.S4] + ADD $0x30, R13, R13 + VLD1.P 16(R13), [V8.S4] + ADD $0x30, R13, R13 + VLD1.P 16(R13), [V9.S4] + ADD $0x30, R13, R13 + JMP state_ready2 + +state_is_zero1: ZERO_STATE() + +state_ready2: MOVD $0, R8 LSL $4, R7, R13 ADD R0, R13, R9 @@ -877,7 +902,7 @@ step_loop: FULL_ROUND(624) FEED_FORWARD() ADD $1, R8, R8 - CMP R14, R8 + CMP R15, R8 BNE step_loop LSL $7, R7, R13 ADD R2, R13, R9 diff --git a/field/koalabear/poseidon2/poseidon2_bench_test.go b/field/koalabear/poseidon2/poseidon2_bench_test.go new file mode 100644 index 000000000..7e383a0dd --- /dev/null +++ b/field/koalabear/poseidon2/poseidon2_bench_test.go @@ -0,0 +1,42 @@ +package poseidon2 + +import ( + "testing" + + fr "github.com/consensys/gnark-crypto/field/koalabear" +) + +func benchmarkCompressx16ColumnsWithState(b *testing.B, accelerated bool) { + const colSize = 512 + state := make([]fr.Element, 16*8) + matrix := make([]fr.Element, 16*colSize) + for i := range state { + state[i].SetUint64(uint64(i*40503 + 17)) + } + for i := range matrix { + matrix[i].SetUint64(uint64(i*2654435761 + 23)) + } + result := make([][8]fr.Element, 16) + h := NewPermutation(16, 6, 21) + if accelerated { + if !h.params.hasFast16_6_21 { + b.Skip("Poseidon2 accelerator unavailable") + } + } else { + h.disableAVX512() + } + + b.ResetTimer() + for b.Loop() { + h.Compressx16ColumnsWithState(state, matrix, colSize, result) + } +} + +func BenchmarkCompressx16ColumnsWithState512(b *testing.B) { + b.Run("accelerated", func(b *testing.B) { + benchmarkCompressx16ColumnsWithState(b, true) + }) + b.Run("generic", func(b *testing.B) { + benchmarkCompressx16ColumnsWithState(b, false) + }) +} diff --git a/field/koalabear/poseidon2/poseidon2_purego.go b/field/koalabear/poseidon2/poseidon2_purego.go index 537b6c796..4c821601f 100644 --- a/field/koalabear/poseidon2/poseidon2_purego.go +++ b/field/koalabear/poseidon2/poseidon2_purego.go @@ -34,6 +34,6 @@ func permutation16x16x512_arm64(matrix *fr.Element, roundKeys [][]fr.Element, re panic("permutation16x16x512_arm64 is not implemented") } -func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64) { +func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64, state *fr.Element) { panic("permutation16x16xN_columns_arm64 is not implemented") } diff --git a/internal/generator/field/asm/arm64/element_vec_F31_poseidon2.go b/internal/generator/field/asm/arm64/element_vec_F31_poseidon2.go index 71a295bff..1883d70b1 100644 --- a/internal/generator/field/asm/arm64/element_vec_F31_poseidon2.go +++ b/internal/generator/field/asm/arm64/element_vec_F31_poseidon2.go @@ -105,7 +105,7 @@ func (f *FFArm64) generatePoseidon2_F31_16x16(params amd64.Poseidon2Parameters, argSize := 8 + 24 + 8 // matrix ptr + roundKeys slice header + result ptr if columns { fnName = "permutation16x16xN_columns_arm64" - argSize += 8 // + nbSteps + argSize += 16 // + nbSteps + state ptr } // Stack frame for temporary storage during each step (8 vectors × 16 bytes = 128 bytes) @@ -172,16 +172,18 @@ func (f *FFArm64) generatePoseidon2_F31_16x16(params amd64.Poseidon2Parameters, batchIdx := registers.Pop() // outer loop counter (0..3) stepIdx := registers.Pop() // inner loop counter (0..N-1) // Pointers to 4 rows for current batch - ptr0 := registers.Pop() // data pointer for batch row 0 - ptr1 := registers.Pop() // data pointer for batch row 1 - ptr2 := registers.Pop() // data pointer for batch row 2 - ptr3 := registers.Pop() // data pointer for batch row 3 - tmpCalc := registers.Pop() // temporary for address calculations + ptr0 := registers.Pop() // data pointer for batch row 0 + ptr1 := registers.Pop() // data pointer for batch row 1 + ptr2 := registers.Pop() // data pointer for batch row 2 + ptr3 := registers.Pop() // data pointer for batch row 3 + tmpCalc := registers.Pop() // temporary for address calculations + addrState := registers.Pop() // optional initial state, in column-major layout var nbSteps arm64.Register // number of steps (columns variant only) if columns { nbSteps = registers.Pop() f.MOVD("nbSteps+40(FP)", nbSteps) + f.MOVD("state+48(FP)", addrState) } // defineOnce defines a macro on the first kernel generation and reuses it on @@ -607,7 +609,25 @@ func (f *FFArm64) generatePoseidon2_F31_16x16(params amd64.Poseidon2Parameters, f.MOVD(0, batchIdx) f.LABEL("batch_loop") - zeroState() + if columns { + stateIsZero := f.NewLabel("state_is_zero") + stateReady := f.NewLabel("state_ready") + f.CBZ(addrState, stateIsZero) + // state[pos*16+lane] is transposed in the same way as the result: + // load four lanes for each of the eight capacity coordinates. + f.WriteLn(fmt.Sprintf(" LSL $4, %s, %s", batchIdx, tmpCalc)) + f.ADD(addrState, tmpCalc, tmpCalc) + for pos := range 8 { + f.VLD1_P(16, tmpCalc, v[pos].S4()) + f.ADD(48, tmpCalc, tmpCalc) + } + f.JMP(stateReady) + f.LABEL(stateIsZero) + zeroState() + f.LABEL(stateReady) + } else { + zeroState() + } // Initialize pointers for 4 parallel inputs const N = 512 / 8 // 64 steps per batch diff --git a/internal/generator/field/template/poseidon2/poseidon2.amd64.go.tmpl b/internal/generator/field/template/poseidon2/poseidon2.amd64.go.tmpl index 8fdb01593..5e4bdbd5f 100644 --- a/internal/generator/field/template/poseidon2/poseidon2.amd64.go.tmpl +++ b/internal/generator/field/template/poseidon2/poseidon2.amd64.go.tmpl @@ -46,7 +46,7 @@ func permutation16x16x512_arm64(matrix *fr.Element, roundKeys [][]fr.Element, re panic("permutation16x16x512_arm64 is not implemented") } -func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64) { +func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64, state *fr.Element) { panic("permutation16x16xN_columns_arm64 is not implemented") } {{- end }} diff --git a/internal/generator/field/template/poseidon2/poseidon2.arm64.go.tmpl b/internal/generator/field/template/poseidon2/poseidon2.arm64.go.tmpl index 1ef5e2280..d25eefd1a 100644 --- a/internal/generator/field/template/poseidon2/poseidon2.arm64.go.tmpl +++ b/internal/generator/field/template/poseidon2/poseidon2.arm64.go.tmpl @@ -34,5 +34,5 @@ func permutation16x16xN_columns_avx512(matrix *fr.Element, roundKeys [][]fr.Elem func permutation16x16x512_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element) //go:noescape -func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64) +func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64, state *fr.Element) {{- end }} diff --git a/internal/generator/field/template/poseidon2/poseidon2.go.tmpl b/internal/generator/field/template/poseidon2/poseidon2.go.tmpl index bfb0317c0..dd60398cd 100644 --- a/internal/generator/field/template/poseidon2/poseidon2.go.tmpl +++ b/internal/generator/field/template/poseidon2/poseidon2.go.tmpl @@ -674,14 +674,18 @@ func (h *Permutation) compressx16Columns( result [][8]fr.Element, ) { - if !h.params.hasFast{{- $wc}}_{{- $fc}}_{{- $pc}} || (runtime.GOARCH == "arm64" && state != nil) { + if !h.params.hasFast{{- $wc}}_{{- $fc}}_{{- $pc}} { h.compressx16ColumnsGeneric(state, matrix, colSize, result) return } nbSteps := uint64(colSize / 8) if runtime.GOARCH == "arm64" { - permutation16x16xN_columns_arm64(&matrix[0], h.params.RoundKeys, &result[0][0], nbSteps) + var statePtr *fr.Element + if state != nil { + statePtr = &state[0] + } + permutation16x16xN_columns_arm64(&matrix[0], h.params.RoundKeys, &result[0][0], nbSteps, statePtr) return } var statePtr *fr.Element diff --git a/internal/generator/field/template/poseidon2/poseidon2.purego.go.tmpl b/internal/generator/field/template/poseidon2/poseidon2.purego.go.tmpl index b45c3f39c..d00cc941b 100644 --- a/internal/generator/field/template/poseidon2/poseidon2.purego.go.tmpl +++ b/internal/generator/field/template/poseidon2/poseidon2.purego.go.tmpl @@ -27,7 +27,7 @@ func permutation16x16x512_arm64(matrix *fr.Element, roundKeys [][]fr.Element, re panic("permutation16x16x512_arm64 is not implemented") } -func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64) { +func permutation16x16xN_columns_arm64(matrix *fr.Element, roundKeys [][]fr.Element, result *fr.Element, nbSteps uint64, state *fr.Element) { panic("permutation16x16xN_columns_arm64 is not implemented") } {{- end }}