mirror of
https://github.com/luxfi/crypto.git
synced 2026-07-27 01:54:50 +00:00
Remove dead code and consolidate GPU ZK operations
- Remove crypto/common/math/ (dead code, using luxfi/math instead) - Consolidate gpu/zk.go to use unified luxfi/gpu package - Delete platform-specific gpu/zk_metal.go in favor of unified approach - Update bls/types.go with GPU bindings - Update dependencies for luxfi/gpu
This commit is contained in:
@@ -1469,39 +1469,44 @@ New package for anonymous group signing (LSAG - Linkable Spontaneous Anonymous G
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## GPU-Accelerated ZK Operations (2026-01-02) - COMPLETED
|
## GPU-Accelerated ZK Operations (2026-01-03) - UNIFIED ARCHITECTURE
|
||||||
|
|
||||||
### Overview
|
### Overview
|
||||||
|
|
||||||
The `gpu/` package provides GPU-accelerated ZK cryptographic operations with automatic threshold-based routing between CPU and GPU execution paths.
|
The `gpu/` package provides GPU-accelerated ZK cryptographic operations with automatic threshold-based routing between CPU and GPU execution paths.
|
||||||
|
|
||||||
### Architecture
|
### Architecture (Unified GPU Stack)
|
||||||
|
|
||||||
```
|
```
|
||||||
┌─────────────────────────────────────────────────────────┐
|
┌─────────────────────────────────────────────────────────┐
|
||||||
│ gpu/zk.go │
|
│ crypto/gpu/zk.go │
|
||||||
│ - Threshold-gated routing (CPU vs GPU) │
|
│ - Threshold-gated routing (CPU vs GPU) │
|
||||||
│ - CPU fallback via gnark-crypto │
|
│ - CPU fallback via gnark-crypto │
|
||||||
│ - GPUHooks interface for platform injection │
|
│ - Uses github.com/luxfi/gpu for GPU ops │
|
||||||
└──────────────────────────┬──────────────────────────────┘
|
└──────────────────────────┬──────────────────────────────┘
|
||||||
│
|
│
|
||||||
┌─────────────────┼─────────────────┐
|
▼
|
||||||
▼ ▼ ▼
|
┌─────────────────────────────────────┐
|
||||||
┌─────────────────┐ ┌─────────────────┐ ┌─────────────────┐
|
│ github.com/luxfi/gpu │
|
||||||
│ gpu/zk_metal.go │ │ gpu/zk_cuda.go │ │ (future backends)│
|
│ Go bindings to luxcpp/gpu (MLX) │
|
||||||
│ (darwin,arm64) │ │ (linux,nvidia) │ │ │
|
│ │
|
||||||
│ CGO + Metal │ │ CGO + CUDA │ │ │
|
│ zk.go (non-CGO stub) │
|
||||||
└────────┬────────┘ └────────┬────────┘ └──────────────────┘
|
│ zk_cgo.go (CGO bindings) │
|
||||||
│ │
|
└──────────────────┬──────────────────┘
|
||||||
└───────────────────┘
|
│
|
||||||
│
|
▼
|
||||||
┌────────▼────────┐
|
┌─────────────────────────────────────┐
|
||||||
│ luxcpp/crypto │
|
│ luxcpp/gpu (MLX) │
|
||||||
│ C/C++ + Metal │
|
│ Unified Metal/CUDA/CPU backend │
|
||||||
│ v0.1.0 │
|
│ │
|
||||||
└─────────────────┘
|
│ mlx/zk/zk.cpp - C++ impl │
|
||||||
|
│ mlx/zk/zk_c_api.h - C API │
|
||||||
|
└─────────────────────────────────────┘
|
||||||
```
|
```
|
||||||
|
|
||||||
|
**Key Change (2026-01-03)**: Removed separate platform files (`zk_metal.go`, `zk_cuda.go`)
|
||||||
|
in favor of the unified `github.com/luxfi/gpu` package which handles all backends via MLX.
|
||||||
|
|
||||||
### Threshold Constants (Tuned for Apple Silicon)
|
### Threshold Constants (Tuned for Apple Silicon)
|
||||||
|
|
||||||
| Operation | Threshold | Description |
|
| Operation | Threshold | Description |
|
||||||
@@ -1517,9 +1522,9 @@ Below threshold: CPU (lower latency). Above threshold: GPU (higher throughput).
|
|||||||
### Core Types
|
### Core Types
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// Fr256 represents a 256-bit field element (BN254 scalar field).
|
// Fr256 is a type alias to luxgpu.Fr256 - 256-bit field element (BN254 scalar field).
|
||||||
// Uses 4 x 64-bit limbs in little-endian order.
|
// Uses 4 x 64-bit limbs in little-endian order.
|
||||||
type Fr256 [4]uint64
|
type Fr256 = luxgpu.Fr256 // [4]uint64
|
||||||
|
|
||||||
// ZKContext provides GPU-accelerated ZK operations with automatic routing.
|
// ZKContext provides GPU-accelerated ZK operations with automatic routing.
|
||||||
type ZKContext struct {
|
type ZKContext struct {
|
||||||
@@ -1528,18 +1533,11 @@ type ZKContext struct {
|
|||||||
gpuCalls int64
|
gpuCalls int64
|
||||||
cpuCalls int64
|
cpuCalls int64
|
||||||
}
|
}
|
||||||
|
|
||||||
// GPUHooks contains function pointers to GPU implementations.
|
|
||||||
// Set by platform-specific init() functions (zk_metal.go, zk_cuda.go).
|
|
||||||
type GPUHooks struct {
|
|
||||||
HashPair func(left, right []Fr256) ([]Fr256, error)
|
|
||||||
MerkleLayer func(nodes []Fr256) ([]Fr256, error)
|
|
||||||
MerkleTree func(leaves []Fr256) ([]Fr256, error)
|
|
||||||
BatchCommitment func(values, blindings, salts []Fr256) ([]Fr256, error)
|
|
||||||
BatchNullifier func(keys, commitments, indices []Fr256) ([]Fr256, error)
|
|
||||||
}
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
**Note**: `Fr256` is now a type alias to `github.com/luxfi/gpu.Fr256`, ensuring consistent
|
||||||
|
type representation across the GPU stack.
|
||||||
|
|
||||||
### Operations
|
### Operations
|
||||||
|
|
||||||
| Function | Description |
|
| Function | Description |
|
||||||
@@ -1597,14 +1595,24 @@ CGO_ENABLED=1 go build -tags "gpu" ./...
|
|||||||
|
|
||||||
| File | Purpose |
|
| File | Purpose |
|
||||||
|------|---------|
|
|------|---------|
|
||||||
| `gpu/zk.go` | Core ZK operations, threshold routing, CPU fallback |
|
| `gpu/zk.go` | Core ZK operations, threshold routing, CPU fallback via gnark-crypto |
|
||||||
| `gpu/zk_metal.go` | Metal CGO bindings (darwin,arm64) |
|
|
||||||
| `gpu/zk_test.go` | Tests for ZK operations |
|
| `gpu/zk_test.go` | Tests for ZK operations |
|
||||||
|
|
||||||
### Dependencies
|
### Dependencies
|
||||||
|
|
||||||
|
- `github.com/luxfi/gpu` - Unified GPU bindings (wraps luxcpp/gpu)
|
||||||
- `github.com/consensys/gnark-crypto` - CPU Poseidon2 via BN254/Fr
|
- `github.com/consensys/gnark-crypto` - CPU Poseidon2 via BN254/Fr
|
||||||
- `luxcpp/crypto v0.1.0` - C/C++ Metal shaders (optional, for GPU)
|
|
||||||
|
### GPU Stack (luxcpp/gpu)
|
||||||
|
|
||||||
|
| File | Purpose |
|
||||||
|
|------|---------|
|
||||||
|
| `lux/gpu/zk.go` | Non-CGO stub (returns ErrZKNotAvailable) |
|
||||||
|
| `lux/gpu/zk_cgo.go` | CGO bindings to luxcpp/gpu |
|
||||||
|
| `luxcpp/gpu/mlx/zk/zk.h` | C++ ZK operations header |
|
||||||
|
| `luxcpp/gpu/mlx/zk/zk.cpp` | C++ ZK implementation using MLX |
|
||||||
|
| `luxcpp/gpu/mlx/zk/zk_c_api.h` | C API for Go bindings |
|
||||||
|
| `luxcpp/gpu/mlx/zk/zk_c_api.cpp` | C API implementation |
|
||||||
|
|
||||||
### Test Coverage
|
### Test Coverage
|
||||||
|
|
||||||
@@ -1615,13 +1623,14 @@ All tests pass in both CGO and non-CGO modes:
|
|||||||
CGO_ENABLED=0 go test ./gpu/...
|
CGO_ENABLED=0 go test ./gpu/...
|
||||||
|
|
||||||
# CGO (with GPU if available)
|
# CGO (with GPU if available)
|
||||||
CGO_ENABLED=1 go test -tags gpu ./gpu/...
|
CGO_ENABLED=1 go test ./gpu/...
|
||||||
```
|
```
|
||||||
|
|
||||||
### Published Versions
|
### Published Versions
|
||||||
|
|
||||||
- `github.com/luxfi/crypto v1.17.35` - Go package with GPU ZK operations
|
- `github.com/luxfi/crypto` - Go package with GPU ZK operations
|
||||||
- `github.com/luxcpp/crypto v0.1.0` - C++ Metal shaders and C API
|
- `github.com/luxfi/gpu` - Go bindings to unified GPU (Metal/CUDA/CPU)
|
||||||
|
- `luxcpp/gpu` - C++ MLX-based GPU library
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
@@ -21,4 +21,6 @@ var (
|
|||||||
ErrNoSignatures = errors.New("no signatures")
|
ErrNoSignatures = errors.New("no signatures")
|
||||||
ErrFailedSignatureAggregation = errors.New("couldn't aggregate signatures")
|
ErrFailedSignatureAggregation = errors.New("couldn't aggregate signatures")
|
||||||
ErrFailedSecretKeyDeserialize = errors.New("couldn't deserialize secret key")
|
ErrFailedSecretKeyDeserialize = errors.New("couldn't deserialize secret key")
|
||||||
|
ErrInvalidInput = errors.New("invalid input")
|
||||||
|
ErrGPUNotAvailable = errors.New("GPU not available")
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,215 +0,0 @@
|
|||||||
// Copyright 2025 The Lux Authors
|
|
||||||
// This file is part of the Lux library.
|
|
||||||
//
|
|
||||||
// The Lux library is free software: you can redistribute it and/or modify
|
|
||||||
// it under the terms of the GNU Lesser General Public License as published by
|
|
||||||
// the Free Software Foundation, either version 3 of the License, or
|
|
||||||
// (at your option) any later version.
|
|
||||||
//
|
|
||||||
// The Lux library is distributed in the hope that it will be useful,
|
|
||||||
// but WITHOUT ANY WARRANTY; without even the implied warranty of
|
|
||||||
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
|
|
||||||
// GNU Lesser General Public License for more details.
|
|
||||||
//
|
|
||||||
// You should have received a copy of the GNU Lesser General Public License
|
|
||||||
// along with the Lux library. If not, see <http://www.gnu.org/licenses/>.
|
|
||||||
|
|
||||||
// Package math provides integer math utilities.
|
|
||||||
package math
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"math/big"
|
|
||||||
"math/bits"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Various big integer limit values.
|
|
||||||
var (
|
|
||||||
tt256 = new(big.Int).Lsh(big.NewInt(1), 256)
|
|
||||||
tt256m1 = new(big.Int).Sub(tt256, big.NewInt(1))
|
|
||||||
tt255 = new(big.Int).Lsh(big.NewInt(1), 255)
|
|
||||||
MaxBig256 = new(big.Int).Sub(new(big.Int).Lsh(big.NewInt(1), 256), big.NewInt(1))
|
|
||||||
MaxBig63 = new(big.Int).Sub(new(big.Int).Lsh(big.NewInt(1), 63), big.NewInt(1))
|
|
||||||
)
|
|
||||||
|
|
||||||
// ParseBig256 parses s as a 256 bit integer in decimal or hexadecimal syntax.
|
|
||||||
// Leading zeros are accepted. The empty string parses as zero.
|
|
||||||
func ParseBig256(s string) (*big.Int, bool) {
|
|
||||||
if s == "" {
|
|
||||||
return new(big.Int), true
|
|
||||||
}
|
|
||||||
var bigint *big.Int
|
|
||||||
var ok bool
|
|
||||||
if len(s) >= 2 && (s[:2] == "0x" || s[:2] == "0X") {
|
|
||||||
bigint, ok = new(big.Int).SetString(s[2:], 16)
|
|
||||||
} else {
|
|
||||||
bigint, ok = new(big.Int).SetString(s, 10)
|
|
||||||
}
|
|
||||||
if ok && bigint.BitLen() > 256 {
|
|
||||||
bigint, ok = nil, false
|
|
||||||
}
|
|
||||||
return bigint, ok
|
|
||||||
}
|
|
||||||
|
|
||||||
// MustParseBig256 parses s as a 256 bit big integer and panics if the string is invalid.
|
|
||||||
func MustParseBig256(s string) *big.Int {
|
|
||||||
v, ok := ParseBig256(s)
|
|
||||||
if !ok {
|
|
||||||
panic("invalid 256 bit integer: " + s)
|
|
||||||
}
|
|
||||||
return v
|
|
||||||
}
|
|
||||||
|
|
||||||
// BigPow returns a ** b as a big integer.
|
|
||||||
func BigPow(a, b int64) *big.Int {
|
|
||||||
r := big.NewInt(a)
|
|
||||||
return r.Exp(r, big.NewInt(b), nil)
|
|
||||||
}
|
|
||||||
|
|
||||||
// BigMax returns the larger of x or y.
|
|
||||||
func BigMax(x, y *big.Int) *big.Int {
|
|
||||||
if x.Cmp(y) < 0 {
|
|
||||||
return y
|
|
||||||
}
|
|
||||||
return x
|
|
||||||
}
|
|
||||||
|
|
||||||
// BigMin returns the smaller of x or y.
|
|
||||||
func BigMin(x, y *big.Int) *big.Int {
|
|
||||||
if x.Cmp(y) > 0 {
|
|
||||||
return y
|
|
||||||
}
|
|
||||||
return x
|
|
||||||
}
|
|
||||||
|
|
||||||
// PaddedBigBytes encodes a big integer as a big-endian byte slice. The length
|
|
||||||
// of the slice is at least n bytes.
|
|
||||||
func PaddedBigBytes(bigint *big.Int, n int) []byte {
|
|
||||||
if bigint.BitLen()/8 >= n {
|
|
||||||
return bigint.Bytes()
|
|
||||||
}
|
|
||||||
ret := make([]byte, n)
|
|
||||||
bigint.FillBytes(ret)
|
|
||||||
return ret
|
|
||||||
}
|
|
||||||
|
|
||||||
// ReadBits encodes the absolute value of bigint as big-endian bytes. Callers must ensure
|
|
||||||
// that bigint is non-negative.
|
|
||||||
func ReadBits(bigint *big.Int, buf []byte) {
|
|
||||||
i := len(buf)
|
|
||||||
for _, d := range bigint.Bits() {
|
|
||||||
for j := 0; j < wordBytes && i > 0; j++ {
|
|
||||||
i--
|
|
||||||
buf[i] = byte(d)
|
|
||||||
d >>= 8
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// U256 encodes as a 256 bit two's complement number. This operation is destructive.
|
|
||||||
func U256(x *big.Int) *big.Int {
|
|
||||||
return x.And(x, tt256m1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// U256Bytes converts a big Int into a 256bit EVM number.
|
|
||||||
// This operation is destructive.
|
|
||||||
func U256Bytes(n *big.Int) []byte {
|
|
||||||
return PaddedBigBytes(U256(n), 32)
|
|
||||||
}
|
|
||||||
|
|
||||||
// S256 interprets x as a two's complement number.
|
|
||||||
// x must not exceed 256 bits (the result is undefined if it does) and is not modified.
|
|
||||||
//
|
|
||||||
// S256(0) = 0
|
|
||||||
// S256(1) = 1
|
|
||||||
// S256(2**255) = -2**255
|
|
||||||
// S256(2**256-1) = -1
|
|
||||||
func S256(x *big.Int) *big.Int {
|
|
||||||
if x.Cmp(tt255) < 0 {
|
|
||||||
return x
|
|
||||||
}
|
|
||||||
return new(big.Int).Sub(x, tt256)
|
|
||||||
}
|
|
||||||
|
|
||||||
// SafeSub returns x-y and checks for overflow.
|
|
||||||
func SafeSub(x, y uint64) (uint64, bool) {
|
|
||||||
diff, borrowOut := bits.Sub64(x, y, 0)
|
|
||||||
return diff, borrowOut != 0
|
|
||||||
}
|
|
||||||
|
|
||||||
// SafeAdd returns x+y and checks for overflow.
|
|
||||||
func SafeAdd(x, y uint64) (uint64, bool) {
|
|
||||||
sum, carryOut := bits.Add64(x, y, 0)
|
|
||||||
return sum, carryOut != 0
|
|
||||||
}
|
|
||||||
|
|
||||||
// SafeMul returns x*y and checks for overflow.
|
|
||||||
func SafeMul(x, y uint64) (uint64, bool) {
|
|
||||||
hi, lo := bits.Mul64(x, y)
|
|
||||||
return lo, hi != 0
|
|
||||||
}
|
|
||||||
|
|
||||||
// SafeDiv returns x/y and checks for division by zero.
|
|
||||||
func SafeDiv(x, y uint64) (uint64, error) {
|
|
||||||
if y == 0 {
|
|
||||||
return 0, fmt.Errorf("division by zero")
|
|
||||||
}
|
|
||||||
return x / y, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Byte returns the byte at position n,
|
|
||||||
// with the supplied padlength in Little Endian encoding.
|
|
||||||
// n==0 returns the MSB
|
|
||||||
// Example: bigint '5', padlength 32, n=31 => 5
|
|
||||||
func Byte(bigint *big.Int, padlength, n int) byte {
|
|
||||||
if n >= padlength {
|
|
||||||
return byte(0)
|
|
||||||
}
|
|
||||||
return bigEndianByteAt(bigint, padlength-1-n)
|
|
||||||
}
|
|
||||||
|
|
||||||
// bigEndianByteAt returns the byte at position n,
|
|
||||||
// in Big Endian encoding
|
|
||||||
// So n==0 returns the least significant byte
|
|
||||||
func bigEndianByteAt(bigint *big.Int, n int) byte {
|
|
||||||
words := bigint.Bits()
|
|
||||||
// Check word-bucket the byte will reside in
|
|
||||||
i := n / wordBytes
|
|
||||||
if i >= len(words) {
|
|
||||||
return byte(0)
|
|
||||||
}
|
|
||||||
word := words[i]
|
|
||||||
// Offset of the byte
|
|
||||||
shift := 8 * uint(n%wordBytes)
|
|
||||||
|
|
||||||
return byte(word >> shift)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Exp implements exponentiation by squaring.
|
|
||||||
// Exp returns a newly-allocated big integer and does not change
|
|
||||||
// base or exponent. The result is truncated to 256 bits.
|
|
||||||
//
|
|
||||||
// Courtesy @karalabe and @chfast
|
|
||||||
func Exp(base, exponent *big.Int) *big.Int {
|
|
||||||
copyBase := new(big.Int).Set(base)
|
|
||||||
result := big.NewInt(1)
|
|
||||||
|
|
||||||
for _, word := range exponent.Bits() {
|
|
||||||
for i := 0; i < wordBits; i++ {
|
|
||||||
if word&1 == 1 {
|
|
||||||
U256(result.Mul(result, copyBase))
|
|
||||||
}
|
|
||||||
U256(copyBase.Mul(copyBase, copyBase))
|
|
||||||
word >>= 1
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return result
|
|
||||||
}
|
|
||||||
|
|
||||||
// Architecture-dependent constants
|
|
||||||
const (
|
|
||||||
// wordBits is the number of bits in a big.Word.
|
|
||||||
wordBits = 32 << (uint64(^big.Word(0)) >> 63)
|
|
||||||
// wordBytes is the number of bytes in a big.Word.
|
|
||||||
wordBytes = wordBits / 8
|
|
||||||
)
|
|
||||||
@@ -16,6 +16,7 @@ require (
|
|||||||
github.com/jedisct1/go-minisign v0.0.0-20230811132847-661be99b8267
|
github.com/jedisct1/go-minisign v0.0.0-20230811132847-661be99b8267
|
||||||
github.com/leanovate/gopter v0.2.11
|
github.com/leanovate/gopter v0.2.11
|
||||||
github.com/luxfi/cache v1.1.0
|
github.com/luxfi/cache v1.1.0
|
||||||
|
github.com/luxfi/gpu v0.29.4
|
||||||
github.com/luxfi/ids v1.2.4
|
github.com/luxfi/ids v1.2.4
|
||||||
github.com/luxfi/log v1.1.26
|
github.com/luxfi/log v1.1.26
|
||||||
github.com/mr-tron/base58 v1.2.0
|
github.com/mr-tron/base58 v1.2.0
|
||||||
|
|||||||
@@ -209,6 +209,8 @@ github.com/leanovate/gopter v0.2.11 h1:vRjThO1EKPb/1NsDXuDrzldR28RLkBflWYcU9CvzW
|
|||||||
github.com/leanovate/gopter v0.2.11/go.mod h1:aK3tzZP/C+p1m3SPRE4SYZFGP7jjkuSI4f7Xvpt0S9c=
|
github.com/leanovate/gopter v0.2.11/go.mod h1:aK3tzZP/C+p1m3SPRE4SYZFGP7jjkuSI4f7Xvpt0S9c=
|
||||||
github.com/luxfi/cache v1.1.0 h1:6LUyGGZ+rrMAJBbAU6+UwkcamXj3zsboRUodIof2Ong=
|
github.com/luxfi/cache v1.1.0 h1:6LUyGGZ+rrMAJBbAU6+UwkcamXj3zsboRUodIof2Ong=
|
||||||
github.com/luxfi/cache v1.1.0/go.mod h1:9GvlEEE9rFPaaWxvVpSPwW8ZMo2+8VMNNcuPa4AwzPg=
|
github.com/luxfi/cache v1.1.0/go.mod h1:9GvlEEE9rFPaaWxvVpSPwW8ZMo2+8VMNNcuPa4AwzPg=
|
||||||
|
github.com/luxfi/gpu v0.29.4 h1:0eCgEbffLayTLaPYLkXx5IbL2490r5jqZAKoSfkVVZI=
|
||||||
|
github.com/luxfi/gpu v0.29.4/go.mod h1:7pFsHqra1Vrmy2aGXVEPKeKsPR+Bn+QmBXuByK3fdPA=
|
||||||
github.com/luxfi/ids v1.2.4 h1:e0OaeSI6xWjS9JnxxxfvCCs9HHDUrwDc3M9ctIhjcDI=
|
github.com/luxfi/ids v1.2.4 h1:e0OaeSI6xWjS9JnxxxfvCCs9HHDUrwDc3M9ctIhjcDI=
|
||||||
github.com/luxfi/ids v1.2.4/go.mod h1:HwvRJSNbuZS+u+0JqeHfPX2S13iFyLCnfGxjBCmt244=
|
github.com/luxfi/ids v1.2.4/go.mod h1:HwvRJSNbuZS+u+0JqeHfPX2S13iFyLCnfGxjBCmt244=
|
||||||
github.com/luxfi/log v1.1.26 h1:ECnJ4wV7TF6WZ5gC6CD6ddZ4OhF+zhQij4bgVUhumDg=
|
github.com/luxfi/log v1.1.26 h1:ECnJ4wV7TF6WZ5gC6CD6ddZ4OhF+zhQij4bgVUhumDg=
|
||||||
|
|||||||
@@ -1,21 +1,18 @@
|
|||||||
// Package gpu provides GPU-accelerated ZK cryptographic operations.
|
// Package gpu provides GPU-accelerated ZK cryptographic operations.
|
||||||
//
|
//
|
||||||
// This file provides threshold-gated routing between CPU and GPU:
|
// This package wraps github.com/luxfi/gpu for unified GPU support:
|
||||||
// - Below threshold: CPU execution (lower latency for small batches)
|
// - Metal (Apple Silicon via MLX)
|
||||||
// - Above threshold: GPU execution (higher throughput for large batches)
|
// - CUDA (NVIDIA via MLX)
|
||||||
|
// - CPU fallback (gnark-crypto)
|
||||||
//
|
//
|
||||||
// Operations:
|
// Operations:
|
||||||
// - Poseidon2 hash (BN254/Fr) for Merkle trees
|
// - Poseidon2 hash (BN254/Fr) for Merkle trees
|
||||||
// - Multi-scalar multiplication (MSM) for commitments
|
// - Multi-scalar multiplication (MSM) for commitments
|
||||||
// - Batch commitment/nullifier operations
|
// - Batch commitment/nullifier operations
|
||||||
//
|
//
|
||||||
// Threshold constants (tuned for Apple Silicon M-series):
|
// Threshold-gated routing:
|
||||||
//
|
// - Below threshold: CPU (lower latency)
|
||||||
// POSEIDON2: 64 hashes - GPU faster above this
|
// - Above threshold: GPU (higher throughput)
|
||||||
// MERKLE: 128 leaf pairs - GPU faster above this
|
|
||||||
// MSM: 256 point-scalar pairs - GPU faster above this
|
|
||||||
// COMMITMENT: 128 commitments - GPU faster above this
|
|
||||||
// FRI: 512 evaluations - GPU faster above this
|
|
||||||
package gpu
|
package gpu
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -23,7 +20,8 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/consensys/gnark-crypto/ecc/bn254/fr"
|
luxgpu "github.com/luxfi/gpu"
|
||||||
|
|
||||||
"github.com/consensys/gnark-crypto/ecc/bn254/fr/poseidon2"
|
"github.com/consensys/gnark-crypto/ecc/bn254/fr/poseidon2"
|
||||||
gnarkHash "github.com/consensys/gnark-crypto/hash"
|
gnarkHash "github.com/consensys/gnark-crypto/hash"
|
||||||
)
|
)
|
||||||
@@ -31,10 +29,6 @@ import (
|
|||||||
// =============================================================================
|
// =============================================================================
|
||||||
// Threshold Constants
|
// Threshold Constants
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
//
|
|
||||||
// These thresholds define the batch size above which GPU acceleration
|
|
||||||
// provides better performance than CPU. Below threshold, the GPU dispatch
|
|
||||||
// overhead exceeds the compute savings.
|
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// ThresholdPoseidon2 is the minimum batch size for GPU Poseidon2 hashing.
|
// ThresholdPoseidon2 is the minimum batch size for GPU Poseidon2 hashing.
|
||||||
@@ -64,87 +58,13 @@ var (
|
|||||||
ErrGPUUnavailable = errors.New("GPU acceleration unavailable")
|
ErrGPUUnavailable = errors.New("GPU acceleration unavailable")
|
||||||
)
|
)
|
||||||
|
|
||||||
// =============================================================================
|
|
||||||
// GPU Function Hooks (set by platform-specific files: zk_metal.go, zk_cuda.go)
|
|
||||||
// =============================================================================
|
|
||||||
|
|
||||||
// GPUHooks contains function pointers to GPU implementations.
|
|
||||||
// These are set by platform-specific init() functions.
|
|
||||||
type GPUHooks struct {
|
|
||||||
// Poseidon2 batch hash
|
|
||||||
HashPair func(left, right []Fr256) ([]Fr256, error)
|
|
||||||
// Merkle layer computation
|
|
||||||
MerkleLayer func(nodes []Fr256) ([]Fr256, error)
|
|
||||||
// Complete Merkle tree
|
|
||||||
MerkleTree func(leaves []Fr256) ([]Fr256, error)
|
|
||||||
// Batch commitment
|
|
||||||
BatchCommitment func(values, blindings, salts []Fr256) ([]Fr256, error)
|
|
||||||
// Batch nullifier
|
|
||||||
BatchNullifier func(keys, commitments, indices []Fr256) ([]Fr256, error)
|
|
||||||
}
|
|
||||||
|
|
||||||
// gpuHooks is the global GPU hooks instance.
|
|
||||||
// Set by zk_metal.go or zk_cuda.go depending on platform.
|
|
||||||
var gpuHooks *GPUHooks
|
|
||||||
|
|
||||||
// RegisterGPUHooks registers GPU implementation functions.
|
|
||||||
// Called by platform-specific init() functions.
|
|
||||||
func RegisterGPUHooks(hooks *GPUHooks) {
|
|
||||||
gpuHooks = hooks
|
|
||||||
}
|
|
||||||
|
|
||||||
// hasGPUHook returns true if a specific GPU hook is available.
|
|
||||||
func hasGPUHook(hook interface{}) bool {
|
|
||||||
return hook != nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
// Field Element Type (BN254 Fr - 256-bit)
|
// Field Element Type (BN254 Fr - 256-bit)
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
|
|
||||||
// Fr256 represents a 256-bit field element (BN254 scalar field).
|
// Fr256 represents a 256-bit field element (BN254 scalar field).
|
||||||
// Uses 4 x 64-bit limbs in little-endian order.
|
// Uses 4 x 64-bit limbs in little-endian order.
|
||||||
type Fr256 [4]uint64
|
type Fr256 = luxgpu.Fr256
|
||||||
|
|
||||||
// ToGnark converts Fr256 to gnark-crypto fr.Element.
|
|
||||||
func (f *Fr256) ToGnark() fr.Element {
|
|
||||||
var e fr.Element
|
|
||||||
e[0] = f[0]
|
|
||||||
e[1] = f[1]
|
|
||||||
e[2] = f[2]
|
|
||||||
e[3] = f[3]
|
|
||||||
return e
|
|
||||||
}
|
|
||||||
|
|
||||||
// FromGnark converts gnark-crypto fr.Element to Fr256.
|
|
||||||
func (f *Fr256) FromGnark(e *fr.Element) {
|
|
||||||
f[0] = e[0]
|
|
||||||
f[1] = e[1]
|
|
||||||
f[2] = e[2]
|
|
||||||
f[3] = e[3]
|
|
||||||
}
|
|
||||||
|
|
||||||
// Bytes returns the Fr256 as a 32-byte big-endian slice.
|
|
||||||
func (f *Fr256) Bytes() []byte {
|
|
||||||
buf := make([]byte, 32)
|
|
||||||
binary.BigEndian.PutUint64(buf[0:8], f[3])
|
|
||||||
binary.BigEndian.PutUint64(buf[8:16], f[2])
|
|
||||||
binary.BigEndian.PutUint64(buf[16:24], f[1])
|
|
||||||
binary.BigEndian.PutUint64(buf[24:32], f[0])
|
|
||||||
return buf
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetBytes sets the Fr256 from a 32-byte big-endian slice.
|
|
||||||
func (f *Fr256) SetBytes(buf []byte) error {
|
|
||||||
if len(buf) != 32 {
|
|
||||||
return ErrInvalidInput
|
|
||||||
}
|
|
||||||
f[3] = binary.BigEndian.Uint64(buf[0:8])
|
|
||||||
f[2] = binary.BigEndian.Uint64(buf[8:16])
|
|
||||||
f[1] = binary.BigEndian.Uint64(buf[16:24])
|
|
||||||
f[0] = binary.BigEndian.Uint64(buf[24:32])
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
// ZK Context
|
// ZK Context
|
||||||
@@ -170,11 +90,10 @@ var (
|
|||||||
// GetZKContext returns the global ZK context.
|
// GetZKContext returns the global ZK context.
|
||||||
func GetZKContext() *ZKContext {
|
func GetZKContext() *ZKContext {
|
||||||
initOnce.Do(func() {
|
initOnce.Do(func() {
|
||||||
// GPU is available if hooks are registered (by zk_metal.go or zk_cuda.go)
|
gpuEnabled := luxgpu.ZKGPUAvailable()
|
||||||
gpuEnabled := gpuHooks != nil
|
deviceName := luxgpu.ZKGetBackend()
|
||||||
deviceName := "CPU (gnark-crypto)"
|
if !gpuEnabled {
|
||||||
if gpuEnabled {
|
deviceName = "CPU (gnark-crypto)"
|
||||||
deviceName = "GPU (Metal/CUDA)"
|
|
||||||
}
|
}
|
||||||
defaultZKContext = &ZKContext{
|
defaultZKContext = &ZKContext{
|
||||||
gpuEnabled: gpuEnabled,
|
gpuEnabled: gpuEnabled,
|
||||||
@@ -217,17 +136,15 @@ var poseidon2HasherPool = sync.Pool{
|
|||||||
}
|
}
|
||||||
|
|
||||||
// poseidon2HashPairCPU computes Poseidon2(left, right) using gnark-crypto.
|
// poseidon2HashPairCPU computes Poseidon2(left, right) using gnark-crypto.
|
||||||
// Uses Merkle-Damgard construction for 2-to-1 compression.
|
|
||||||
func poseidon2HashPairCPU(left, right *Fr256) Fr256 {
|
func poseidon2HashPairCPU(left, right *Fr256) Fr256 {
|
||||||
// Get a hasher from the pool
|
|
||||||
h := poseidon2HasherPool.Get().(gnarkHash.StateStorer)
|
h := poseidon2HasherPool.Get().(gnarkHash.StateStorer)
|
||||||
defer poseidon2HasherPool.Put(h)
|
defer poseidon2HasherPool.Put(h)
|
||||||
|
|
||||||
h.Reset()
|
h.Reset()
|
||||||
|
|
||||||
// Write left and right as bytes
|
// Write left and right as bytes
|
||||||
leftBytes := left.Bytes()
|
leftBytes := fr256ToBytes(left)
|
||||||
rightBytes := right.Bytes()
|
rightBytes := fr256ToBytes(right)
|
||||||
|
|
||||||
_, _ = h.Write(leftBytes)
|
_, _ = h.Write(leftBytes)
|
||||||
_, _ = h.Write(rightBytes)
|
_, _ = h.Write(rightBytes)
|
||||||
@@ -237,10 +154,36 @@ func poseidon2HashPairCPU(left, right *Fr256) Fr256 {
|
|||||||
|
|
||||||
// Convert back to Fr256
|
// Convert back to Fr256
|
||||||
var out Fr256
|
var out Fr256
|
||||||
_ = out.SetBytes(resultBytes[:32])
|
_ = fr256FromBytes(&out, resultBytes[:32])
|
||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// fr256ToBytes converts Fr256 to 32-byte big-endian slice.
|
||||||
|
func fr256ToBytes(f *Fr256) []byte {
|
||||||
|
buf := make([]byte, 32)
|
||||||
|
binary.BigEndian.PutUint64(buf[0:8], f[3])
|
||||||
|
binary.BigEndian.PutUint64(buf[8:16], f[2])
|
||||||
|
binary.BigEndian.PutUint64(buf[16:24], f[1])
|
||||||
|
binary.BigEndian.PutUint64(buf[24:32], f[0])
|
||||||
|
return buf
|
||||||
|
}
|
||||||
|
|
||||||
|
// fr256FromBytes sets Fr256 from 32-byte big-endian slice.
|
||||||
|
func fr256FromBytes(f *Fr256, buf []byte) error {
|
||||||
|
if len(buf) != 32 {
|
||||||
|
return ErrInvalidInput
|
||||||
|
}
|
||||||
|
f[3] = binary.BigEndian.Uint64(buf[0:8])
|
||||||
|
f[2] = binary.BigEndian.Uint64(buf[8:16])
|
||||||
|
f[1] = binary.BigEndian.Uint64(buf[16:24])
|
||||||
|
f[0] = binary.BigEndian.Uint64(buf[24:32])
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// =============================================================================
|
||||||
|
// Public API - Poseidon2
|
||||||
|
// =============================================================================
|
||||||
|
|
||||||
// Poseidon2HashPair computes Poseidon2(left, right) with automatic routing.
|
// Poseidon2HashPair computes Poseidon2(left, right) with automatic routing.
|
||||||
func (z *ZKContext) Poseidon2HashPair(left, right *Fr256) Fr256 {
|
func (z *ZKContext) Poseidon2HashPair(left, right *Fr256) Fr256 {
|
||||||
// Single hash always uses CPU (GPU overhead not worth it)
|
// Single hash always uses CPU (GPU overhead not worth it)
|
||||||
@@ -261,11 +204,11 @@ func (z *ZKContext) Poseidon2BatchHashPair(left, right []Fr256) ([]Fr256, error)
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check threshold for GPU routing
|
// Check threshold for GPU routing
|
||||||
if z.GPUEnabled() && n >= ThresholdPoseidon2 && gpuHooks != nil && gpuHooks.HashPair != nil {
|
if z.GPUEnabled() && n >= ThresholdPoseidon2 {
|
||||||
z.mu.Lock()
|
z.mu.Lock()
|
||||||
z.gpuCalls++
|
z.gpuCalls++
|
||||||
z.mu.Unlock()
|
z.mu.Unlock()
|
||||||
return gpuHooks.HashPair(left, right)
|
return luxgpu.Poseidon2Hash(left, right)
|
||||||
}
|
}
|
||||||
|
|
||||||
// CPU path using gnark-crypto
|
// CPU path using gnark-crypto
|
||||||
@@ -286,8 +229,6 @@ func (z *ZKContext) Poseidon2BatchHashPair(left, right []Fr256) ([]Fr256, error)
|
|||||||
// =============================================================================
|
// =============================================================================
|
||||||
|
|
||||||
// Poseidon2MerkleLayer computes one layer of a Merkle tree.
|
// Poseidon2MerkleLayer computes one layer of a Merkle tree.
|
||||||
// Input: current layer nodes (must be even count)
|
|
||||||
// Output: parent nodes (half the size)
|
|
||||||
func (z *ZKContext) Poseidon2MerkleLayer(nodes []Fr256) ([]Fr256, error) {
|
func (z *ZKContext) Poseidon2MerkleLayer(nodes []Fr256) ([]Fr256, error) {
|
||||||
n := len(nodes)
|
n := len(nodes)
|
||||||
if n == 0 || n%2 != 0 {
|
if n == 0 || n%2 != 0 {
|
||||||
@@ -297,11 +238,11 @@ func (z *ZKContext) Poseidon2MerkleLayer(nodes []Fr256) ([]Fr256, error) {
|
|||||||
parentCount := n / 2
|
parentCount := n / 2
|
||||||
|
|
||||||
// Check threshold for GPU routing
|
// Check threshold for GPU routing
|
||||||
if z.GPUEnabled() && parentCount >= ThresholdMerkle && gpuHooks != nil && gpuHooks.MerkleLayer != nil {
|
if z.GPUEnabled() && parentCount >= ThresholdMerkle {
|
||||||
z.mu.Lock()
|
z.mu.Lock()
|
||||||
z.gpuCalls++
|
z.gpuCalls++
|
||||||
z.mu.Unlock()
|
z.mu.Unlock()
|
||||||
return gpuHooks.MerkleLayer(nodes)
|
return luxgpu.MerkleLayer(nodes)
|
||||||
}
|
}
|
||||||
|
|
||||||
// CPU path
|
// CPU path
|
||||||
@@ -318,7 +259,6 @@ func (z *ZKContext) Poseidon2MerkleLayer(nodes []Fr256) ([]Fr256, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Poseidon2MerkleRoot computes the Merkle root from leaves.
|
// Poseidon2MerkleRoot computes the Merkle root from leaves.
|
||||||
// Leaves must be a power of 2.
|
|
||||||
func (z *ZKContext) Poseidon2MerkleRoot(leaves []Fr256) (Fr256, error) {
|
func (z *ZKContext) Poseidon2MerkleRoot(leaves []Fr256) (Fr256, error) {
|
||||||
n := len(leaves)
|
n := len(leaves)
|
||||||
if n == 0 || (n&(n-1)) != 0 {
|
if n == 0 || (n&(n-1)) != 0 {
|
||||||
@@ -329,7 +269,15 @@ func (z *ZKContext) Poseidon2MerkleRoot(leaves []Fr256) (Fr256, error) {
|
|||||||
return leaves[0], nil
|
return leaves[0], nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build tree layer by layer
|
// Use GPU if available and above threshold
|
||||||
|
if z.GPUEnabled() && n >= ThresholdMerkle*2 {
|
||||||
|
z.mu.Lock()
|
||||||
|
z.gpuCalls++
|
||||||
|
z.mu.Unlock()
|
||||||
|
return luxgpu.MerkleRoot(leaves)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CPU path - build tree layer by layer
|
||||||
current := leaves
|
current := leaves
|
||||||
for len(current) > 1 {
|
for len(current) > 1 {
|
||||||
next, err := z.Poseidon2MerkleLayer(current)
|
next, err := z.Poseidon2MerkleLayer(current)
|
||||||
@@ -342,9 +290,7 @@ func (z *ZKContext) Poseidon2MerkleRoot(leaves []Fr256) (Fr256, error) {
|
|||||||
return current[0], nil
|
return current[0], nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Poseidon2MerkleTree builds a complete Merkle tree and returns all internal nodes.
|
// Poseidon2MerkleTree builds a complete Merkle tree.
|
||||||
// Returns: [layer_n-1, layer_n-2, ..., layer_0, root]
|
|
||||||
// Total internal nodes: num_leaves - 1
|
|
||||||
func (z *ZKContext) Poseidon2MerkleTree(leaves []Fr256) ([]Fr256, error) {
|
func (z *ZKContext) Poseidon2MerkleTree(leaves []Fr256) ([]Fr256, error) {
|
||||||
n := len(leaves)
|
n := len(leaves)
|
||||||
if n == 0 || (n&(n-1)) != 0 {
|
if n == 0 || (n&(n-1)) != 0 {
|
||||||
@@ -355,7 +301,15 @@ func (z *ZKContext) Poseidon2MerkleTree(leaves []Fr256) ([]Fr256, error) {
|
|||||||
return []Fr256{leaves[0]}, nil
|
return []Fr256{leaves[0]}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Collect all internal nodes
|
// Use GPU if available and above threshold
|
||||||
|
if z.GPUEnabled() && n >= ThresholdMerkle*2 {
|
||||||
|
z.mu.Lock()
|
||||||
|
z.gpuCalls++
|
||||||
|
z.mu.Unlock()
|
||||||
|
return luxgpu.MerkleTree(leaves)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CPU path - collect all internal nodes
|
||||||
allNodes := make([]Fr256, 0, n-1)
|
allNodes := make([]Fr256, 0, n-1)
|
||||||
current := leaves
|
current := leaves
|
||||||
|
|
||||||
@@ -375,16 +329,14 @@ func (z *ZKContext) Poseidon2MerkleTree(leaves []Fr256) ([]Fr256, error) {
|
|||||||
// Commitment and Nullifier Operations
|
// Commitment and Nullifier Operations
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
|
|
||||||
// Poseidon2Commitment computes commitment = Poseidon2(value, blinding, salt).
|
// Poseidon2Commitment computes commitment = Poseidon2(Poseidon2(value, blinding), salt).
|
||||||
func (z *ZKContext) Poseidon2Commitment(value, blinding, salt *Fr256) Fr256 {
|
func (z *ZKContext) Poseidon2Commitment(value, blinding, salt *Fr256) Fr256 {
|
||||||
// Hash: Poseidon2(Poseidon2(value, blinding), salt)
|
|
||||||
intermediate := poseidon2HashPairCPU(value, blinding)
|
intermediate := poseidon2HashPairCPU(value, blinding)
|
||||||
return poseidon2HashPairCPU(&intermediate, salt)
|
return poseidon2HashPairCPU(&intermediate, salt)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Poseidon2Nullifier computes nullifier = Poseidon2(key, commitment, index).
|
// Poseidon2Nullifier computes nullifier = Poseidon2(Poseidon2(key, commitment), index).
|
||||||
func (z *ZKContext) Poseidon2Nullifier(key, commitment, index *Fr256) Fr256 {
|
func (z *ZKContext) Poseidon2Nullifier(key, commitment, index *Fr256) Fr256 {
|
||||||
// Hash: Poseidon2(Poseidon2(key, commitment), index)
|
|
||||||
intermediate := poseidon2HashPairCPU(key, commitment)
|
intermediate := poseidon2HashPairCPU(key, commitment)
|
||||||
return poseidon2HashPairCPU(&intermediate, index)
|
return poseidon2HashPairCPU(&intermediate, index)
|
||||||
}
|
}
|
||||||
@@ -400,11 +352,11 @@ func (z *ZKContext) BatchCommitment(values, blindings, salts []Fr256) ([]Fr256,
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check threshold for GPU routing
|
// Check threshold for GPU routing
|
||||||
if z.GPUEnabled() && n >= ThresholdCommitment && gpuHooks != nil && gpuHooks.BatchCommitment != nil {
|
if z.GPUEnabled() && n >= ThresholdCommitment {
|
||||||
z.mu.Lock()
|
z.mu.Lock()
|
||||||
z.gpuCalls++
|
z.gpuCalls++
|
||||||
z.mu.Unlock()
|
z.mu.Unlock()
|
||||||
return gpuHooks.BatchCommitment(values, blindings, salts)
|
return luxgpu.BatchCommitment(values, blindings, salts)
|
||||||
}
|
}
|
||||||
|
|
||||||
// CPU path
|
// CPU path
|
||||||
@@ -431,11 +383,11 @@ func (z *ZKContext) BatchNullifier(keys, commitments, indices []Fr256) ([]Fr256,
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check threshold for GPU routing
|
// Check threshold for GPU routing
|
||||||
if z.GPUEnabled() && n >= ThresholdCommitment && gpuHooks != nil && gpuHooks.BatchNullifier != nil {
|
if z.GPUEnabled() && n >= ThresholdCommitment {
|
||||||
z.mu.Lock()
|
z.mu.Lock()
|
||||||
z.gpuCalls++
|
z.gpuCalls++
|
||||||
z.mu.Unlock()
|
z.mu.Unlock()
|
||||||
return gpuHooks.BatchNullifier(keys, commitments, indices)
|
return luxgpu.BatchNullifier(keys, commitments, indices)
|
||||||
}
|
}
|
||||||
|
|
||||||
// CPU path
|
// CPU path
|
||||||
|
|||||||
-300
@@ -1,300 +0,0 @@
|
|||||||
//go:build darwin && arm64 && cgo && gpu
|
|
||||||
|
|
||||||
// Package gpu provides GPU-accelerated ZK operations via Metal.
|
|
||||||
//
|
|
||||||
// This file implements CGO bindings to the luxcpp/crypto Metal ZK library.
|
|
||||||
// Build with: CGO_ENABLED=1 go build -tags "darwin,arm64,cgo"
|
|
||||||
package gpu
|
|
||||||
|
|
||||||
/*
|
|
||||||
#cgo CFLAGS: -I/Users/z/work/luxcpp/crypto/include
|
|
||||||
#cgo LDFLAGS: -L/Users/z/work/luxcpp/crypto/build -lluxcrypto -framework Metal -framework Foundation -lc++
|
|
||||||
|
|
||||||
#include <lux/crypto/metal_zk.h>
|
|
||||||
#include <stdlib.h>
|
|
||||||
*/
|
|
||||||
import "C"
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"sync"
|
|
||||||
"unsafe"
|
|
||||||
)
|
|
||||||
|
|
||||||
// =============================================================================
|
|
||||||
// Metal ZK Context (CGO)
|
|
||||||
// =============================================================================
|
|
||||||
|
|
||||||
// metalContext holds the CGO context for Metal operations.
|
|
||||||
type metalContext struct {
|
|
||||||
ctx *C.MetalZKContext
|
|
||||||
mu sync.Mutex
|
|
||||||
enabled bool
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
|
||||||
metalCtx *metalContext
|
|
||||||
metalCtxOnce sync.Once
|
|
||||||
)
|
|
||||||
|
|
||||||
// getMetalContext returns the singleton Metal context.
|
|
||||||
func getMetalContext() *metalContext {
|
|
||||||
metalCtxOnce.Do(func() {
|
|
||||||
metalCtx = &metalContext{}
|
|
||||||
metalCtx.ctx = C.metal_zk_init()
|
|
||||||
metalCtx.enabled = metalCtx.ctx != nil
|
|
||||||
})
|
|
||||||
return metalCtx
|
|
||||||
}
|
|
||||||
|
|
||||||
// MetalAvailable returns true if Metal ZK acceleration is available.
|
|
||||||
func MetalAvailable() bool {
|
|
||||||
return bool(C.metal_zk_available())
|
|
||||||
}
|
|
||||||
|
|
||||||
// =============================================================================
|
|
||||||
// Fr256 CGO Helpers
|
|
||||||
// =============================================================================
|
|
||||||
|
|
||||||
// toCFr256 converts Fr256 slice to C array pointer.
|
|
||||||
// Returns the C array and a cleanup function.
|
|
||||||
func toCFr256Slice(elements []Fr256) (*C.Fr256, func()) {
|
|
||||||
if len(elements) == 0 {
|
|
||||||
return nil, func() {}
|
|
||||||
}
|
|
||||||
// Fr256 in Go is [4]uint64, same layout as C Fr256
|
|
||||||
return (*C.Fr256)(unsafe.Pointer(&elements[0])), func() {}
|
|
||||||
}
|
|
||||||
|
|
||||||
// allocCFr256 allocates a C array for Fr256 results.
|
|
||||||
func allocCFr256(count int) []Fr256 {
|
|
||||||
return make([]Fr256, count)
|
|
||||||
}
|
|
||||||
|
|
||||||
// =============================================================================
|
|
||||||
// GPU Poseidon2 Hash Operations
|
|
||||||
// =============================================================================
|
|
||||||
|
|
||||||
// poseidon2HashPairGPU computes batch Poseidon2 hashes on GPU.
|
|
||||||
func poseidon2HashPairGPU(left, right []Fr256) ([]Fr256, error) {
|
|
||||||
ctx := getMetalContext()
|
|
||||||
if !ctx.enabled {
|
|
||||||
return nil, errors.New("Metal not available")
|
|
||||||
}
|
|
||||||
|
|
||||||
n := len(left)
|
|
||||||
if n != len(right) || n == 0 {
|
|
||||||
return nil, ErrSizeMismatch
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx.mu.Lock()
|
|
||||||
defer ctx.mu.Unlock()
|
|
||||||
|
|
||||||
// Allocate output
|
|
||||||
output := allocCFr256(n)
|
|
||||||
|
|
||||||
// Get C pointers
|
|
||||||
outPtr, _ := toCFr256Slice(output)
|
|
||||||
leftPtr, _ := toCFr256Slice(left)
|
|
||||||
rightPtr, _ := toCFr256Slice(right)
|
|
||||||
|
|
||||||
// Call Metal function
|
|
||||||
ret := C.metal_zk_poseidon2_hash_pair(
|
|
||||||
ctx.ctx,
|
|
||||||
outPtr,
|
|
||||||
leftPtr,
|
|
||||||
rightPtr,
|
|
||||||
C.uint32_t(n),
|
|
||||||
)
|
|
||||||
|
|
||||||
if ret != C.METAL_ZK_SUCCESS {
|
|
||||||
return nil, metalError(ret)
|
|
||||||
}
|
|
||||||
|
|
||||||
return output, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// poseidon2MerkleLayerGPU computes one Merkle layer on GPU.
|
|
||||||
func poseidon2MerkleLayerGPU(nodes []Fr256) ([]Fr256, error) {
|
|
||||||
ctx := getMetalContext()
|
|
||||||
if !ctx.enabled {
|
|
||||||
return nil, errors.New("Metal not available")
|
|
||||||
}
|
|
||||||
|
|
||||||
n := len(nodes)
|
|
||||||
if n == 0 || n%2 != 0 {
|
|
||||||
return nil, ErrInvalidInput
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx.mu.Lock()
|
|
||||||
defer ctx.mu.Unlock()
|
|
||||||
|
|
||||||
parentCount := n / 2
|
|
||||||
output := allocCFr256(parentCount)
|
|
||||||
|
|
||||||
outPtr, _ := toCFr256Slice(output)
|
|
||||||
nodesPtr, _ := toCFr256Slice(nodes)
|
|
||||||
|
|
||||||
ret := C.metal_zk_poseidon2_merkle_layer(
|
|
||||||
ctx.ctx,
|
|
||||||
outPtr,
|
|
||||||
nodesPtr,
|
|
||||||
C.uint32_t(n),
|
|
||||||
)
|
|
||||||
|
|
||||||
if ret != C.METAL_ZK_SUCCESS {
|
|
||||||
return nil, metalError(ret)
|
|
||||||
}
|
|
||||||
|
|
||||||
return output, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// poseidon2MerkleTreeGPU builds a complete Merkle tree on GPU.
|
|
||||||
func poseidon2MerkleTreeGPU(leaves []Fr256) ([]Fr256, error) {
|
|
||||||
ctx := getMetalContext()
|
|
||||||
if !ctx.enabled {
|
|
||||||
return nil, errors.New("Metal not available")
|
|
||||||
}
|
|
||||||
|
|
||||||
n := len(leaves)
|
|
||||||
if n == 0 || (n&(n-1)) != 0 {
|
|
||||||
return nil, ErrNotPowerOfTwo
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx.mu.Lock()
|
|
||||||
defer ctx.mu.Unlock()
|
|
||||||
|
|
||||||
// Tree has n-1 internal nodes
|
|
||||||
tree := allocCFr256(n - 1)
|
|
||||||
|
|
||||||
treePtr, _ := toCFr256Slice(tree)
|
|
||||||
leavesPtr, _ := toCFr256Slice(leaves)
|
|
||||||
|
|
||||||
ret := C.metal_zk_poseidon2_merkle_tree(
|
|
||||||
ctx.ctx,
|
|
||||||
treePtr,
|
|
||||||
leavesPtr,
|
|
||||||
C.uint32_t(n),
|
|
||||||
)
|
|
||||||
|
|
||||||
if ret != C.METAL_ZK_SUCCESS {
|
|
||||||
return nil, metalError(ret)
|
|
||||||
}
|
|
||||||
|
|
||||||
return tree, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// batchCommitmentGPU computes batch commitments on GPU.
|
|
||||||
func batchCommitmentGPU(values, blindings, salts []Fr256) ([]Fr256, error) {
|
|
||||||
ctx := getMetalContext()
|
|
||||||
if !ctx.enabled {
|
|
||||||
return nil, errors.New("Metal not available")
|
|
||||||
}
|
|
||||||
|
|
||||||
n := len(values)
|
|
||||||
if n != len(blindings) || n != len(salts) || n == 0 {
|
|
||||||
return nil, ErrSizeMismatch
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx.mu.Lock()
|
|
||||||
defer ctx.mu.Unlock()
|
|
||||||
|
|
||||||
output := allocCFr256(n)
|
|
||||||
|
|
||||||
outPtr, _ := toCFr256Slice(output)
|
|
||||||
valuesPtr, _ := toCFr256Slice(values)
|
|
||||||
blindingsPtr, _ := toCFr256Slice(blindings)
|
|
||||||
saltsPtr, _ := toCFr256Slice(salts)
|
|
||||||
|
|
||||||
ret := C.metal_zk_batch_commitment(
|
|
||||||
ctx.ctx,
|
|
||||||
outPtr,
|
|
||||||
valuesPtr,
|
|
||||||
blindingsPtr,
|
|
||||||
saltsPtr,
|
|
||||||
C.uint32_t(n),
|
|
||||||
)
|
|
||||||
|
|
||||||
if ret != C.METAL_ZK_SUCCESS {
|
|
||||||
return nil, metalError(ret)
|
|
||||||
}
|
|
||||||
|
|
||||||
return output, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// batchNullifierGPU computes batch nullifiers on GPU.
|
|
||||||
func batchNullifierGPU(keys, commitments, indices []Fr256) ([]Fr256, error) {
|
|
||||||
ctx := getMetalContext()
|
|
||||||
if !ctx.enabled {
|
|
||||||
return nil, errors.New("Metal not available")
|
|
||||||
}
|
|
||||||
|
|
||||||
n := len(keys)
|
|
||||||
if n != len(commitments) || n != len(indices) || n == 0 {
|
|
||||||
return nil, ErrSizeMismatch
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx.mu.Lock()
|
|
||||||
defer ctx.mu.Unlock()
|
|
||||||
|
|
||||||
output := allocCFr256(n)
|
|
||||||
|
|
||||||
outPtr, _ := toCFr256Slice(output)
|
|
||||||
keysPtr, _ := toCFr256Slice(keys)
|
|
||||||
commitmentsPtr, _ := toCFr256Slice(commitments)
|
|
||||||
indicesPtr, _ := toCFr256Slice(indices)
|
|
||||||
|
|
||||||
ret := C.metal_zk_batch_nullifier(
|
|
||||||
ctx.ctx,
|
|
||||||
outPtr,
|
|
||||||
keysPtr,
|
|
||||||
commitmentsPtr,
|
|
||||||
indicesPtr,
|
|
||||||
C.uint32_t(n),
|
|
||||||
)
|
|
||||||
|
|
||||||
if ret != C.METAL_ZK_SUCCESS {
|
|
||||||
return nil, metalError(ret)
|
|
||||||
}
|
|
||||||
|
|
||||||
return output, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// =============================================================================
|
|
||||||
// Error Handling
|
|
||||||
// =============================================================================
|
|
||||||
|
|
||||||
func metalError(code C.int) error {
|
|
||||||
switch code {
|
|
||||||
case C.METAL_ZK_ERROR_NO_DEVICE:
|
|
||||||
return errors.New("metal: no device available")
|
|
||||||
case C.METAL_ZK_ERROR_NO_SHADER:
|
|
||||||
return errors.New("metal: shader not found")
|
|
||||||
case C.METAL_ZK_ERROR_ALLOC:
|
|
||||||
return errors.New("metal: allocation failed")
|
|
||||||
case C.METAL_ZK_ERROR_NULL_PTR:
|
|
||||||
return errors.New("metal: null pointer")
|
|
||||||
case C.METAL_ZK_ERROR_INVALID:
|
|
||||||
return errors.New("metal: invalid argument")
|
|
||||||
case C.METAL_ZK_ERROR_SIZE:
|
|
||||||
return errors.New("metal: invalid size")
|
|
||||||
default:
|
|
||||||
return errors.New("metal: unknown error")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// =============================================================================
|
|
||||||
// GPU Hook Registration
|
|
||||||
// =============================================================================
|
|
||||||
|
|
||||||
// init registers Metal GPU hooks when available.
|
|
||||||
func init() {
|
|
||||||
if MetalAvailable() {
|
|
||||||
RegisterGPUHooks(&GPUHooks{
|
|
||||||
HashPair: poseidon2HashPairGPU,
|
|
||||||
MerkleLayer: poseidon2MerkleLayerGPU,
|
|
||||||
MerkleTree: poseidon2MerkleTreeGPU,
|
|
||||||
BatchCommitment: batchCommitmentGPU,
|
|
||||||
BatchNullifier: batchNullifierGPU,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+33
-29
@@ -2,50 +2,54 @@ package gpu
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/consensys/gnark-crypto/ecc/bn254/fr"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// TestFr256Conversion tests Fr256 <-> gnark-crypto conversion.
|
// TestFr256Type tests that Fr256 is correctly defined.
|
||||||
func TestFr256Conversion(t *testing.T) {
|
func TestFr256Type(t *testing.T) {
|
||||||
// Create a random field element
|
|
||||||
var e fr.Element
|
|
||||||
e.SetRandom()
|
|
||||||
|
|
||||||
// Convert to Fr256
|
|
||||||
var f Fr256
|
|
||||||
f.FromGnark(&e)
|
|
||||||
|
|
||||||
// Convert back
|
|
||||||
e2 := f.ToGnark()
|
|
||||||
|
|
||||||
// Should be equal
|
|
||||||
if !e.Equal(&e2) {
|
|
||||||
t.Errorf("Conversion roundtrip failed: got %v, want %v", e2, e)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestFr256Bytes tests Fr256 byte serialization.
|
|
||||||
func TestFr256Bytes(t *testing.T) {
|
|
||||||
var f Fr256
|
var f Fr256
|
||||||
f[0] = 0x1234567890abcdef
|
f[0] = 0x1234567890abcdef
|
||||||
f[1] = 0xfedcba0987654321
|
f[1] = 0xfedcba0987654321
|
||||||
f[2] = 0xabcdef0123456789
|
f[2] = 0xabcdef0123456789
|
||||||
f[3] = 0x9876543210fedcba
|
f[3] = 0x9876543210fedcba
|
||||||
|
|
||||||
// Serialize
|
// Verify 4 limbs of 64 bits each
|
||||||
buf := f.Bytes()
|
if len(f) != 4 {
|
||||||
|
t.Fatalf("Expected 4 limbs, got %d", len(f))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test copy semantics
|
||||||
|
f2 := f
|
||||||
|
if f != f2 {
|
||||||
|
t.Error("Copy failed")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Modify copy should not affect original
|
||||||
|
f2[0] = 0
|
||||||
|
if f[0] == 0 {
|
||||||
|
t.Error("Copy modified original")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFr256ByteHelpers tests internal byte conversion helpers.
|
||||||
|
func TestFr256ByteHelpers(t *testing.T) {
|
||||||
|
var f Fr256
|
||||||
|
f[0] = 0x1234567890abcdef
|
||||||
|
f[1] = 0xfedcba0987654321
|
||||||
|
f[2] = 0xabcdef0123456789
|
||||||
|
f[3] = 0x9876543210fedcba
|
||||||
|
|
||||||
|
// Use internal helpers via roundtrip through hash
|
||||||
|
// The byte conversion is exercised in poseidon2HashPairCPU
|
||||||
|
buf := fr256ToBytes(&f)
|
||||||
if len(buf) != 32 {
|
if len(buf) != 32 {
|
||||||
t.Fatalf("Expected 32 bytes, got %d", len(buf))
|
t.Fatalf("Expected 32 bytes, got %d", len(buf))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Deserialize
|
|
||||||
var f2 Fr256
|
var f2 Fr256
|
||||||
if err := f2.SetBytes(buf); err != nil {
|
if err := fr256FromBytes(&f2, buf); err != nil {
|
||||||
t.Fatalf("SetBytes failed: %v", err)
|
t.Fatalf("fr256FromBytes failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Should be equal
|
|
||||||
if f != f2 {
|
if f != f2 {
|
||||||
t.Errorf("Byte roundtrip failed: got %v, want %v", f2, f)
|
t.Errorf("Byte roundtrip failed: got %v, want %v", f2, f)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user