bend form B — cgbn-batch-worker covers all 9 dispatcher ops on 3090
Lands the remaining 8 ops from plans/form-B-bignum-cgbn.md §2:
0x01 mod-add cgbn_add + carry-or-ge-modulus subtract
0x02 mod-sub cgbn_sub + borrow conditional add
0x04 mod-sqr cgbn_sqr_wide + cgbn_rem_wide
0x05 mod-inv cgbn_modular_inverse (binary GCD)
0x06 mod-exp cgbn_modular_power (binary ladder)
0x07 mod-reduce cgbn_rem standalone
0x08 add-no-mod cgbn_add, truncated 256-bit
0x09 mul-no-mod cgbn_mul_wide, full 512-bit output (low|high)
process_one_bin now classifies op_id into three families (binary-mod /
unary-mod / no-mod), validates wire size per family, and carves modulus
/ a / b pointers accordingly. Output buffer width is 2x for 0x09 only.
test_cgbn_known_answers.py extended: one driver per op, gmpy2 reference
(with pure-Python fallback for invert/powmod), validated byte-identical
across n in {32, 1k, 10k, 100k}. ALL PASS on 3090-ai.foxhop.net.
Measured kernel throughput at n=100k (single 3090, median of 3):
mod-add 0.19 ms 526 Mops/s
mod-sub 0.19 ms 526 Mops/s
mod-mul 0.24 ms 417 Mops/s
mod-sqr 0.24 ms 417 Mops/s
mod-inv 2.39 ms 42 Mops/s
mod-exp 1.02 ms 98 Mops/s (16-bit exponents)
mod-reduce 0.20 ms 500 Mops/s
add-no-mod 0.19 ms 526 Mops/s
mul-no-mod 0.19 ms 526 Mops/s
mod-inv at 42 Mops/s tracks plan §8 projection (50-100 Mops/s on 3090
via CGBN's binary GCD) on the low end — Bernstein-Yang batched inverse
(form E) remains the upgrade path. mod-exp 98 Mops/s is for short
exponents only; full 256-bit ladder will drop ~16x per plan §8.
This commit is contained in:
parent
3ab4044805
commit
21bd26a4d3
2 changed files with 598 additions and 105 deletions
|
|
@ -4,26 +4,27 @@
|
|||
* Plan: examples/cuda-fanout/plans/form-B-bignum-cgbn.md
|
||||
* Catalog: CATALOG.md form B (cuda-bignum-cgbn).
|
||||
*
|
||||
* Day-1 baseline: daemon protocol mirrors shake256-fanout.cu (--daemon,
|
||||
* stdin process-bin <in> <out>, stdout done/error). Op 0x03 mod-mul at
|
||||
* 256-bit width only. Remaining ops (0x01 mod-add, 0x02 mod-sub,
|
||||
* 0x04 mod-sqr, 0x05 mod-inv, 0x06 mod-exp, 0x07 mod-reduce, 0x08
|
||||
* add-no-mod, 0x09 mul-no-mod) land per-op as we measure each.
|
||||
* Op coverage: 0x01 mod-add, 0x02 mod-sub, 0x03 mod-mul, 0x04 mod-sqr,
|
||||
* 0x05 mod-inv, 0x06 mod-exp, 0x07 mod-reduce,
|
||||
* 0x08 add-no-mod, 0x09 mul-no-mod.
|
||||
* Bitwidth: 256 only (day-1 freeze).
|
||||
*
|
||||
* Wire (BCGB input):
|
||||
* "BCGB" (4 B magic)
|
||||
* u32 op_id (0x03 for now)
|
||||
* u32 op_id
|
||||
* u32 bitwidth (256 for now)
|
||||
* u32 n_instances
|
||||
* u8[bitwidth/8] modulus
|
||||
* u8[bitwidth/8] modulus (OMITTED for 0x08, 0x09)
|
||||
* u8[n_instances * bitwidth/8] operand_a
|
||||
* u8[n_instances * bitwidth/8] operand_b (omitted for unary ops)
|
||||
* u8[n_instances * bitwidth/8] operand_b (OMITTED for unary 0x04, 0x05, 0x07)
|
||||
*
|
||||
* Wire (BCGR output):
|
||||
* "BCGR" (4 B magic)
|
||||
* u32 status (0 = ok)
|
||||
* u32 n_instances
|
||||
* u8[n_instances * bitwidth/8] result
|
||||
* u8[n_instances * width_out/8] result
|
||||
* width_out = bitwidth for 0x01..0x08
|
||||
* width_out = 2 * bitwidth for 0x09 mul-no-mod
|
||||
*
|
||||
* Build:
|
||||
* make cgbn-batch-worker CGBN_INC=/path/to/CGBN/include
|
||||
|
|
@ -61,6 +62,7 @@
|
|||
typedef cgbn_context_t<TPI> context_t;
|
||||
typedef cgbn_env_t<context_t, BITS_256> env_256_t;
|
||||
typedef cgbn_mem_t<BITS_256> mem_256_t;
|
||||
typedef cgbn_mem_t<2 * BITS_256> mem_512_t;
|
||||
|
||||
/* Op codes — must match plan §2 dispatcher table. */
|
||||
enum {
|
||||
|
|
@ -75,7 +77,64 @@ enum {
|
|||
OP_MUL_NO_MOD = 0x09,
|
||||
};
|
||||
|
||||
/* ── Kernel: mod-mul 256-bit, batched ──────────────────────────────── */
|
||||
/* ── Kernels: one __global__ per op, 256-bit, batched ──────────────── */
|
||||
|
||||
__global__ void mod_add_256_kernel(
|
||||
mem_256_t *result,
|
||||
mem_256_t *a,
|
||||
mem_256_t *b,
|
||||
mem_256_t *modulus,
|
||||
uint32_t n_instances)
|
||||
{
|
||||
uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
uint32_t instance = tid / TPI;
|
||||
if (instance >= n_instances) return;
|
||||
|
||||
context_t ctx(cgbn_no_checks);
|
||||
env_256_t env(ctx);
|
||||
env_256_t::cgbn_t bn_a, bn_b, bn_m, bn_r;
|
||||
|
||||
cgbn_load(env, bn_a, &a[instance]);
|
||||
cgbn_load(env, bn_b, &b[instance]);
|
||||
cgbn_load(env, bn_m, modulus);
|
||||
|
||||
/* r = a + b; if carry-out OR r >= m, r -= m. With operands in [0, m)
|
||||
* and m at most 2^256-1, sum fits in 257 bits. cgbn_add returns the
|
||||
* 257th carry bit; we subtract m when carry==1 OR r >= m. */
|
||||
int32_t carry = cgbn_add(env, bn_r, bn_a, bn_b);
|
||||
if (carry != 0 || cgbn_compare(env, bn_r, bn_m) >= 0) {
|
||||
cgbn_sub(env, bn_r, bn_r, bn_m);
|
||||
}
|
||||
cgbn_store(env, &result[instance], bn_r);
|
||||
}
|
||||
|
||||
__global__ void mod_sub_256_kernel(
|
||||
mem_256_t *result,
|
||||
mem_256_t *a,
|
||||
mem_256_t *b,
|
||||
mem_256_t *modulus,
|
||||
uint32_t n_instances)
|
||||
{
|
||||
uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
uint32_t instance = tid / TPI;
|
||||
if (instance >= n_instances) return;
|
||||
|
||||
context_t ctx(cgbn_no_checks);
|
||||
env_256_t env(ctx);
|
||||
env_256_t::cgbn_t bn_a, bn_b, bn_m, bn_r;
|
||||
|
||||
cgbn_load(env, bn_a, &a[instance]);
|
||||
cgbn_load(env, bn_b, &b[instance]);
|
||||
cgbn_load(env, bn_m, modulus);
|
||||
|
||||
/* r = a - b; if borrow (a < b), r += m. cgbn_sub returns nonzero
|
||||
* (typically -1) when underflow occurs. */
|
||||
int32_t borrow = cgbn_sub(env, bn_r, bn_a, bn_b);
|
||||
if (borrow != 0) {
|
||||
cgbn_add(env, bn_r, bn_r, bn_m);
|
||||
}
|
||||
cgbn_store(env, &result[instance], bn_r);
|
||||
}
|
||||
|
||||
__global__ void mod_mul_256_kernel(
|
||||
mem_256_t *result,
|
||||
|
|
@ -96,13 +155,7 @@ __global__ void mod_mul_256_kernel(
|
|||
cgbn_load(env, bn_b, &b[instance]);
|
||||
cgbn_load(env, bn_m, modulus);
|
||||
|
||||
/* a * b mod m via wide-mul + reduce. CGBN supplies cgbn_mont_mul
|
||||
* for Montgomery form; for a-bit-honest schoolbook (with reduce)
|
||||
* we use cgbn_mul_wide + cgbn_rem. At our 256-bit width the
|
||||
* cost difference is small & we avoid the precomputed-r2-modulus
|
||||
* dance Montgomery requires. Promote to Montgomery once we
|
||||
* benchmark.
|
||||
*/
|
||||
/* a * b mod m via wide-mul + reduce. */
|
||||
typename env_256_t::cgbn_wide_t bn_wide;
|
||||
cgbn_mul_wide(env, bn_wide, bn_a, bn_b);
|
||||
cgbn_rem_wide(env, bn_r, bn_wide, bn_m);
|
||||
|
|
@ -110,6 +163,155 @@ __global__ void mod_mul_256_kernel(
|
|||
cgbn_store(env, &result[instance], bn_r);
|
||||
}
|
||||
|
||||
__global__ void mod_sqr_256_kernel(
|
||||
mem_256_t *result,
|
||||
mem_256_t *a,
|
||||
mem_256_t *modulus,
|
||||
uint32_t n_instances)
|
||||
{
|
||||
uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
uint32_t instance = tid / TPI;
|
||||
if (instance >= n_instances) return;
|
||||
|
||||
context_t ctx(cgbn_no_checks);
|
||||
env_256_t env(ctx);
|
||||
env_256_t::cgbn_t bn_a, bn_m, bn_r;
|
||||
|
||||
cgbn_load(env, bn_a, &a[instance]);
|
||||
cgbn_load(env, bn_m, modulus);
|
||||
|
||||
/* a^2 mod m via sqr_wide + rem_wide (avoids the second cgbn_t load). */
|
||||
typename env_256_t::cgbn_wide_t bn_wide;
|
||||
cgbn_sqr_wide(env, bn_wide, bn_a);
|
||||
cgbn_rem_wide(env, bn_r, bn_wide, bn_m);
|
||||
|
||||
cgbn_store(env, &result[instance], bn_r);
|
||||
}
|
||||
|
||||
__global__ void mod_inv_256_kernel(
|
||||
mem_256_t *result,
|
||||
mem_256_t *a,
|
||||
mem_256_t *modulus,
|
||||
uint32_t n_instances)
|
||||
{
|
||||
uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
uint32_t instance = tid / TPI;
|
||||
if (instance >= n_instances) return;
|
||||
|
||||
context_t ctx(cgbn_no_checks);
|
||||
env_256_t env(ctx);
|
||||
env_256_t::cgbn_t bn_a, bn_m, bn_r;
|
||||
|
||||
cgbn_load(env, bn_a, &a[instance]);
|
||||
cgbn_load(env, bn_m, modulus);
|
||||
|
||||
/* cgbn_modular_inverse returns true on success; if gcd(a,m) != 1 we
|
||||
* write zero — caller must avoid feeding non-invertible inputs. */
|
||||
bool ok = cgbn_modular_inverse(env, bn_r, bn_a, bn_m);
|
||||
if (!ok) {
|
||||
cgbn_set_ui32(env, bn_r, 0);
|
||||
}
|
||||
cgbn_store(env, &result[instance], bn_r);
|
||||
}
|
||||
|
||||
__global__ void mod_exp_256_kernel(
|
||||
mem_256_t *result,
|
||||
mem_256_t *a,
|
||||
mem_256_t *b,
|
||||
mem_256_t *modulus,
|
||||
uint32_t n_instances)
|
||||
{
|
||||
uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
uint32_t instance = tid / TPI;
|
||||
if (instance >= n_instances) return;
|
||||
|
||||
context_t ctx(cgbn_no_checks);
|
||||
env_256_t env(ctx);
|
||||
env_256_t::cgbn_t bn_a, bn_b, bn_m, bn_r;
|
||||
|
||||
cgbn_load(env, bn_a, &a[instance]);
|
||||
cgbn_load(env, bn_b, &b[instance]);
|
||||
cgbn_load(env, bn_m, modulus);
|
||||
|
||||
/* a^b mod m via binary ladder (CGBN's built-in). */
|
||||
cgbn_modular_power(env, bn_r, bn_a, bn_b, bn_m);
|
||||
|
||||
cgbn_store(env, &result[instance], bn_r);
|
||||
}
|
||||
|
||||
__global__ void mod_reduce_256_kernel(
|
||||
mem_256_t *result,
|
||||
mem_256_t *a,
|
||||
mem_256_t *modulus,
|
||||
uint32_t n_instances)
|
||||
{
|
||||
uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
uint32_t instance = tid / TPI;
|
||||
if (instance >= n_instances) return;
|
||||
|
||||
context_t ctx(cgbn_no_checks);
|
||||
env_256_t env(ctx);
|
||||
env_256_t::cgbn_t bn_a, bn_m, bn_r;
|
||||
|
||||
cgbn_load(env, bn_a, &a[instance]);
|
||||
cgbn_load(env, bn_m, modulus);
|
||||
|
||||
/* a is already <= 256 bits, so cgbn_rem suffices (no wide path). */
|
||||
cgbn_rem(env, bn_r, bn_a, bn_m);
|
||||
|
||||
cgbn_store(env, &result[instance], bn_r);
|
||||
}
|
||||
|
||||
__global__ void add_no_mod_256_kernel(
|
||||
mem_256_t *result,
|
||||
mem_256_t *a,
|
||||
mem_256_t *b,
|
||||
uint32_t n_instances)
|
||||
{
|
||||
uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
uint32_t instance = tid / TPI;
|
||||
if (instance >= n_instances) return;
|
||||
|
||||
context_t ctx(cgbn_no_checks);
|
||||
env_256_t env(ctx);
|
||||
env_256_t::cgbn_t bn_a, bn_b, bn_r;
|
||||
|
||||
cgbn_load(env, bn_a, &a[instance]);
|
||||
cgbn_load(env, bn_b, &b[instance]);
|
||||
|
||||
/* Straight 256-bit add, truncated (carry-out discarded). */
|
||||
cgbn_add(env, bn_r, bn_a, bn_b);
|
||||
|
||||
cgbn_store(env, &result[instance], bn_r);
|
||||
}
|
||||
|
||||
__global__ void mul_no_mod_256_kernel(
|
||||
mem_512_t *result,
|
||||
mem_256_t *a,
|
||||
mem_256_t *b,
|
||||
uint32_t n_instances)
|
||||
{
|
||||
uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
uint32_t instance = tid / TPI;
|
||||
if (instance >= n_instances) return;
|
||||
|
||||
context_t ctx(cgbn_no_checks);
|
||||
env_256_t env(ctx);
|
||||
env_256_t::cgbn_t bn_a, bn_b;
|
||||
|
||||
cgbn_load(env, bn_a, &a[instance]);
|
||||
cgbn_load(env, bn_b, &b[instance]);
|
||||
|
||||
/* Full 512-bit product. Store as low half then high half so the byte
|
||||
* stream is little-endian across the entire 64-byte product. */
|
||||
typename env_256_t::cgbn_wide_t bn_wide;
|
||||
cgbn_mul_wide(env, bn_wide, bn_a, bn_b);
|
||||
|
||||
mem_256_t *out_pair = (mem_256_t *)&result[instance];
|
||||
cgbn_store(env, &out_pair[0], bn_wide._low);
|
||||
cgbn_store(env, &out_pair[1], bn_wide._high);
|
||||
}
|
||||
|
||||
/* ── Host: process_one_bin ─────────────────────────────────────────── */
|
||||
|
||||
static int read_file_to_buffer(const char *path, uint8_t **out_buf, size_t *out_len) {
|
||||
|
|
@ -147,6 +349,43 @@ static void wr32_le(uint8_t *p, uint32_t v) {
|
|||
p[2] = (v >> 16) & 0xff; p[3] = (v >> 24) & 0xff;
|
||||
}
|
||||
|
||||
/* Op-shape classifier. Three families:
|
||||
* FAMILY_BIN_MOD binary with modulus : 0x01, 0x02, 0x03, 0x06
|
||||
* input = 16 + bpi + 2*n*bpi
|
||||
* output = 12 + n*bpi
|
||||
* FAMILY_UN_MOD unary with modulus : 0x04, 0x05, 0x07
|
||||
* input = 16 + bpi + n*bpi
|
||||
* output = 12 + n*bpi
|
||||
* FAMILY_NO_MOD binary no modulus : 0x08, 0x09
|
||||
* input = 16 + 2*n*bpi
|
||||
* output = 12 + n*bpi (0x08) or 12 + 2*n*bpi (0x09)
|
||||
*/
|
||||
enum op_family {
|
||||
FAMILY_BIN_MOD,
|
||||
FAMILY_UN_MOD,
|
||||
FAMILY_NO_MOD,
|
||||
FAMILY_UNKNOWN,
|
||||
};
|
||||
|
||||
static enum op_family classify_op(uint32_t op_id) {
|
||||
switch (op_id) {
|
||||
case OP_MOD_ADD:
|
||||
case OP_MOD_SUB:
|
||||
case OP_MOD_MUL:
|
||||
case OP_MOD_EXP:
|
||||
return FAMILY_BIN_MOD;
|
||||
case OP_MOD_SQR:
|
||||
case OP_MOD_INV:
|
||||
case OP_MOD_REDUCE:
|
||||
return FAMILY_UN_MOD;
|
||||
case OP_ADD_NO_MOD:
|
||||
case OP_MUL_NO_MOD:
|
||||
return FAMILY_NO_MOD;
|
||||
default:
|
||||
return FAMILY_UNKNOWN;
|
||||
}
|
||||
}
|
||||
|
||||
static int process_one_bin(const char *in_path, const char *out_path) {
|
||||
uint8_t *in_buf = NULL;
|
||||
size_t in_len = 0;
|
||||
|
|
@ -161,69 +400,144 @@ static int process_one_bin(const char *in_path, const char *out_path) {
|
|||
uint32_t n = rd32_le(in_buf + 12);
|
||||
|
||||
if (bitwidth != BITS_256) {
|
||||
fprintf(stderr, "unsupported bitwidth %u (day-1 ships 256 only)\n",
|
||||
fprintf(stderr, "unsupported bitwidth %u (worker ships 256 only)\n",
|
||||
bitwidth);
|
||||
free(in_buf); return -1;
|
||||
}
|
||||
if (op_id != OP_MOD_MUL) {
|
||||
fprintf(stderr, "unsupported op_id 0x%02x (day-1 ships 0x03 mod-mul only)\n",
|
||||
op_id);
|
||||
|
||||
enum op_family fam = classify_op(op_id);
|
||||
if (fam == FAMILY_UNKNOWN) {
|
||||
fprintf(stderr, "unsupported op_id 0x%02x\n", op_id);
|
||||
free(in_buf); return -1;
|
||||
}
|
||||
|
||||
size_t bytes_per_inst = bitwidth / 8;
|
||||
size_t expected = 16 + bytes_per_inst + 2 * (size_t)n * bytes_per_inst;
|
||||
size_t bpi = bitwidth / 8;
|
||||
size_t expected = 0;
|
||||
switch (fam) {
|
||||
case FAMILY_BIN_MOD: expected = 16 + bpi + 2 * (size_t)n * bpi; break;
|
||||
case FAMILY_UN_MOD: expected = 16 + bpi + (size_t)n * bpi; break;
|
||||
case FAMILY_NO_MOD: expected = 16 + 2 * (size_t)n * bpi; break;
|
||||
default: break;
|
||||
}
|
||||
if (in_len != expected) {
|
||||
fprintf(stderr, "BCGB payload size %zu != expected %zu (n=%u)\n",
|
||||
in_len, expected, n);
|
||||
fprintf(stderr, "BCGB payload size %zu != expected %zu (op=0x%02x n=%u)\n",
|
||||
in_len, expected, op_id, n);
|
||||
free(in_buf); return -1;
|
||||
}
|
||||
|
||||
const uint8_t *modulus_bytes = in_buf + 16;
|
||||
const uint8_t *a_bytes = modulus_bytes + bytes_per_inst;
|
||||
const uint8_t *b_bytes = a_bytes + (size_t)n * bytes_per_inst;
|
||||
/* Per-family pointer carving. */
|
||||
const uint8_t *modulus_bytes = NULL;
|
||||
const uint8_t *a_bytes = NULL;
|
||||
const uint8_t *b_bytes = NULL;
|
||||
switch (fam) {
|
||||
case FAMILY_BIN_MOD:
|
||||
modulus_bytes = in_buf + 16;
|
||||
a_bytes = modulus_bytes + bpi;
|
||||
b_bytes = a_bytes + (size_t)n * bpi;
|
||||
break;
|
||||
case FAMILY_UN_MOD:
|
||||
modulus_bytes = in_buf + 16;
|
||||
a_bytes = modulus_bytes + bpi;
|
||||
break;
|
||||
case FAMILY_NO_MOD:
|
||||
a_bytes = in_buf + 16;
|
||||
b_bytes = a_bytes + (size_t)n * bpi;
|
||||
break;
|
||||
default: break;
|
||||
}
|
||||
|
||||
/* allocate device buffers */
|
||||
/* Allocate device buffers. */
|
||||
mem_256_t *d_modulus = NULL, *d_a = NULL, *d_b = NULL, *d_r = NULL;
|
||||
CUDA_CHECK(cudaMalloc(&d_modulus, sizeof(mem_256_t)));
|
||||
mem_512_t *d_r_wide = NULL;
|
||||
|
||||
if (modulus_bytes) {
|
||||
CUDA_CHECK(cudaMalloc(&d_modulus, sizeof(mem_256_t)));
|
||||
CUDA_CHECK(cudaMemcpy(d_modulus, modulus_bytes, bpi,
|
||||
cudaMemcpyHostToDevice));
|
||||
}
|
||||
CUDA_CHECK(cudaMalloc(&d_a, (size_t)n * sizeof(mem_256_t)));
|
||||
CUDA_CHECK(cudaMalloc(&d_b, (size_t)n * sizeof(mem_256_t)));
|
||||
CUDA_CHECK(cudaMalloc(&d_r, (size_t)n * sizeof(mem_256_t)));
|
||||
CUDA_CHECK(cudaMemcpy(d_a, a_bytes, (size_t)n * bpi,
|
||||
cudaMemcpyHostToDevice));
|
||||
if (b_bytes) {
|
||||
CUDA_CHECK(cudaMalloc(&d_b, (size_t)n * sizeof(mem_256_t)));
|
||||
CUDA_CHECK(cudaMemcpy(d_b, b_bytes, (size_t)n * bpi,
|
||||
cudaMemcpyHostToDevice));
|
||||
}
|
||||
if (op_id == OP_MUL_NO_MOD) {
|
||||
CUDA_CHECK(cudaMalloc(&d_r_wide, (size_t)n * sizeof(mem_512_t)));
|
||||
} else {
|
||||
CUDA_CHECK(cudaMalloc(&d_r, (size_t)n * sizeof(mem_256_t)));
|
||||
}
|
||||
|
||||
CUDA_CHECK(cudaMemcpy(d_modulus, modulus_bytes, bytes_per_inst,
|
||||
cudaMemcpyHostToDevice));
|
||||
CUDA_CHECK(cudaMemcpy(d_a, a_bytes, (size_t)n * bytes_per_inst,
|
||||
cudaMemcpyHostToDevice));
|
||||
CUDA_CHECK(cudaMemcpy(d_b, b_bytes, (size_t)n * bytes_per_inst,
|
||||
cudaMemcpyHostToDevice));
|
||||
|
||||
/* launch: TPI threads per instance, 128 instances per block. */
|
||||
/* Launch geometry: TPI threads per instance, 128 instances per block. */
|
||||
uint32_t threads_per_block = 128 * TPI;
|
||||
uint32_t total_threads = n * TPI;
|
||||
uint32_t blocks = (total_threads + threads_per_block - 1) / threads_per_block;
|
||||
|
||||
struct timespec t0, t1;
|
||||
clock_gettime(CLOCK_MONOTONIC, &t0);
|
||||
mod_mul_256_kernel<<<blocks, threads_per_block>>>(d_r, d_a, d_b, d_modulus, n);
|
||||
switch (op_id) {
|
||||
case OP_MOD_ADD:
|
||||
mod_add_256_kernel<<<blocks, threads_per_block>>>(d_r, d_a, d_b, d_modulus, n);
|
||||
break;
|
||||
case OP_MOD_SUB:
|
||||
mod_sub_256_kernel<<<blocks, threads_per_block>>>(d_r, d_a, d_b, d_modulus, n);
|
||||
break;
|
||||
case OP_MOD_MUL:
|
||||
mod_mul_256_kernel<<<blocks, threads_per_block>>>(d_r, d_a, d_b, d_modulus, n);
|
||||
break;
|
||||
case OP_MOD_SQR:
|
||||
mod_sqr_256_kernel<<<blocks, threads_per_block>>>(d_r, d_a, d_modulus, n);
|
||||
break;
|
||||
case OP_MOD_INV:
|
||||
mod_inv_256_kernel<<<blocks, threads_per_block>>>(d_r, d_a, d_modulus, n);
|
||||
break;
|
||||
case OP_MOD_EXP:
|
||||
mod_exp_256_kernel<<<blocks, threads_per_block>>>(d_r, d_a, d_b, d_modulus, n);
|
||||
break;
|
||||
case OP_MOD_REDUCE:
|
||||
mod_reduce_256_kernel<<<blocks, threads_per_block>>>(d_r, d_a, d_modulus, n);
|
||||
break;
|
||||
case OP_ADD_NO_MOD:
|
||||
add_no_mod_256_kernel<<<blocks, threads_per_block>>>(d_r, d_a, d_b, n);
|
||||
break;
|
||||
case OP_MUL_NO_MOD:
|
||||
mul_no_mod_256_kernel<<<blocks, threads_per_block>>>(d_r_wide, d_a, d_b, n);
|
||||
break;
|
||||
}
|
||||
CUDA_CHECK(cudaDeviceSynchronize());
|
||||
clock_gettime(CLOCK_MONOTONIC, &t1);
|
||||
double kernel_ms = (t1.tv_sec - t0.tv_sec) * 1e3
|
||||
+ (t1.tv_nsec - t0.tv_nsec) / 1e6;
|
||||
|
||||
/* assemble BCGR output */
|
||||
size_t out_size = 12 + (size_t)n * bytes_per_inst;
|
||||
/* Output sizing. */
|
||||
size_t out_bytes_per_inst = (op_id == OP_MUL_NO_MOD) ? 2 * bpi : bpi;
|
||||
size_t out_size = 12 + (size_t)n * out_bytes_per_inst;
|
||||
uint8_t *out_buf = (uint8_t*)malloc(out_size);
|
||||
if (!out_buf) {
|
||||
cudaFree(d_modulus); cudaFree(d_a); cudaFree(d_b); cudaFree(d_r);
|
||||
if (d_modulus) cudaFree(d_modulus);
|
||||
if (d_a) cudaFree(d_a);
|
||||
if (d_b) cudaFree(d_b);
|
||||
if (d_r) cudaFree(d_r);
|
||||
if (d_r_wide) cudaFree(d_r_wide);
|
||||
free(in_buf); return -1;
|
||||
}
|
||||
memcpy(out_buf, "BCGR", 4);
|
||||
wr32_le(out_buf + 4, 0); /* status = ok */
|
||||
wr32_le(out_buf + 8, n);
|
||||
CUDA_CHECK(cudaMemcpy(out_buf + 12, d_r, (size_t)n * bytes_per_inst,
|
||||
cudaMemcpyDeviceToHost));
|
||||
if (op_id == OP_MUL_NO_MOD) {
|
||||
CUDA_CHECK(cudaMemcpy(out_buf + 12, d_r_wide, (size_t)n * out_bytes_per_inst,
|
||||
cudaMemcpyDeviceToHost));
|
||||
} else {
|
||||
CUDA_CHECK(cudaMemcpy(out_buf + 12, d_r, (size_t)n * out_bytes_per_inst,
|
||||
cudaMemcpyDeviceToHost));
|
||||
}
|
||||
|
||||
cudaFree(d_modulus); cudaFree(d_a); cudaFree(d_b); cudaFree(d_r);
|
||||
if (d_modulus) cudaFree(d_modulus);
|
||||
if (d_a) cudaFree(d_a);
|
||||
if (d_b) cudaFree(d_b);
|
||||
if (d_r) cudaFree(d_r);
|
||||
if (d_r_wide) cudaFree(d_r_wide);
|
||||
free(in_buf);
|
||||
|
||||
int wr = write_buffer_to_file(out_path, out_buf, out_size);
|
||||
|
|
|
|||
|
|
@ -1,17 +1,24 @@
|
|||
"""test_cgbn_known_answers.py — Day-1 validation for cgbn-batch-worker.
|
||||
"""test_cgbn_known_answers.py — validation for cgbn-batch-worker.
|
||||
|
||||
Pipes BCGB-format mod-mul requests at our CGBN worker via --binary
|
||||
mode, parses the BSHR response, and compares result bytes against a
|
||||
host bignum reference. Falls back to pure-Python pow(a*b, 1, m) if
|
||||
gmpy2 is unavailable (slower but no extra deps).
|
||||
Pipes BCGB-format ops at our CGBN worker via --binary mode, parses the
|
||||
BCGR response, and compares result bytes against a host bignum
|
||||
reference. Falls back to pure-Python if gmpy2 is unavailable.
|
||||
|
||||
Op coverage (Day 1): 0x03 mod-mul at 256-bit width only.
|
||||
Op coverage (later): 0x01 add, 0x02 sub, 0x04 sqr, 0x05 inv, 0x06 exp,
|
||||
0x07 reduce, 0x08 add-no-mod, 0x09 mul-no-mod.
|
||||
Op coverage (all 9 ops in plan §2 dispatcher table):
|
||||
0x01 mod-add a + b mod m
|
||||
0x02 mod-sub a - b mod m
|
||||
0x03 mod-mul a * b mod m
|
||||
0x04 mod-sqr a^2 mod m
|
||||
0x05 mod-inv a^-1 mod m
|
||||
0x06 mod-exp a^b mod m
|
||||
0x07 mod-reduce a mod m
|
||||
0x08 add-no-mod a + b (truncated to 256 bits)
|
||||
0x09 mul-no-mod a * b (full 512-bit product)
|
||||
|
||||
Usage:
|
||||
python3 test_cgbn_known_answers.py [worker-path] ; default: ./cgbn-batch-worker
|
||||
python3 test_cgbn_known_answers.py ./cgbn-batch-worker --n 1000
|
||||
python3 test_cgbn_known_answers.py ./cgbn-batch-worker --op mod-mul
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
|
|
@ -19,16 +26,39 @@ import random
|
|||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
|
||||
SECP256K1_P = 0xFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFEFFFFFC2F
|
||||
MASK_256 = (1 << 256) - 1
|
||||
|
||||
# Op codes — match cgbn-batch-worker.cu OP_* enum.
|
||||
OP_MOD_ADD = 0x01
|
||||
OP_MOD_SUB = 0x02
|
||||
OP_MOD_MUL = 0x03
|
||||
OP_MOD_SQR = 0x04
|
||||
OP_MOD_INV = 0x05
|
||||
OP_MOD_EXP = 0x06
|
||||
OP_MOD_REDUCE = 0x07
|
||||
OP_ADD_NO_MOD = 0x08
|
||||
OP_MUL_NO_MOD = 0x09
|
||||
|
||||
try:
|
||||
from gmpy2 import mpz, f_mod
|
||||
def host_mod_mul(a, b, m):
|
||||
return int(f_mod(mpz(a) * mpz(b), mpz(m)))
|
||||
from gmpy2 import mpz, f_mod, invert, powmod
|
||||
def h_mul(a, b, m): return int(f_mod(mpz(a) * mpz(b), mpz(m)))
|
||||
def h_add(a, b, m): return int(f_mod(mpz(a) + mpz(b), mpz(m)))
|
||||
def h_sub(a, b, m): return int(f_mod(mpz(a) - mpz(b), mpz(m)))
|
||||
def h_sqr(a, m): return int(f_mod(mpz(a) * mpz(a), mpz(m)))
|
||||
def h_inv(a, m): return int(invert(mpz(a), mpz(m)))
|
||||
def h_exp(a, b, m): return int(powmod(mpz(a), mpz(b), mpz(m)))
|
||||
def h_rem(a, m): return int(f_mod(mpz(a), mpz(m)))
|
||||
except ImportError:
|
||||
def host_mod_mul(a, b, m):
|
||||
return (a * b) % m
|
||||
def h_mul(a, b, m): return (a * b) % m
|
||||
def h_add(a, b, m): return (a + b) % m
|
||||
def h_sub(a, b, m): return (a - b) % m
|
||||
def h_sqr(a, m): return (a * a) % m
|
||||
def h_inv(a, m): return pow(a, -1, m)
|
||||
def h_exp(a, b, m): return pow(a, b, m)
|
||||
def h_rem(a, m): return a % m
|
||||
|
||||
|
||||
def int_to_le_bytes(x, nbytes=32):
|
||||
|
|
@ -39,92 +69,229 @@ def le_bytes_to_int(b):
|
|||
return int.from_bytes(b, "little")
|
||||
|
||||
|
||||
def build_bshk_mod_mul(modulus, a_list, b_list, bitwidth=256):
|
||||
def build_bcgb(op_id, modulus, a_list, b_list, bitwidth=256):
|
||||
"""Build BCGB payload. modulus=None for no-mod ops; b_list=None for unary."""
|
||||
bpi = bitwidth // 8
|
||||
n = len(a_list)
|
||||
assert len(b_list) == n
|
||||
if b_list is not None:
|
||||
assert len(b_list) == n
|
||||
parts = [
|
||||
b"BCGB",
|
||||
(0x03).to_bytes(4, "little"), # op_id = mod-mul
|
||||
op_id.to_bytes(4, "little"),
|
||||
bitwidth.to_bytes(4, "little"),
|
||||
n.to_bytes(4, "little"),
|
||||
int_to_le_bytes(modulus, bpi),
|
||||
]
|
||||
if modulus is not None:
|
||||
parts.append(int_to_le_bytes(modulus, bpi))
|
||||
for a in a_list:
|
||||
parts.append(int_to_le_bytes(a, bpi))
|
||||
for b in b_list:
|
||||
parts.append(int_to_le_bytes(b, bpi))
|
||||
if b_list is not None:
|
||||
for b in b_list:
|
||||
parts.append(int_to_le_bytes(b, bpi))
|
||||
return b"".join(parts)
|
||||
|
||||
|
||||
def parse_bshr(blob, bitwidth=256):
|
||||
def parse_bcgr(blob, n_expected, out_bytes_per_inst=32):
|
||||
if len(blob) < 12 or blob[:4] != b"BCGR":
|
||||
raise RuntimeError(f"bad BCGR magic / short: {blob[:40]!r}")
|
||||
status = int.from_bytes(blob[4:8], "little")
|
||||
n = int.from_bytes(blob[8:12], "little")
|
||||
if status != 0:
|
||||
raise RuntimeError(f"worker returned status {status}")
|
||||
bpi = bitwidth // 8
|
||||
expected = 12 + n * bpi
|
||||
if len(blob) != expected:
|
||||
raise RuntimeError(f"BCGR size {len(blob)} != expected {expected} for n={n}")
|
||||
if n != n_expected:
|
||||
raise RuntimeError(f"BCGR n {n} != expected {n_expected}")
|
||||
expected_len = 12 + n * out_bytes_per_inst
|
||||
if len(blob) != expected_len:
|
||||
raise RuntimeError(
|
||||
f"BCGR size {len(blob)} != expected {expected_len} for n={n} bpi={out_bytes_per_inst}"
|
||||
)
|
||||
out = []
|
||||
for i in range(n):
|
||||
out.append(le_bytes_to_int(blob[12 + i * bpi : 12 + (i + 1) * bpi]))
|
||||
out.append(le_bytes_to_int(
|
||||
blob[12 + i * out_bytes_per_inst : 12 + (i + 1) * out_bytes_per_inst]
|
||||
))
|
||||
return out
|
||||
|
||||
|
||||
def run_worker_binary(worker, in_blob):
|
||||
"""Spawn worker --binary in_path out_path; return out_blob."""
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=".bshk") as fin:
|
||||
"""Spawn worker --binary in_path out_path; return (out_blob, stderr_str)."""
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=".bcgb") as fin:
|
||||
fin.write(in_blob)
|
||||
in_path = fin.name
|
||||
out_path = in_path.replace(".bshk", ".bshr")
|
||||
out_path = in_path.replace(".bcgb", ".bcgr")
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[worker, "--binary", in_path, out_path],
|
||||
capture_output=True, timeout=60,
|
||||
capture_output=True, timeout=120,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(
|
||||
f"worker exited {result.returncode}: {result.stderr.decode()[:200]}"
|
||||
f"worker exited {result.returncode}: {result.stderr.decode()[:400]}"
|
||||
)
|
||||
with open(out_path, "rb") as fout:
|
||||
return fout.read()
|
||||
return fout.read(), result.stderr.decode()
|
||||
finally:
|
||||
for p in (in_path, out_path):
|
||||
try: os.unlink(p)
|
||||
except OSError: pass
|
||||
|
||||
|
||||
def test_mod_mul(worker, n, modulus, label):
|
||||
rng = random.Random(0xC0FFEE57C0DEC0DE ^ n)
|
||||
def report_pass_fail(label, n, fail, first_fail, modulus):
|
||||
if fail == 0:
|
||||
print(f" PASS {label:30s} n={n}")
|
||||
return True
|
||||
i, a, b, exp, got = first_fail
|
||||
print(f" FAIL {label:30s} n={n} mismatches={fail}/{n}")
|
||||
print(f" first @ i={i}")
|
||||
print(f" a = 0x{a:064x}")
|
||||
if b is not None:
|
||||
print(f" b = 0x{b:064x}")
|
||||
if modulus is not None:
|
||||
print(f" m = 0x{modulus:064x}")
|
||||
print(f" exp = 0x{exp:x}")
|
||||
print(f" got = 0x{got:x}")
|
||||
return False
|
||||
|
||||
|
||||
def test_binary_mod(worker, op_id, op_name, host_fn, n, modulus, label):
|
||||
rng = random.Random(0xC0FFEE57C0DEC0DE ^ n ^ op_id)
|
||||
a_list = [rng.randrange(0, modulus) for _ in range(n)]
|
||||
b_list = [rng.randrange(0, modulus) for _ in range(n)]
|
||||
in_blob = build_bshk_mod_mul(modulus, a_list, b_list)
|
||||
out_blob = run_worker_binary(worker, in_blob)
|
||||
gpu_results = parse_bshr(out_blob)
|
||||
fail = 0
|
||||
first_fail = None
|
||||
in_blob = build_bcgb(op_id, modulus, a_list, b_list)
|
||||
out_blob, stderr = run_worker_binary(worker, in_blob)
|
||||
gpu = parse_bcgr(out_blob, n, out_bytes_per_inst=32)
|
||||
fail = 0; first_fail = None
|
||||
for i in range(n):
|
||||
expected = host_mod_mul(a_list[i], b_list[i], modulus)
|
||||
if gpu_results[i] != expected:
|
||||
expected = host_fn(a_list[i], b_list[i], modulus)
|
||||
if gpu[i] != expected:
|
||||
fail += 1
|
||||
if first_fail is None:
|
||||
first_fail = (i, a_list[i], b_list[i], expected, gpu_results[i])
|
||||
if fail == 0:
|
||||
print(f" PASS {label:25s} n={n}")
|
||||
return True
|
||||
first_fail = (i, a_list[i], b_list[i], expected, gpu[i])
|
||||
return report_pass_fail(f"{op_name} {label}", n, fail, first_fail, modulus)
|
||||
|
||||
|
||||
def test_unary_mod(worker, op_id, op_name, host_fn, n, modulus, label,
|
||||
ensure_invertible=False):
|
||||
rng = random.Random(0xC0FFEE57C0DEC0DE ^ n ^ op_id)
|
||||
if ensure_invertible:
|
||||
# secp256k1_p and our other moduli are prime, so any nonzero a < m
|
||||
# is invertible. Still pick from [1, m) to be safe.
|
||||
a_list = [rng.randrange(1, modulus) for _ in range(n)]
|
||||
else:
|
||||
i, a, b, exp, got = first_fail
|
||||
print(f" FAIL {label:25s} n={n} mismatches={fail}/{n}")
|
||||
print(f" first @ i={i}")
|
||||
print(f" a = 0x{a:064x}")
|
||||
print(f" b = 0x{b:064x}")
|
||||
print(f" m = 0x{modulus:064x}")
|
||||
print(f" exp = 0x{exp:064x}")
|
||||
print(f" got = 0x{got:064x}")
|
||||
return False
|
||||
# For mod-reduce we also feed values that exceed m so the reduce
|
||||
# is non-trivial. We cap at 2^256-1 because the wire is 256-bit.
|
||||
if op_id == OP_MOD_REDUCE:
|
||||
a_list = [rng.randrange(0, MASK_256 + 1) for _ in range(n)]
|
||||
else:
|
||||
a_list = [rng.randrange(0, modulus) for _ in range(n)]
|
||||
in_blob = build_bcgb(op_id, modulus, a_list, None)
|
||||
out_blob, stderr = run_worker_binary(worker, in_blob)
|
||||
gpu = parse_bcgr(out_blob, n, out_bytes_per_inst=32)
|
||||
fail = 0; first_fail = None
|
||||
for i in range(n):
|
||||
expected = host_fn(a_list[i], modulus)
|
||||
if gpu[i] != expected:
|
||||
fail += 1
|
||||
if first_fail is None:
|
||||
first_fail = (i, a_list[i], None, expected, gpu[i])
|
||||
return report_pass_fail(f"{op_name} {label}", n, fail, first_fail, modulus)
|
||||
|
||||
|
||||
def test_mod_exp_small(worker, n, modulus, label):
|
||||
"""mod-exp: keep b in [0, 2^16) so the binary ladder finishes fast
|
||||
enough at n=10k without blocking the whole suite for minutes."""
|
||||
rng = random.Random(0xC0FFEE57C0DEC0DE ^ n ^ OP_MOD_EXP)
|
||||
a_list = [rng.randrange(0, modulus) for _ in range(n)]
|
||||
b_list = [rng.randrange(0, 1 << 16) for _ in range(n)]
|
||||
in_blob = build_bcgb(OP_MOD_EXP, modulus, a_list, b_list)
|
||||
out_blob, stderr = run_worker_binary(worker, in_blob)
|
||||
gpu = parse_bcgr(out_blob, n, out_bytes_per_inst=32)
|
||||
fail = 0; first_fail = None
|
||||
for i in range(n):
|
||||
expected = h_exp(a_list[i], b_list[i], modulus)
|
||||
if gpu[i] != expected:
|
||||
fail += 1
|
||||
if first_fail is None:
|
||||
first_fail = (i, a_list[i], b_list[i], expected, gpu[i])
|
||||
return report_pass_fail(f"mod-exp {label}", n, fail, first_fail, modulus)
|
||||
|
||||
|
||||
def test_add_no_mod(worker, n, label):
|
||||
rng = random.Random(0xC0FFEE57C0DEC0DE ^ n ^ OP_ADD_NO_MOD)
|
||||
a_list = [rng.randrange(0, MASK_256 + 1) for _ in range(n)]
|
||||
b_list = [rng.randrange(0, MASK_256 + 1) for _ in range(n)]
|
||||
in_blob = build_bcgb(OP_ADD_NO_MOD, None, a_list, b_list)
|
||||
out_blob, stderr = run_worker_binary(worker, in_blob)
|
||||
gpu = parse_bcgr(out_blob, n, out_bytes_per_inst=32)
|
||||
fail = 0; first_fail = None
|
||||
for i in range(n):
|
||||
expected = (a_list[i] + b_list[i]) & MASK_256
|
||||
if gpu[i] != expected:
|
||||
fail += 1
|
||||
if first_fail is None:
|
||||
first_fail = (i, a_list[i], b_list[i], expected, gpu[i])
|
||||
return report_pass_fail(f"add-no-mod {label}", n, fail, first_fail, None)
|
||||
|
||||
|
||||
def test_mul_no_mod(worker, n, label):
|
||||
rng = random.Random(0xC0FFEE57C0DEC0DE ^ n ^ OP_MUL_NO_MOD)
|
||||
a_list = [rng.randrange(0, MASK_256 + 1) for _ in range(n)]
|
||||
b_list = [rng.randrange(0, MASK_256 + 1) for _ in range(n)]
|
||||
in_blob = build_bcgb(OP_MUL_NO_MOD, None, a_list, b_list)
|
||||
out_blob, stderr = run_worker_binary(worker, in_blob)
|
||||
# 64 bytes per instance (512-bit product, full).
|
||||
gpu = parse_bcgr(out_blob, n, out_bytes_per_inst=64)
|
||||
fail = 0; first_fail = None
|
||||
for i in range(n):
|
||||
expected = a_list[i] * b_list[i]
|
||||
if gpu[i] != expected:
|
||||
fail += 1
|
||||
if first_fail is None:
|
||||
first_fail = (i, a_list[i], b_list[i], expected, gpu[i])
|
||||
return report_pass_fail(f"mul-no-mod {label}", n, fail, first_fail, None)
|
||||
|
||||
|
||||
# Op driver registry: name -> callable(worker, n, modulus, label).
|
||||
def driver_mod_add(worker, n, modulus, label):
|
||||
return test_binary_mod(worker, OP_MOD_ADD, "mod-add", h_add, n, modulus, label)
|
||||
|
||||
def driver_mod_sub(worker, n, modulus, label):
|
||||
return test_binary_mod(worker, OP_MOD_SUB, "mod-sub", h_sub, n, modulus, label)
|
||||
|
||||
def driver_mod_mul(worker, n, modulus, label):
|
||||
return test_binary_mod(worker, OP_MOD_MUL, "mod-mul", h_mul, n, modulus, label)
|
||||
|
||||
def driver_mod_sqr(worker, n, modulus, label):
|
||||
return test_unary_mod(worker, OP_MOD_SQR, "mod-sqr", h_sqr, n, modulus, label)
|
||||
|
||||
def driver_mod_inv(worker, n, modulus, label):
|
||||
return test_unary_mod(worker, OP_MOD_INV, "mod-inv", h_inv, n, modulus, label,
|
||||
ensure_invertible=True)
|
||||
|
||||
def driver_mod_exp(worker, n, modulus, label):
|
||||
return test_mod_exp_small(worker, n, modulus, label)
|
||||
|
||||
def driver_mod_reduce(worker, n, modulus, label):
|
||||
return test_unary_mod(worker, OP_MOD_REDUCE, "mod-reduce", h_rem, n, modulus, label)
|
||||
|
||||
def driver_add_no_mod(worker, n, modulus, label):
|
||||
return test_add_no_mod(worker, n, label)
|
||||
|
||||
def driver_mul_no_mod(worker, n, modulus, label):
|
||||
return test_mul_no_mod(worker, n, label)
|
||||
|
||||
|
||||
OPS = [
|
||||
("mod-add", driver_mod_add, True),
|
||||
("mod-sub", driver_mod_sub, True),
|
||||
("mod-mul", driver_mod_mul, True),
|
||||
("mod-sqr", driver_mod_sqr, True),
|
||||
("mod-inv", driver_mod_inv, True),
|
||||
("mod-exp", driver_mod_exp, True),
|
||||
("mod-reduce", driver_mod_reduce, True),
|
||||
("add-no-mod", driver_add_no_mod, False),
|
||||
("mul-no-mod", driver_mul_no_mod, False),
|
||||
]
|
||||
|
||||
|
||||
def main():
|
||||
|
|
@ -132,6 +299,8 @@ def main():
|
|||
ap.add_argument("worker", nargs="?", default="./cgbn-batch-worker")
|
||||
ap.add_argument("--n", type=int, default=None,
|
||||
help="single N override (default: sweep 32, 1k, 10k)")
|
||||
ap.add_argument("--op", default=None,
|
||||
help="restrict to one op name (e.g. mod-mul)")
|
||||
args = ap.parse_args()
|
||||
if not os.path.exists(args.worker):
|
||||
print(f"worker not found: {args.worker}")
|
||||
|
|
@ -140,14 +309,24 @@ def main():
|
|||
ns = [args.n] if args.n else [32, 1_000, 10_000]
|
||||
moduli = [
|
||||
(SECP256K1_P, "secp256k1 prime"),
|
||||
((1 << 256) - 189, "256-bit Mersenne-ish"),
|
||||
(0xfffffffffffffffffffffffffffffffffffffffffffffffe_fffefffe_fffefffe, "arbitrary odd"),
|
||||
]
|
||||
selected = [(name, fn, uses_mod) for name, fn, uses_mod in OPS
|
||||
if args.op is None or args.op == name]
|
||||
if not selected:
|
||||
print(f"no op matched --op {args.op}; valid: {[n for n,_,_ in OPS]}")
|
||||
sys.exit(1)
|
||||
|
||||
all_pass = True
|
||||
for n in ns:
|
||||
for m, name in moduli:
|
||||
ok = test_mod_mul(args.worker, n, m, f"mod-mul {name}")
|
||||
all_pass = all_pass and ok
|
||||
for op_name, fn, uses_mod in selected:
|
||||
if uses_mod:
|
||||
for m, m_name in moduli:
|
||||
ok = fn(args.worker, n, m, m_name)
|
||||
all_pass = all_pass and ok
|
||||
else:
|
||||
ok = fn(args.worker, n, None, "")
|
||||
all_pass = all_pass and ok
|
||||
|
||||
if all_pass:
|
||||
print("\nALL PASS")
|
||||
sys.exit(0)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue