spb/prime-mystery-engine
Public
Python 64.3%
TeX 35.7%
1#!/usr/bin/env python32# =============================================================================3# gpu_sieve.py — Cycle 4: segmented sieve on Apple GPU via MLX Metal kernel4# Author: Simon-Pierre Boucher — contact@spboucher.ai5# =============================================================================6# One thread per (block, prime): marks multiples of base_primes[y] inside7# block x of the segment byte-array. Write races are benign (all store 0).8# Marking starts at max(p*p, first multiple >= block start) — identical9# semantics to core.sieve_segment (validated in __main__ before any use).10# =============================================================================1112import numpy as np1314import mlx.core as mx1516_SRC = r"""17 uint b = thread_position_in_grid.x; // block index18 uint pi = thread_position_in_grid.y; // prime index19 int64_t lo = params[0];20 int64_t hi = params[1];21 int64_t B = params[2];22 int64_t p = primes[pi];23 if (p * p >= hi) return;24 int64_t blo = lo + (int64_t)b * B;25 int64_t bhi = blo + B < hi ? blo + B : hi;26 if (blo >= bhi) return;27 int64_t start = ((blo + p - 1) / p) * p;28 if (start < p * p) start = p * p;29 for (int64_t m = start; m < bhi; m += p) {30 out[m - lo] = 0;31 }32"""3334_kernel = mx.fast.metal_kernel(35 name="segmented_sieve",36 input_names=["primes", "params"],37 output_names=["out"],38 source=_SRC,39)4041_BLOCK = 1_000_000424344def sieve_segment_gpu(lo: int, hi: int, base_primes) -> np.ndarray:45 """Primes in [lo, hi) — GPU-marked composites, CPU extraction."""46 lo = max(lo, 2)47 n = hi - lo48 if n <= 0:49 return np.array([], dtype=np.int64)50 primes_mx = mx.array(np.asarray(base_primes, dtype=np.int64))51 params = mx.array(np.array([lo, hi, _BLOCK], dtype=np.int64))52 nblocks = (n + _BLOCK - 1) // _BLOCK53 (out,) = _kernel(54 inputs=[primes_mx, params],55 grid=(nblocks, len(base_primes), 1),56 threadgroup=(min(nblocks, 32), 8, 1),57 output_shapes=[(n,)],58 output_dtypes=[mx.uint8],59 init_value=1,60 )61 flags = np.array(out, copy=False)62 if lo <= 2:63 # positions of 0 and 1 if present (lo==2 means none)64 pass65 res = lo + np.nonzero(flags)[0].astype(np.int64)66 # base primes < sqrt(hi) inside the segment survive (marking starts at p^2)67 return res686970if __name__ == "__main__":71 import sys72 import time73 from core import primes_upto, sieve_segment7475 print("device:", mx.default_device())76 # correctness gate: GPU == CPU on varied windows77 ok = True78 windows = [(2, 10**6), (999_000_000, 1_001_000_000),79 (10**12, 10**12 + 5 * 10**7), (10**13 - 5 * 10**7, 10**13)]80 for lo, hi in windows:81 base = primes_upto(int(hi ** 0.5) + 1)82 g = sieve_segment_gpu(lo, hi, base)83 c = sieve_segment(lo, hi, base)84 same = np.array_equal(g, c)85 ok &= same86 print(f" [{'PASS' if same else 'FAIL'}] [{lo:.2e},{hi:.2e}): "87 f"{len(g)} primes (gpu) vs {len(c)} (cpu)")88 if not ok:89 sys.exit(1)90 # benchmark: one 1e9 chunk near 1e12, GPU vs single-core CPU91 lo, hi = 10**12, 10**12 + 10**992 base = primes_upto(int(hi ** 0.5) + 1)93 for name, fn in (("gpu", sieve_segment_gpu), ("cpu", sieve_segment)):94 t0 = time.time()95 tot = 096 for s in range(lo, hi, 50_000_000):97 tot += len(fn(s, min(s + 50_000_000, hi), base))98 print(f" bench {name}: 1e9-chunk near 1e12 in {time.time()-t0:.1f}s "99 f"({tot} primes)")100 print("GPU SIEVE VALIDATED")101