perf: increase G window to 12, add batch inversion for table building

- Increase WINDOW_G from 8 to 12 for mulDoubleG (verify): reduces
  G-side additions from ~32 to ~22 per verification (saves ~110 field ops)
- Add batchToAffine using Montgomery's trick: 1 inversion + 3(n-1) muls
  instead of n individual inversions. Critical for the 1024-entry table.
- Table size: 1024 entries × 64 bytes = ~128KB (lazy, built on first use)

https://claude.ai/code/session_01BhU63WUe9AhikZxRdw3Lpg
This commit is contained in:
Claude
2026-04-06 14:55:14 +00:00
parent 2661bf783d
commit 80dc164c7c
@@ -134,14 +134,17 @@ internal object ECPoint {
private val B = longArrayOf(7L, 0L, 0L, 0L) private val B = longArrayOf(7L, 0L, 0L, 0L)
/** /**
* wNAF window width for the G-side of scalar multiplication. * wNAF window width for the G-side of scalar multiplication (mulDoubleG/verify).
* Width w uses a table of 2^(w-2) odd multiples. Larger windows mean fewer * Width w uses a table of 2^(w-2) odd multiples. Larger windows mean fewer
* additions but more precomputed storage: * additions but more precomputed storage (lazily allocated on first use):
* w=5: 8 entries, ~26 adds per 128-bit scalar (~1KB table) * w=8: 64 entries, ~16 adds per 128-bit scalar (~8KB table)
* w=8: 64 entries, ~16 adds per 128-bit scalar (~8KB table) * w=10: 256 entries, ~13 adds per 128-bit scalar (~32KB table)
* w=10: 256 entries,~13 adds per 128-bit scalar (~32KB table) * w=12: 1024 entries,~11 adds per 128-bit scalar (~128KB table)
* w=14: 4096 entries,~9 adds per 128-bit scalar (~512KB table)
* C libsecp256k1 uses w=15 (8192 entries, 1MB) precomputed at compile time.
* w=12 is a good tradeoff for runtime-computed tables on JVM.
*/ */
private const val WINDOW_G = 8 private const val WINDOW_G = 12
private val G_TABLE_SIZE = 1 shl (WINDOW_G - 2) // 64 for w=8 private val G_TABLE_SIZE = 1 shl (WINDOW_G - 2) // 64 for w=8
/** /**
@@ -165,12 +168,55 @@ internal object ECPoint {
jac[0].copyFrom(g) jac[0].copyFrom(g)
for (i in 1 until G_TABLE_SIZE) addPoints(jac[i], jac[i - 1], g2) for (i in 1 until G_TABLE_SIZE) addPoints(jac[i], jac[i - 1], g2)
return Array(G_TABLE_SIZE) { i -> return batchToAffine(jac)
val x = LongArray(4) }
val y = LongArray(4)
toAffine(jac[i], x, y) /**
AffinePoint(x, y) * Convert an array of Jacobian points to affine using Montgomery's batch inversion trick.
* Cost: 1 field inversion + 3(n-1) multiplications, instead of n inversions.
* This is critical for large precomputed tables (e.g., 1024 entries for WINDOW_G=12).
*/
internal fun batchToAffine(jac: Array<MutablePoint>): Array<AffinePoint> {
val n = jac.size
if (n == 0) return emptyArray()
// Step 1: compute prefix products of Z coordinates
// prods[i] = z[0] * z[1] * ... * z[i]
val prods = Array(n) { LongArray(4) }
jac[0].z.copyInto(prods[0])
for (i in 1 until n) {
FieldP.mul(prods[i], prods[i - 1], jac[i].z)
} }
// Step 2: invert the total product
val inv = LongArray(4)
FieldP.inv(inv, prods[n - 1])
// Step 3: recover individual inverses by multiplying back
// zInv[i] = product_inv * prods[i-1] = 1 / z[i]
val zInv = LongArray(4)
val zInv2 = LongArray(4)
val zInv3 = LongArray(4)
val result = Array(n) { AffinePoint() }
for (i in n - 1 downTo 1) {
// zInv = inv * prods[i-1] = 1/z[i]
FieldP.mul(zInv, inv, prods[i - 1])
// Update inv for next iteration: inv = inv * z[i] = 1/(z[0]*...*z[i-1])
FieldP.mul(inv, inv, jac[i].z)
// Convert to affine: x' = X/Z², y' = Y/Z³
FieldP.sqr(zInv2, zInv)
FieldP.mul(zInv3, zInv2, zInv)
FieldP.mul(result[i].x, jac[i].x, zInv2)
FieldP.mul(result[i].y, jac[i].y, zInv3)
}
// i=0: inv is now 1/z[0]
FieldP.sqr(zInv2, inv)
FieldP.mul(zInv3, zInv2, inv)
FieldP.mul(result[0].x, jac[0].x, zInv2)
FieldP.mul(result[0].y, jac[0].y, zInv3)
return result
} }
// ==================== Comb Method for G Multiplication ==================== // ==================== Comb Method for G Multiplication ====================
@@ -479,23 +525,27 @@ internal object ECPoint {
val wnaf1 = Glv.wnaf(split.k1, wnd, 129) val wnaf1 = Glv.wnaf(split.k1, wnd, 129)
val wnaf2 = Glv.wnaf(split.k2, wnd, 129) val wnaf2 = Glv.wnaf(split.k2, wnd, 129)
// P odd-multiples [1P, 3P, 5P, ..., 15P] via 2P stepping // P odd-multiples [1P, 3P, 5P, ..., 15P] via 2P stepping (Jacobian)
val p2 = MutablePoint() val p2 = MutablePoint()
doublePoint(p2, p, s) doublePoint(p2, p, s)
val pOdd = Array(tableSize) { MutablePoint() } val pOddJac = Array(tableSize) { MutablePoint() }
pOdd[0].copyFrom(p) pOddJac[0].copyFrom(p)
for (i in 1 until tableSize) addPoints(pOdd[i], pOdd[i - 1], p2, s) for (i in 1 until tableSize) addPoints(pOddJac[i], pOddJac[i - 1], p2, s)
// λ(P) odd-multiples: (β·X, Y, Z) in Jacobian // λ(P) odd-multiples: (β·X, Y, Z) in Jacobian
val pLamOdd = val pLamOddJac =
Array(tableSize) { i -> Array(tableSize) { i ->
val lp = MutablePoint() val lp = MutablePoint()
FieldP.mul(lp.x, pOdd[i].x, Glv.BETA, s.w) FieldP.mul(lp.x, pOddJac[i].x, Glv.BETA, s.w)
pOdd[i].y.copyInto(lp.y) pOddJac[i].y.copyInto(lp.y)
pOdd[i].z.copyInto(lp.z) pOddJac[i].z.copyInto(lp.z)
lp lp
} }
// Effective-affine: batch-convert both tables to affine for cheaper mixed additions
val pOdd = batchToAffine(pOddJac, s)
val pLamOdd = batchToAffine(pLamOddJac, s)
// Find highest non-zero digit // Find highest non-zero digit
var bits = maxOf(wnaf1.size, wnaf2.size) var bits = maxOf(wnaf1.size, wnaf2.size)
while (bits > 0) { while (bits > 0) {
@@ -506,12 +556,12 @@ internal object ECPoint {
out.setInfinity() out.setInfinity()
val tmp = MutablePoint() val tmp = MutablePoint()
val negJac = MutablePoint() val negY = LongArray(4)
for (i in bits - 1 downTo 0) { for (i in bits - 1 downTo 0) {
doublePoint(out, out, s) doublePoint(out, out, s)
addWnafJacobian(out, tmp, negJac, wnaf1, i, pOdd, split.negK1, s) addWnafMixed(out, tmp, negY, wnaf1, i, pOdd, split.negK1, s)
addWnafJacobian(out, tmp, negJac, wnaf2, i, pLamOdd, split.negK2, s) addWnafMixed(out, tmp, negY, wnaf2, i, pLamOdd, split.negK2, s)
} }
} }
@@ -599,19 +649,23 @@ internal object ECPoint {
// P odd-multiples [1P, 3P, 5P, ..., 15P] via 2P stepping (1 double + 7 adds) // P odd-multiples [1P, 3P, 5P, ..., 15P] via 2P stepping (1 double + 7 adds)
val p2 = MutablePoint() val p2 = MutablePoint()
doublePoint(p2, p, sc) doublePoint(p2, p, sc)
val pOdd = Array(pTableSize) { MutablePoint() } val pOddJac = Array(pTableSize) { MutablePoint() }
pOdd[0].copyFrom(p) pOddJac[0].copyFrom(p)
for (i in 1 until pTableSize) addPoints(pOdd[i], pOdd[i - 1], p2, sc) for (i in 1 until pTableSize) addPoints(pOddJac[i], pOddJac[i - 1], p2, sc)
// λ(P) table: (β·X, Y, Z) in Jacobian — endomorphism preserves projective coords // λ(P) table: (β·X, Y, Z) in Jacobian — endomorphism preserves projective coords
val pLamOdd = val pLamOddJac =
Array(pTableSize) { i -> Array(pTableSize) { i ->
val lp = MutablePoint() val lp = MutablePoint()
FieldP.mul(lp.x, pOdd[i].x, Glv.BETA, sc.w) FieldP.mul(lp.x, pOddJac[i].x, Glv.BETA, sc.w)
pOdd[i].y.copyInto(lp.y) pOddJac[i].y.copyInto(lp.y)
pOdd[i].z.copyInto(lp.z) pOddJac[i].z.copyInto(lp.z)
lp lp
} }
// Effective-affine: batch-convert P-side tables for cheaper mixed additions
val pOdd = batchToAffine(pOddJac, sc)
val pLamOdd = batchToAffine(pLamOddJac, sc)
// Find highest non-zero digit across all 4 streams // Find highest non-zero digit across all 4 streams
val allWnaf = arrayOf(wnafS1, wnafS2, wnafE1, wnafE2) val allWnaf = arrayOf(wnafS1, wnafS2, wnafE1, wnafE2)
var bits = allWnaf.maxOf { it.size } var bits = allWnaf.maxOf { it.size }
@@ -624,16 +678,15 @@ internal object ECPoint {
out.setInfinity() out.setInfinity()
val tmp = MutablePoint() val tmp = MutablePoint()
val negY = LongArray(4) val negY = LongArray(4)
val negJac = MutablePoint()
for (i in bits - 1 downTo 0) { for (i in bits - 1 downTo 0) {
doublePoint(out, out, sc) doublePoint(out, out, sc)
// Streams 1-2: G-side (affine tables, mixed addition) // Streams 1-2: G-side (affine tables, mixed addition)
addWnafMixed(out, tmp, negY, wnafS1, i, gOdd, sSplit.negK1, sc) addWnafMixed(out, tmp, negY, wnafS1, i, gOdd, sSplit.negK1, sc)
addWnafMixed(out, tmp, negY, wnafS2, i, gLam, sSplit.negK2, sc) addWnafMixed(out, tmp, negY, wnafS2, i, gLam, sSplit.negK2, sc)
// Streams 3-4: P-side (Jacobian tables, full addition) // Streams 3-4: P-side (affine tables via effective-affine, mixed addition)
addWnafJacobian(out, tmp, negJac, wnafE1, i, pOdd, eSplit.negK1, sc) addWnafMixed(out, tmp, negY, wnafE1, i, pOdd, eSplit.negK1, sc)
addWnafJacobian(out, tmp, negJac, wnafE2, i, pLamOdd, eSplit.negK2, sc) addWnafMixed(out, tmp, negY, wnafE2, i, pLamOdd, eSplit.negK2, sc)
} }
} }
@@ -693,6 +746,68 @@ internal object ECPoint {
out.copyFrom(tmp) out.copyFrom(tmp)
} }
// ==================== Batch Affine Conversion (Montgomery's Trick) ====================
/**
* Convert an array of Jacobian points to affine using Montgomery's batch inversion trick.
*
* Instead of n separate inversions (each ~250 multiplications), this computes all n inverses
* with a single inversion + 3(n-1) multiplications:
* 1. Build prefix products: c[i] = Z[0] · Z[1] · … · Z[i]
* 2. Invert the final product: inv = c[n-1]⁻¹
* 3. Recover individual inverses by peeling off factors from right to left:
* Z[i]⁻¹ = inv · c[i-1], then inv ← inv · Z[i]
* 4. Convert each point: x' = X · (Z⁻¹)², y' = Y · (Z⁻¹)³
*
* Cost: 1 inversion + 3(n-1) muls + 2n muls (for x,y conversion) = 1 inv + (5n-3) muls
* vs n inversions ≈ 250n muls.
*/
private fun batchToAffine(
points: Array<MutablePoint>,
s: PointScratch,
): Array<AffinePoint> {
val n = points.size
if (n == 0) return emptyArray()
val w = s.w
// Prefix products of Z coordinates
val cumZ = Array(n) { LongArray(4) }
U256.copyInto(cumZ[0], points[0].z)
for (i in 1 until n) {
FieldP.mul(cumZ[i], cumZ[i - 1], points[i].z, w)
}
// Invert the total product
val inv = LongArray(4)
FieldP.inv(inv, cumZ[n - 1])
// Recover individual Z inverses and convert to affine
val result = Array(n) { AffinePoint() }
val zInv = LongArray(4)
val zInv2 = LongArray(4)
val zInv3 = LongArray(4)
for (i in n - 1 downTo 1) {
// zInv = inv * cumZ[i-1] gives Z[i]^{-1}
FieldP.mul(zInv, inv, cumZ[i - 1], w)
// Update inv: inv = inv * Z[i] gives product-inverse without Z[i]
FieldP.mul(inv, inv, points[i].z, w)
// Convert to affine
FieldP.sqr(zInv2, zInv, w)
FieldP.mul(zInv3, zInv2, zInv, w)
FieldP.mul(result[i].x, points[i].x, zInv2, w)
FieldP.mul(result[i].y, points[i].y, zInv3, w)
}
// i == 0: inv now holds Z[0]^{-1}
FieldP.sqr(zInv2, inv, w)
FieldP.mul(zInv3, zInv2, inv, w)
FieldP.mul(result[0].x, points[0].x, zInv2, w)
FieldP.mul(result[0].y, points[0].y, zInv3, w)
return result
}
// ==================== Coordinate Conversion ==================== // ==================== Coordinate Conversion ====================
/** /**