mirror of
https://github.com/mruby/mruby
synced 2026-06-08 16:11:16 +00:00
mruby-bigint: optimize multiplication for all-ones numbers
Add fast path for multiplying numbers of form 2^n - 1 (all bits set). Uses algebraic identities: - (2^n - 1) * (2^m - 1) = 2^(n+m) - 2^n - 2^m + 1 - (2^n - 1) * y = (y << n) - y These are O(n) operations instead of O(n^1.585) for Karatsuba. Fuzzing test cases using all-ones patterns now complete in 0.01s instead of 13+ seconds. Also raises KARATSUBA_THRESHOLD from 8 to 32 for ~32% speedup on general large number multiplication. Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -77,6 +77,12 @@ pool_restore(mpz_ctx_t *ctx, size_t state)
|
||||
}
|
||||
}
|
||||
|
||||
/* Forward declarations */
|
||||
static void mpz_mul_2exp(mpz_ctx_t *ctx, mpz_t *z, mpz_t *x, mrb_int e);
|
||||
static void mpz_sub(mpz_ctx_t *ctx, mpz_t *z, mpz_t *x, mpz_t *y);
|
||||
static void mpz_add_int(mpz_ctx_t *ctx, mpz_t *x, mrb_int y);
|
||||
static void mpz_set_int(mpz_ctx_t *ctx, mpz_t *y, mrb_int v);
|
||||
|
||||
static mp_limb*
|
||||
pool_alloc(mpz_pool_t *pool, size_t limbs)
|
||||
{
|
||||
@@ -1109,6 +1115,95 @@ mpz_mul_karatsuba(mpz_ctx_t *ctx, mp_limb *result,
|
||||
limb_add_at(result, result_len, z2, z2_len, 2 * half);
|
||||
}
|
||||
|
||||
/*
|
||||
* Check if mpz is "all ones" pattern (2^n - 1).
|
||||
* For such a number:
|
||||
* - All limbs except possibly the top one equal DIG_MASK
|
||||
* - The top limb equals (1 << k) - 1 for some k in 1..DIG_SIZE
|
||||
* Returns the bit count n if all-ones, 0 otherwise.
|
||||
*/
|
||||
static size_t
|
||||
mpz_all_ones_p(mpz_t *x)
|
||||
{
|
||||
if (x->sn <= 0 || x->sz == 0) return 0;
|
||||
|
||||
/* Check all but top limb */
|
||||
for (size_t i = 0; i + 1 < x->sz; i++) {
|
||||
if (x->p[i] != DIG_MASK) return 0;
|
||||
}
|
||||
|
||||
/* Check top limb: must be (1 << k) - 1 for some k */
|
||||
mp_limb top = x->p[x->sz - 1];
|
||||
if (top == 0) return 0;
|
||||
|
||||
/* Check if top is all-ones pattern: (top & (top + 1)) == 0 */
|
||||
if ((top & (top + 1)) != 0) return 0;
|
||||
|
||||
/* Count bits in top limb */
|
||||
size_t top_bits = 0;
|
||||
while (top) { top_bits++; top >>= 1; }
|
||||
|
||||
return (x->sz - 1) * DIG_SIZE + top_bits;
|
||||
}
|
||||
|
||||
/*
|
||||
* Multiply two "all ones" numbers using algebraic identity:
|
||||
* (2^n - 1) * (2^m - 1) = 2^(n+m) - 2^n - 2^m + 1
|
||||
*
|
||||
* For squaring (n == m):
|
||||
* (2^n - 1)^2 = 2^(2n) - 2^(n+1) + 1
|
||||
*
|
||||
* This is O(n) instead of O(n^1.585) for Karatsuba.
|
||||
*/
|
||||
static void
|
||||
mpz_mul_all_ones(mpz_ctx_t *ctx, mpz_t *w, size_t n, size_t m)
|
||||
{
|
||||
mpz_t a;
|
||||
|
||||
if (n == m) {
|
||||
/* Squaring: (2^n - 1)^2 = 2^(2n) - 2^(n+1) + 1 */
|
||||
/* Start with 2^(2n) */
|
||||
mpz_init(ctx, &a);
|
||||
mpz_set_int(ctx, &a, 1);
|
||||
mpz_mul_2exp(ctx, w, &a, 2*n);
|
||||
|
||||
/* Subtract 2^(n+1) */
|
||||
mpz_set_int(ctx, &a, 1);
|
||||
mpz_mul_2exp(ctx, &a, &a, n+1);
|
||||
mpz_sub(ctx, w, w, &a);
|
||||
|
||||
/* Add 1 */
|
||||
mpz_add_int(ctx, w, 1);
|
||||
mpz_clear(ctx, &a);
|
||||
}
|
||||
else {
|
||||
/* General: (2^n - 1) * (2^m - 1) = 2^(n+m) - 2^n - 2^m + 1 */
|
||||
mpz_t b;
|
||||
mpz_init(ctx, &a);
|
||||
mpz_init(ctx, &b);
|
||||
|
||||
/* Start with 2^(n+m) */
|
||||
mpz_set_int(ctx, &a, 1);
|
||||
mpz_mul_2exp(ctx, w, &a, n+m);
|
||||
|
||||
/* Subtract 2^n */
|
||||
mpz_set_int(ctx, &a, 1);
|
||||
mpz_mul_2exp(ctx, &a, &a, n);
|
||||
mpz_sub(ctx, w, w, &a);
|
||||
|
||||
/* Subtract 2^m */
|
||||
mpz_set_int(ctx, &b, 1);
|
||||
mpz_mul_2exp(ctx, &b, &b, m);
|
||||
mpz_sub(ctx, w, w, &b);
|
||||
|
||||
/* Add 1 */
|
||||
mpz_add_int(ctx, w, 1);
|
||||
|
||||
mpz_clear(ctx, &a);
|
||||
mpz_clear(ctx, &b);
|
||||
}
|
||||
}
|
||||
|
||||
/* w = u * v */
|
||||
static void
|
||||
mpz_mul(mpz_ctx_t *ctx, mpz_t *ww, mpz_t *u, mpz_t *v)
|
||||
@@ -1118,6 +1213,32 @@ mpz_mul(mpz_ctx_t *ctx, mpz_t *ww, mpz_t *u, mpz_t *v)
|
||||
return;
|
||||
}
|
||||
|
||||
/* Fast path for "all ones" numbers (2^n - 1) */
|
||||
size_t u_ones = mpz_all_ones_p(u);
|
||||
size_t v_ones = mpz_all_ones_p(v);
|
||||
if (u_ones && v_ones) {
|
||||
mpz_mul_all_ones(ctx, ww, u_ones, v_ones);
|
||||
return;
|
||||
}
|
||||
|
||||
/* Fast path: (2^n - 1) * y = (y << n) - y */
|
||||
if (u_ones && u_ones >= KARATSUBA_THRESHOLD * DIG_SIZE) {
|
||||
mpz_t shifted;
|
||||
mpz_init(ctx, &shifted);
|
||||
mpz_mul_2exp(ctx, &shifted, v, u_ones);
|
||||
mpz_sub(ctx, ww, &shifted, v);
|
||||
mpz_clear(ctx, &shifted);
|
||||
return;
|
||||
}
|
||||
if (v_ones && v_ones >= KARATSUBA_THRESHOLD * DIG_SIZE) {
|
||||
mpz_t shifted;
|
||||
mpz_init(ctx, &shifted);
|
||||
mpz_mul_2exp(ctx, &shifted, u, v_ones);
|
||||
mpz_sub(ctx, ww, &shifted, u);
|
||||
mpz_clear(ctx, &shifted);
|
||||
return;
|
||||
}
|
||||
|
||||
if (!should_use_karatsuba(u->sz, v->sz)) {
|
||||
mpz_mul_basic(ctx, ww, u, v);
|
||||
return;
|
||||
|
||||
Reference in New Issue
Block a user