#!/usr/bin/env python3 # ============================================================================= # gpu_sieve.py — Cycle 4: segmented sieve on Apple GPU via MLX Metal kernel # Author: Simon-Pierre Boucher — contact@spboucher.ai # ============================================================================= # One thread per (block, prime): marks multiples of base_primes[y] inside # block x of the segment byte-array. Write races are benign (all store 0). # Marking starts at max(p*p, first multiple >= block start) — identical # semantics to core.sieve_segment (validated in __main__ before any use). # ============================================================================= import numpy as np import mlx.core as mx _SRC = r""" uint b = thread_position_in_grid.x; // block index uint pi = thread_position_in_grid.y; // prime index int64_t lo = params[0]; int64_t hi = params[1]; int64_t B = params[2]; int64_t p = primes[pi]; if (p * p >= hi) return; int64_t blo = lo + (int64_t)b * B; int64_t bhi = blo + B < hi ? blo + B : hi; if (blo >= bhi) return; int64_t start = ((blo + p - 1) / p) * p; if (start < p * p) start = p * p; for (int64_t m = start; m < bhi; m += p) { out[m - lo] = 0; } """ _kernel = mx.fast.metal_kernel( name="segmented_sieve", input_names=["primes", "params"], output_names=["out"], source=_SRC, ) _BLOCK = 1_000_000 def sieve_segment_gpu(lo: int, hi: int, base_primes) -> np.ndarray: """Primes in [lo, hi) — GPU-marked composites, CPU extraction.""" lo = max(lo, 2) n = hi - lo if n <= 0: return np.array([], dtype=np.int64) primes_mx = mx.array(np.asarray(base_primes, dtype=np.int64)) params = mx.array(np.array([lo, hi, _BLOCK], dtype=np.int64)) nblocks = (n + _BLOCK - 1) // _BLOCK (out,) = _kernel( inputs=[primes_mx, params], grid=(nblocks, len(base_primes), 1), threadgroup=(min(nblocks, 32), 8, 1), output_shapes=[(n,)], output_dtypes=[mx.uint8], init_value=1, ) flags = np.array(out, copy=False) if lo <= 2: # positions of 0 and 1 if present (lo==2 means none) pass res = lo + np.nonzero(flags)[0].astype(np.int64) # base primes < sqrt(hi) inside the segment survive (marking starts at p^2) return res if __name__ == "__main__": import sys import time from core import primes_upto, sieve_segment print("device:", mx.default_device()) # correctness gate: GPU == CPU on varied windows ok = True windows = [(2, 10**6), (999_000_000, 1_001_000_000), (10**12, 10**12 + 5 * 10**7), (10**13 - 5 * 10**7, 10**13)] for lo, hi in windows: base = primes_upto(int(hi ** 0.5) + 1) g = sieve_segment_gpu(lo, hi, base) c = sieve_segment(lo, hi, base) same = np.array_equal(g, c) ok &= same print(f" [{'PASS' if same else 'FAIL'}] [{lo:.2e},{hi:.2e}): " f"{len(g)} primes (gpu) vs {len(c)} (cpu)") if not ok: sys.exit(1) # benchmark: one 1e9 chunk near 1e12, GPU vs single-core CPU lo, hi = 10**12, 10**12 + 10**9 base = primes_upto(int(hi ** 0.5) + 1) for name, fn in (("gpu", sieve_segment_gpu), ("cpu", sieve_segment)): t0 = time.time() tot = 0 for s in range(lo, hi, 50_000_000): tot += len(fn(s, min(s + 50_000_000, hi), base)) print(f" bench {name}: 1e9-chunk near 1e12 in {time.time()-t0:.1f}s " f"({tot} primes)") print("GPU SIEVE VALIDATED")