diff --git a/src/gpu.cu b/src/gpu.cu index 3f3a96a..75ec325 100644 --- a/src/gpu.cu +++ b/src/gpu.cu @@ -1761,7 +1761,11 @@ void GpuThread::run() { auto start = std::chrono::steady_clock::now(); for (uint32_t i = 0; !should_stop(); i++) { - uint64_t start_seed = input.next(KernelFilterSeeds::threads_per_run); + auto seed = input.next(KernelFilterSeeds::threads_per_run); + if (!seed) { + break; + } + uint64_t start_seed = *seed; TRY_CUDA(cudaMemsetAsync(device_buffer_lens, 0, sizeof(*device_buffer_lens), stream)); diff --git a/src/gpu.h b/src/gpu.h index 971a229..1b3a28b 100644 --- a/src/gpu.h +++ b/src/gpu.h @@ -1,16 +1,25 @@ #pragma once #include "common.h" +#include struct SeedIterator { + uint64_t start; std::atomic_uint64_t pos; + std::optional end; - SeedIterator(uint64_t start) : pos(start) { + SeedIterator(uint64_t start, std::optional end = {}) : start(start), pos(start), end(end) { } - uint64_t next(uint64_t count) { - return pos.fetch_add(count); + std::optional next(uint64_t count) { + uint64_t seed = pos.fetch_add(count); + if (end && seed - start >= *end - start) return {}; + return seed; + } + + bool exhausted() const { + return end && pos.load(std::memory_order_relaxed) - start >= *end - start; } }; diff --git a/src/main.cpp b/src/main.cpp index 4a8037c..fe125f9 100644 --- a/src/main.cpp +++ b/src/main.cpp @@ -103,6 +103,7 @@ struct Args { std::optional server; std::optional output_file; std::optional start_seed; + std::optional end_seed; std::optional min_size; bool parse(int argc, const char **const argv) { @@ -151,6 +152,8 @@ struct Args { output_file = argv[i++]; } else if (std::strcmp("--start", arg) == 0) { if (!parse_argument_int(argc, argv, i, start_seed, [](int64_t start_seed){ return true; }, arg)) return false; + } else if (std::strcmp("--end", arg) == 0) { + if (!parse_argument_int(argc, argv, i, end_seed, [](int64_t end_seed){ return true; }, arg)) return false; } else if (std::strcmp("--size", arg) == 0) { if (!parse_argument_int(argc, argv, i, min_size, [](int32_t min_size){ return min_size >= 0; }, arg)) return false; } else { @@ -178,6 +181,16 @@ struct Args { return false; } + if (end_seed && !start_seed) { + std::fprintf(stderr, "--end requires --start\n"); + return false; + } + + if (end_seed && end_seed.value() == start_seed.value()) { + std::fprintf(stderr, "--end must not equal --start\n"); + return false; + } + if (min_size && !threads && client) { std::fprintf(stderr, "--size does nothing when not running cpu threads\n"); return false; @@ -195,7 +208,7 @@ uint64_t random_start_seed() { int main_inner(int argc, char **argv) { Args args{}; if (!args.parse(argc, const_cast(argv))) { - std::fprintf(stderr, "Usage:\n%s [--device ,,...] [--threads ] [--client ] [--server ] [--output ] [--start ] [--size ]\n", argv[0]); + std::fprintf(stderr, "Usage:\n%s [--device ,,...] [--threads ] [--client ] [--server ] [--output ] [--start ] [--end ] [--size ]\n", argv[0]); return 1; } @@ -237,7 +250,11 @@ int main_inner(int argc, char **argv) { #ifndef NO_GPU uint64_t start_seed = args.start_seed.value_or(random_start_seed()); - SeedIterator seed_range(start_seed); + std::optional end_seed; + if (args.end_seed) { + end_seed = (uint64_t)args.end_seed.value(); + } + SeedIterator seed_range(start_seed, end_seed); std::vector> gpu_threads; for (int device : args.devices) { @@ -287,6 +304,12 @@ int main_inner(int argc, char **argv) { std::printf("gpu_outputs.queue.size() = %zu\n", gpu_outputs.queue.size()); } +#ifndef NO_GPU + if (seed_range.exhausted()) { + running.store(false, std::memory_order_relaxed); + } +#endif + std::this_thread::sleep_for(std::chrono::seconds(1)); }