Removed mmap mode because it was consistently slower than fd and a pain to maintain. Minor code cleanuop.
This commit is contained in:
+11
-30
@@ -34,9 +34,6 @@ def bench_mode(
|
||||
else:
|
||||
raise ValueError(f"Unknown io_mode: {io_mode}")
|
||||
|
||||
if "mmap" in io_mode:
|
||||
iocmd.append("--mmap")
|
||||
|
||||
# Print iocmd at start of row
|
||||
print(f"{' '.join(iocmd)[:20]:<20}", end="", flush=True)
|
||||
|
||||
@@ -118,35 +115,19 @@ def run_benchmark(args):
|
||||
print()
|
||||
print("-" * (20 + 8 * len(tcounts)))
|
||||
|
||||
for io_mode in ["dry", "dry-mmap", "null", "file", "file-mmap"]:
|
||||
for io_mode in ["dry", "null", "file"]:
|
||||
results = bench_mode(tcounts, io_mode, length, alg=args.alg, bench_file=bench_file)
|
||||
all_results[io_mode] = results
|
||||
|
||||
print("-" * (20 + 8 * len(tcounts)))
|
||||
|
||||
# Find best for file output
|
||||
best_speed = 0.0
|
||||
best_iocmd = None
|
||||
best_threads = 0
|
||||
for io_mode in ["file", "file-mmap"]:
|
||||
for w, sp, iocmd in all_results.get(io_mode, []):
|
||||
if sp > best_speed:
|
||||
best_speed = sp
|
||||
best_iocmd = iocmd
|
||||
best_threads = w
|
||||
|
||||
# Find fastest generation speed
|
||||
gen_speed = 0.0
|
||||
for io_mode in all_results:
|
||||
for _w, sp, _iocmd in all_results.get(io_mode, []):
|
||||
if sp > gen_speed:
|
||||
gen_speed = sp
|
||||
|
||||
if best_iocmd:
|
||||
threads = f" -t{best_threads}" if best_threads != 1 else ""
|
||||
print(
|
||||
f"\n>>> Fastest wrote {best_speed:.2f} GB/s, plain RNG {gen_speed:.0f} GB/s\n"
|
||||
f"randquik {' '.join(best_iocmd)}{threads}\n"
|
||||
)
|
||||
else:
|
||||
print("\nNo file output results collected.", file=sys.stderr)
|
||||
# Find fastest configuration and RNG speed
|
||||
gen_speed = max(r[1] for res in all_results.values() for r in res)
|
||||
best_speed, best_threads, best_iocmd = max(
|
||||
[(sp, w, iocmd) for w, sp, iocmd in all_results["file"]],
|
||||
)
|
||||
threads = f" -t{best_threads}" if best_threads != 1 else ""
|
||||
print(
|
||||
f"\n>>> Fastest wrote {best_speed:.2f} GB/s, plain RNG {gen_speed:.0f} GB/s\n"
|
||||
f"randquik {' '.join(best_iocmd)}{threads}\n"
|
||||
)
|
||||
|
||||
+1
-62
@@ -2,7 +2,6 @@
|
||||
|
||||
import argparse
|
||||
import gc
|
||||
import mmap
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
@@ -20,7 +19,6 @@ from randquik.utils import (
|
||||
from randquik.workers import (
|
||||
BLOCK_SIZE,
|
||||
FdProducer,
|
||||
MmapProducer,
|
||||
)
|
||||
|
||||
__all__ = ["main"]
|
||||
@@ -58,18 +56,6 @@ def parse_seeks(args):
|
||||
return iseek, oseek
|
||||
|
||||
|
||||
def _mmap(output: str | None, oseek: int, length: int, *, dry=False) -> mmap.mmap:
|
||||
with open_fd(output, length, oseek=oseek, dry=dry) as fd:
|
||||
try:
|
||||
return mmap.mmap(fd, length)
|
||||
except (OSError, ValueError) as e:
|
||||
if fd == -1:
|
||||
raise ValueError("Cannot mmap all memory: specify --len SIZE") from e
|
||||
if not output:
|
||||
raise ValueError("Cannot mmap stdout: remove --mmap or use -o FILE") from e
|
||||
raise ValueError(f"Cannot mmap {output}: {e}") from e
|
||||
|
||||
|
||||
def singlethreaded(args, total_bytes, oseek, start_time, key, seed_for_display):
|
||||
if args.verbose:
|
||||
mode_desc = "infinite output" if total_bytes is None else "workers=0"
|
||||
@@ -146,11 +132,6 @@ def _main():
|
||||
type=str,
|
||||
default=None,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mmap",
|
||||
action="store_true",
|
||||
help="Use file-backed mmap for output instead of writing via fd",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--benchmark",
|
||||
action="store_true",
|
||||
@@ -224,50 +205,8 @@ def _main():
|
||||
return singlethreaded(args, total_bytes, oseek, start_time, key, seed_for_display)
|
||||
|
||||
workers = args.threads if args.threads is not None else 1
|
||||
# File-backed mmap output
|
||||
if args.mmap:
|
||||
with (
|
||||
_mmap(
|
||||
args.output,
|
||||
oseek,
|
||||
length=oseek + (total_bytes if total_bytes is not None else 0),
|
||||
dry=args.dry,
|
||||
) as mm,
|
||||
):
|
||||
producer = MmapProducer(
|
||||
workers,
|
||||
key,
|
||||
ciph,
|
||||
total_bytes,
|
||||
mm,
|
||||
use_madvise=True,
|
||||
oseek=oseek,
|
||||
iseek=iseek,
|
||||
)
|
||||
progress_state = {"written": 0}
|
||||
producer.progress_state = progress_state
|
||||
producer.start()
|
||||
|
||||
progress = ProgressDisplay(
|
||||
total_bytes, start_time, progress_state, seed=seed_for_display
|
||||
)
|
||||
if not args.quiet:
|
||||
progress.start()
|
||||
|
||||
try:
|
||||
producer.join()
|
||||
finally:
|
||||
progress.stop()
|
||||
elapsed = time.perf_counter() - start_time
|
||||
if not args.quiet:
|
||||
print_summary(producer.written, elapsed, "wrote", seed=seed_for_display)
|
||||
if args.verbose:
|
||||
stats = producer.get_combined_stats()
|
||||
print(stats.format_report(f"Workers (×{workers})"), file=sys.stderr)
|
||||
producer.cleanup()
|
||||
return
|
||||
|
||||
# Standard file mode with ring buffers
|
||||
# File mode with ring buffers
|
||||
with open_fd(args.output, total_bytes, dry=args.dry, oseek=oseek) as fd:
|
||||
producer = FdProducer(
|
||||
workers, key, ciph, total_bytes, fd, dry=args.dry, iseek=iseek, profile=args.verbose
|
||||
|
||||
+23
-54
@@ -1,73 +1,34 @@
|
||||
"""File I/O helpers for output handling and mmap operations."""
|
||||
"""File I/O helpers for output handling."""
|
||||
|
||||
import contextlib
|
||||
import ctypes
|
||||
import mmap
|
||||
import errno
|
||||
import os
|
||||
import pathlib
|
||||
import sys
|
||||
from collections.abc import Generator
|
||||
|
||||
__all__ = [
|
||||
"HAS_MADVISE",
|
||||
"madvise_buffer",
|
||||
"madvise_file",
|
||||
"open_fd",
|
||||
"open_memoryview",
|
||||
]
|
||||
|
||||
# Platform-specific constants for madvise
|
||||
MADV_DONTNEED = 4 # Linux/macOS
|
||||
MADV_SEQUENTIAL = 2 # Hint for sequential access
|
||||
MADV_WILLNEED = 3 # Pre-fault pages
|
||||
MADV_RANDOM = getattr(mmap, "MADV_RANDOM", 1) # Not available on Windows
|
||||
|
||||
# Try to get O_DIRECT (Linux only, not available on macOS)
|
||||
O_DIRECT = getattr(os, "O_DIRECT", 0)
|
||||
|
||||
# Load libc for madvise (not available on Windows)
|
||||
HAS_MADVISE = False
|
||||
_madvise = None
|
||||
if sys.platform != "win32":
|
||||
try:
|
||||
if sys.platform == "darwin":
|
||||
_libc = ctypes.CDLL("libc.dylib", use_errno=True)
|
||||
else:
|
||||
_libc = ctypes.CDLL("libc.so.6", use_errno=True)
|
||||
_madvise = _libc.madvise
|
||||
_madvise.argtypes = [ctypes.c_void_p, ctypes.c_size_t, ctypes.c_int]
|
||||
_madvise.restype = ctypes.c_int
|
||||
HAS_MADVISE = True
|
||||
except (OSError, AttributeError):
|
||||
pass
|
||||
|
||||
|
||||
def _madvise_call(mm: mmap.mmap, advice: int, offset: int = 0, length: int = 0):
|
||||
"""Call madvise with specified advice."""
|
||||
if length == 0:
|
||||
length = len(mm)
|
||||
mm.madvise(advice, offset, length)
|
||||
|
||||
|
||||
def madvise_buffer(mm: mmap.mmap, offset: int = 0, length: int = 0):
|
||||
"""Mark mmap region for random access (avoid caching)."""
|
||||
_madvise_call(mm, MADV_RANDOM, offset, length)
|
||||
|
||||
|
||||
def madvise_file(mm: mmap.mmap, offset: int = 0, length: int = 0):
|
||||
"""Mark mmap region for sequential access (pre-fault pages)."""
|
||||
_madvise_call(mm, MADV_WILLNEED, offset, length)
|
||||
|
||||
|
||||
def _open_output(
|
||||
output_path: str,
|
||||
total_bytes: int | None,
|
||||
oseek: int = 0,
|
||||
) -> int:
|
||||
"""Open output file descriptor, preallocate and apply platform hints."""
|
||||
) -> tuple[int, bool]:
|
||||
"""Open output file descriptor, preallocate and apply platform hints.
|
||||
|
||||
Returns:
|
||||
Tuple of (file descriptor, whether we created the file)
|
||||
"""
|
||||
created = False
|
||||
if output_path:
|
||||
flags = os.O_RDWR | os.O_CREAT
|
||||
fd = os.open(str(pathlib.Path(output_path)), flags, 0o644)
|
||||
path = pathlib.Path(output_path)
|
||||
created = not path.exists()
|
||||
flags = os.O_WRONLY | os.O_CREAT
|
||||
fd = os.open(str(path), flags, 0o644)
|
||||
else:
|
||||
if sys.stdout.isatty():
|
||||
raise ValueError("Refusing to write binary data to terminal. Use -o to specify a file.")
|
||||
@@ -97,7 +58,7 @@ def _open_output(
|
||||
except (OSError, AttributeError, ImportError):
|
||||
pass
|
||||
|
||||
return fd
|
||||
return fd, created
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
@@ -124,9 +85,17 @@ def open_fd(
|
||||
if not output_path:
|
||||
yield sys.stdout.fileno()
|
||||
return
|
||||
fd = _open_output(output_path, total_bytes, oseek)
|
||||
fd, created = _open_output(output_path, total_bytes, oseek)
|
||||
try:
|
||||
yield fd
|
||||
except OSError as e:
|
||||
if e.errno == errno.ENOSPC:
|
||||
# Clean up file we created on disk full
|
||||
if created:
|
||||
with contextlib.suppress(Exception):
|
||||
os.unlink(output_path)
|
||||
raise ValueError(f"No space left on device: {output_path}") from None
|
||||
raise
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
os.close(fd)
|
||||
|
||||
@@ -2,17 +2,14 @@
|
||||
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
|
||||
from randquik.io import madvise_file
|
||||
from randquik.utils import stopwatch
|
||||
|
||||
__all__ = [
|
||||
"BLOCK_SIZE",
|
||||
"FdProducer",
|
||||
"MmapProducer",
|
||||
"WorkerStats",
|
||||
]
|
||||
|
||||
@@ -324,176 +321,3 @@ class FdProducer:
|
||||
lines.append(f" write time: {self.write_time * 1000:.1f}ms")
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
class MmapProducer:
|
||||
"""Multi-threaded producer that writes directly into mmap."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
workers: int,
|
||||
key: bytes,
|
||||
ciph,
|
||||
total_bytes: int | None,
|
||||
mm_raw=None,
|
||||
use_madvise: bool = False,
|
||||
oseek: int = 0,
|
||||
iseek: int = 0,
|
||||
block_size: int = BLOCK_SIZE,
|
||||
dry: bool = False,
|
||||
):
|
||||
self.workers = workers
|
||||
self.key = key
|
||||
self.ciph = ciph
|
||||
self.total_bytes = total_bytes
|
||||
self.use_madvise = use_madvise
|
||||
self.oseek = oseek
|
||||
self.iseek = iseek
|
||||
self.block_size = block_size
|
||||
self.dry = dry
|
||||
self.written = 0
|
||||
|
||||
# For dry runs, create an anonymous mmap
|
||||
import mmap as mmap_module
|
||||
|
||||
if dry:
|
||||
length = oseek + (total_bytes if total_bytes is not None else 0)
|
||||
self.mm_raw = mmap_module.mmap(-1, length)
|
||||
self._owns_mmap = True
|
||||
else:
|
||||
self.mm_raw = mm_raw
|
||||
self._owns_mmap = False
|
||||
|
||||
# iseek handling: which block to start at, and offset within first block
|
||||
self.start_block = iseek // block_size
|
||||
self.start_offset = iseek % block_size
|
||||
|
||||
if total_bytes is None:
|
||||
self.num_blocks = sys.maxsize // block_size
|
||||
else:
|
||||
self.num_blocks = (total_bytes + block_size - 1) // block_size
|
||||
# If first block is partial, we need one more block to cover total_bytes
|
||||
if self.start_offset > 0 and total_bytes > 0:
|
||||
self.num_blocks = (self.start_offset + total_bytes + block_size - 1) // block_size
|
||||
self._lock = threading.Lock()
|
||||
self._next_block = 0
|
||||
self._quit = threading.Event()
|
||||
self._threads: list[threading.Thread] = []
|
||||
self.progress_state: dict | None = None
|
||||
|
||||
# Per-worker stats, collected after threads finish
|
||||
self._worker_stats: list[WorkerStats] = []
|
||||
self._stats_lock = threading.Lock()
|
||||
|
||||
def start(self):
|
||||
"""Start worker threads."""
|
||||
for _ in range(self.workers):
|
||||
t = threading.Thread(target=self._worker, daemon=True)
|
||||
self._threads.append(t)
|
||||
t.start()
|
||||
|
||||
def _worker(self):
|
||||
# Create thread-local memoryview from mmap
|
||||
view = memoryview(self.mm_raw)
|
||||
# Temporary buffer for first/last partial blocks
|
||||
tmp_buf = bytearray(self.block_size)
|
||||
stats = WorkerStats()
|
||||
timer = stopwatch()
|
||||
try:
|
||||
while True:
|
||||
next(timer) # reset
|
||||
with self._lock:
|
||||
stats.lock_acquire_time += next(timer)
|
||||
if self._next_block >= self.num_blocks:
|
||||
return
|
||||
block_num = self._next_block
|
||||
self._next_block += 1
|
||||
|
||||
# Calculate actual nonce (accounting for iseek)
|
||||
actual_block = self.start_block + block_num
|
||||
nonce = actual_block.to_bytes(self.ciph.NONCEBYTES, "little")
|
||||
|
||||
# Calculate byte range within output
|
||||
# block_num=0 corresponds to output byte 0
|
||||
# But if start_offset > 0, first block is partial
|
||||
next(timer) # reset for crypto timing
|
||||
if block_num == 0 and self.start_offset > 0:
|
||||
# First partial block: generate full block, copy from start_offset
|
||||
self.ciph.stream(self.key, nonce, self.block_size, into=tmp_buf)
|
||||
if self.total_bytes is None:
|
||||
copy_len = self.block_size - self.start_offset
|
||||
else:
|
||||
copy_len = min(self.block_size - self.start_offset, self.total_bytes)
|
||||
out_start = self.oseek
|
||||
out_end = out_start + copy_len
|
||||
view[out_start:out_end] = tmp_buf[
|
||||
self.start_offset : self.start_offset + copy_len
|
||||
]
|
||||
written = copy_len
|
||||
else:
|
||||
# Full block or last partial block
|
||||
# Output position: account for first block being partial
|
||||
if self.start_offset > 0:
|
||||
out_start = (
|
||||
self.oseek
|
||||
+ (self.block_size - self.start_offset)
|
||||
+ (block_num - 1) * self.block_size
|
||||
)
|
||||
else:
|
||||
out_start = self.oseek + block_num * self.block_size
|
||||
if self.total_bytes is None:
|
||||
out_end = out_start + self.block_size
|
||||
else:
|
||||
out_end = min(out_start + self.block_size, self.oseek + self.total_bytes)
|
||||
size = out_end - out_start
|
||||
self.ciph.stream(self.key, nonce, size, into=view[out_start:out_end])
|
||||
written = size
|
||||
|
||||
stats.crypto_time += next(timer)
|
||||
stats.blocks_processed += 1
|
||||
stats.bytes_generated += written
|
||||
|
||||
if self.use_madvise and self.mm_raw:
|
||||
next(timer) # reset
|
||||
madvise_file(self.mm_raw, out_start, written)
|
||||
stats.madvise_time += next(timer)
|
||||
|
||||
next(timer) # reset
|
||||
with self._lock:
|
||||
self.written += written
|
||||
if self.progress_state is not None:
|
||||
self.progress_state["written"] = self.written
|
||||
stats.notify_time += next(timer)
|
||||
except BaseException as e:
|
||||
logging.exception("Worker thread exception: %s", e)
|
||||
finally:
|
||||
view.release()
|
||||
# Collect stats before signaling quit
|
||||
with self._stats_lock:
|
||||
self._worker_stats.append(stats)
|
||||
self._quit.set()
|
||||
self._quit.wait()
|
||||
|
||||
def join(self):
|
||||
"""Wait for all worker threads to finish."""
|
||||
for t in self._threads:
|
||||
t.join()
|
||||
|
||||
def stop(self):
|
||||
"""Signal workers to stop."""
|
||||
self._quit.set()
|
||||
|
||||
def cleanup(self):
|
||||
"""Release references to mmap resources."""
|
||||
if self._owns_mmap and self.mm_raw is not None:
|
||||
self.mm_raw.close()
|
||||
self.mm_raw = None
|
||||
self._threads.clear()
|
||||
|
||||
def get_combined_stats(self) -> WorkerStats:
|
||||
"""Get combined stats from all workers. Call after join()."""
|
||||
combined = WorkerStats()
|
||||
with self._stats_lock:
|
||||
for s in self._worker_stats:
|
||||
combined.merge(s)
|
||||
return combined
|
||||
|
||||
Reference in New Issue
Block a user