SPB Git forge
3commits 1branches 0releases
1.0 MBsize
maindefault branch
1 mo agolast push
Python 64.3% TeX 35.7%
3.5 KB · 101 lines python
Raw Blame History
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