Files
corona/threshold/bench_gpu_test.go
T
Hanzo AI 2e22541859 gpu: opt corona threshold signing into lattice/ring GPU NTT dispatch
Add corona/gpu package — the single, decomplecting point where corona
opts into the lattice library's per-SubRing GPU NTT dispatcher.

Architecture (decomplected).

The lattice library already owns ALL build-tag plumbing for GPU NTT:
ring.SetGPUDispatchers (subring_ops.go) is the canonical hook; the
lattice/gpu package installs it under `cgo && gpu` build tags and
provides a real CPU fallback under !cgo or !gpu. Output is byte-equal
to ring.SubRing.NTT by lattice's own contract.

corona/gpu adds the corona-side bridge: UseAccelerator() flips a
global flag, NewParams() across corona consults the flag via
MaybeRegister and binds each created Ring's SubRings into the lattice
GPU registry. Single source of truth for the opt-in; no build tags
inside corona.

Threshold gating (honest).

Single-poly Metal NTT at corona's production N=256 is roughly 4-6x
SLOWER than pure-Go ring.SubRing.NTT (measured: BenchmarkPulsarSign_
5of7 force-GPU 7.1s vs CPU 1.1s; 14of21 force-GPU 23.5s vs CPU 5.9s).
The GPU win exists only in BATCHED dispatch (many polynomials per
kernel launch), which requires future engine-layer plumbing of
lattice/gpu.MontgomeryNTTContext.Forward(data, batch>=4) bypassing
the per-poly r.NTT() pinch point.

Therefore UseAccelerator() picks defaultThreshold=1024 — above
corona's N=256 — so the SubRing dispatch is armed but does not fire
on single-poly NTT. The registry remains primed for any future batch
caller (e.g. FHE bootstraps in thresholdvm sharing this library).

UseAcceleratorForce() (threshold=1) is provided strictly for the
correctness gate: every NTT call routes through the GPU so the
byte-equality test in threshold/threshold_gpu_test.go exercises the
GPU path end-to-end. Production callers use UseAccelerator() instead.

Byte-equality.

TestThresholdSign_CPU_vs_GPU_ByteIdentical runs the full 2-round
Pulsar signing protocol with GPU dispatch off and forced-on (same
deterministic dealer randomness, same message) and asserts byte-equal
sig.C / sig.Z / sig.Delta. Passes under CGO_ENABLED=0, CGO_ENABLED=1,
and CGO_ENABLED=1 -tags gpu. Existing TestDKG2_GPU_ByteEqual coverage
extends across n=3,5,7,11,21 (production shape).

Wiring.

NewParams() in sign-bound packages — threshold, dkg2, dkg, reshare —
calls corona/gpu.MaybeRegister(r) for the main Q ring. RXi and RNu
are power-of-two moduli; the NTT path is not taken on them.

Tests.

Full corona test suite passes under both build modes:
  - CGO_ENABLED=0 go test ./...    => all green
  - CGO_ENABLED=1 -tags gpu test   => all green
2026-05-21 13:41:15 -07:00

104 lines
3.6 KiB
Go

// Copyright (C) 2025-2026, Lux Industries Inc. All rights reserved.
// See the file LICENSE for licensing terms.
package threshold
import (
"crypto/rand"
"io"
"testing"
cgpu "github.com/luxfi/corona/gpu"
)
// BenchmarkPulsarSign measures the wall-clock cost of the 2-round
// Pulsar threshold protocol's *online* phase (Round1 + Round2 +
// Finalize, given a fresh GenerateKeys epoch). The IEEE S&P 2025
// Pulsar evaluation calls out a 0.6 s online phase across 5
// continents at the production shape; this bench gives the local
// upper bound (network RTT is excluded; the cost here is pure CPU /
// GPU compute).
//
// HONEST PERFORMANCE NOTE.
//
// At corona's production ring degree N=256 the single-poly Metal NTT
// is slower than the pure-Go ring.SubRing.NTT (lattice/gpu's own
// header documents this for every N up to 16384). The default
// corona/gpu threshold = 1024 keeps single-poly dispatch OFF at
// N=256, so BenchmarkPulsarSign_*_GPU here is effectively the same
// kernel as the CPU bench with a small dispatcher branch cost; bench
// noise dominates the comparison.
//
// The GPU win for Pulsar requires a BATCHED dispatch path that calls
// lattice/gpu.MontgomeryNTTContext.Forward(data, batch>=4) at the
// engine layer, bypassing the per-poly r.NTT() pinch point. That
// kernel slot exists (see lattice/gpu/gpu_cgo.go::BatchNTT) but is
// not yet plumbed through r.NTT — that's the v0.6+ NIST submission
// pipeline work referenced in corona/dkg2/dkg2_gpu_accel.go.
//
// CPU vs GPU pairs (one bench function each) let `go test -bench .`
// emit a side-by-side comparison without bench-fixture trickery; the
// GPU bench remains useful as a regression watchdog (any change that
// adds non-trivial dispatcher overhead will show up here).
func BenchmarkPulsarSign_2of3_CPU(b *testing.B) { benchPulsarSign(b, 3, 2, false) }
func BenchmarkPulsarSign_2of3_GPU(b *testing.B) { benchPulsarSign(b, 3, 2, true) }
func BenchmarkPulsarSign_5of7_CPU(b *testing.B) { benchPulsarSign(b, 7, 5, false) }
func BenchmarkPulsarSign_5of7_GPU(b *testing.B) { benchPulsarSign(b, 7, 5, true) }
func BenchmarkPulsarSign_7of11_CPU(b *testing.B) { benchPulsarSign(b, 11, 7, false) }
func BenchmarkPulsarSign_7of11_GPU(b *testing.B) { benchPulsarSign(b, 11, 7, true) }
// 14-of-21 — production Lux consensus shape.
func BenchmarkPulsarSign_14of21_CPU(b *testing.B) { benchPulsarSign(b, 21, 14, false) }
func BenchmarkPulsarSign_14of21_GPU(b *testing.B) { benchPulsarSign(b, 21, 14, true) }
func benchPulsarSign(b *testing.B, n, thr int, gpuOn bool) {
if gpuOn {
if err := cgpu.UseAccelerator(); err != nil {
b.Fatalf("UseAccelerator: %v", err)
}
} else {
cgpu.DisableAccelerator()
}
b.Cleanup(cgpu.DisableAccelerator)
shares, _, err := GenerateKeys(thr, n, rand.Reader)
if err != nil {
b.Fatal(err)
}
signers := make([]*Signer, n)
for i, sh := range shares {
signers[i] = NewSigner(sh)
}
signerIDs := make([]int, n)
for i := range signerIDs {
signerIDs[i] = i
}
prfKey := make([]byte, 32)
if _, err := io.ReadFull(rand.Reader, prfKey); err != nil {
b.Fatal(err)
}
msg := "bench-pulsar-sign-online"
b.ResetTimer()
for i := 0; i < b.N; i++ {
sid := i + 1
round1 := make(map[int]*Round1Data, n)
for _, s := range signers {
d := s.Round1(sid, prfKey, signerIDs)
round1[d.PartyID] = d
}
round2 := make(map[int]*Round2Data, n)
for _, s := range signers {
d, err := s.Round2(sid, msg, prfKey, signerIDs, round1)
if err != nil {
b.Fatal(err)
}
round2[d.PartyID] = d
}
if _, err := signers[0].Finalize(round2); err != nil {
b.Fatal(err)
}
}
}