Skip to documentation
SLOP

tiny.simd.thread.pool

Reference tiny.simd thread pool

Defined in thread.

API (22)

Actions

Public operations.

Types and contracts

Public types and contracts.

Values and defaults

Public values and defaults.

No direct callersNo direct callsthreadpool
Static calls · unresolved targets: unknown · external targets: unknown.

Source

Called byCallsNo direct callsprivate sourcelib.simd.src.thread.pool.ThreadPooldivideRangetest sourcelib.simd.src.thread.pooltest: Highway worker range division i...thread.poolworkerRange
Static calls · unresolved targets: 0 · external targets: 0.

Source: lib/simd/src/thread/pool.zig

zig
const std = @import("std");const sys = @import("sys");const simd = @import("../root.zig");const spin = @import("spin.zig");const wait = @import("wait.zig");const autotune = simd.autotune;const topology = simd.topology;pub const max_clusters: usize = 33;pub const all_clusters: usize = max_clusters - 1;pub const max_threads: usize = 127;pub const max_workers: usize = max_threads + 1;pub const max_callers: usize = 60;pub const caller_name_capacity: usize = 64;pub const max_victims: usize = 4;pub const max_configs: usize = 4;pub const PoolWaitMode = enum(u8) {    block = 1,    spin,};pub const WaitType = enum(u8) {    block,    spin_shared,    spin_separate,};pub fn waitName(wait_type: WaitType) []const u8 {    return switch (wait_type) {        .block => "Block",        .spin_shared => "Single",        .spin_separate => "Separate",    };}pub const Exit = enum(u32) {    none,    loop,    thread,};pub const Config = struct {    spin_type: spin.SpinType = .pause,    wait_type: WaitType = .spin_separate,    reserved: [2]u8 = @splat(0),    pub fn formatName(        self: Config,        storage: []u8,    ) error{NoSpaceLeft}![]const u8 {        return std.fmt.bufPrint(            storage,            "{s:<14} {s:<9}",            .{ spin.name(self.spin_type), waitName(self.wait_type) },        );    }    pub fn candidates(        mode: PoolWaitMode,        storage: *[max_configs]Config,    ) []const Config {        return switch (mode) {            .block => blk: {                storage[0] = .{ .wait_type = .block };                break :blk storage[0..1];            },            .spin => blk: {                const detected = spin.detectSpin(0);                const spin_types: [2]spin.SpinType = .{ detected, .pause };                const spin_count: usize = if (detected == .pause) 1 else 2;                var count: usize = 0;                for (spin_types[0..spin_count]) |spin_type| {                    storage[count] = .{                        .spin_type = spin_type,                        .wait_type = .spin_shared,                    };                    count += 1;                    storage[count] = .{                        .spin_type = spin_type,                        .wait_type = .spin_separate,                    };                    count += 1;                }                break :blk storage[0..count];            },        };    }    fn encode(self: Config) u16 {        return @as(u16, @backingInt(self.spin_type)) |            @as(u16, @backingInt(self.wait_type)) << 8;    }    fn decode(bits: u16) Config {        return .{            .spin_type = @fromBackingInt(@intCast(@as(u8, @truncate(bits)))),            .wait_type = @fromBackingInt(@intCast(@as(u8, @truncate(bits >> 8)))),        };    }};pub const PoolWorkerMapping = struct {    cluster_index: u8 = 0,    max_cluster_workers: usize = 0,    pub fn init(        cluster_index: usize,        max_cluster_workers: usize,    ) error{ InvalidCluster, EmptyCluster }!PoolWorkerMapping {        if (cluster_index > all_clusters) return error.InvalidCluster;        if (max_cluster_workers == 0) return error.EmptyCluster;        return .{            .cluster_index = @intCast(cluster_index),            .max_cluster_workers = max_cluster_workers,        };    }    pub fn clusterIndex(self: PoolWorkerMapping) usize {        return self.cluster_index;    }    pub fn maxClusterWorkers(self: PoolWorkerMapping) usize {        return self.max_cluster_workers;    }    pub fn globalIndex(        self: PoolWorkerMapping,        worker_index: usize,    ) usize {        if (self.max_cluster_workers == 0) return worker_index;        if (self.cluster_index == all_clusters) {            std.debug.assert(worker_index < all_clusters);            return worker_index * self.max_cluster_workers;        }        std.debug.assert(worker_index < self.max_cluster_workers);        return self.cluster_index * self.max_cluster_workers + worker_index;    }};pub const ShuffledIota = struct {    coprime: u32 = 1,    pub fn init(coprime: u32) ShuffledIota {        return .{ .coprime = coprime };    }    pub fn next(        self: ShuffledIota,        current: u32,        size: u32,    ) u32 {        std.debug.assert(size != 0);        std.debug.assert(current < size);        return @intCast(            (@as(u64, current) + self.coprime) % @as(u64, size),        );    }    pub fn coprimeNonzero(a_value: u32, b_value: u32) bool {        std.debug.assert(a_value != 0);        std.debug.assert(b_value != 0);        var a = a_value;        var b = b_value;        const trailing_a = @ctz(a);        const trailing_b = @ctz(b);        if (@min(trailing_a, trailing_b) != 0) return false;        a >>= @intCast(trailing_a);        b >>= @intCast(trailing_b);        while (true) {            const previous_a = a;            a = @max(previous_a, b);            b = @min(previous_a, b);            if (b == 1) return true;            a -= b;            if (a == 0) return false;            a >>= @intCast(@ctz(a));        }    }    pub fn findAnotherCoprime(size: u32, start: u32) u32 {        std.debug.assert(size != 0);        if (size <= 2) return 1;        const increment: u32 = if (size & 1 == 0) 2 else 1;        var candidate = start | 1;        var attempts: u64 = 0;        const max_attempts = @as(u64, size) * 16;        while (attempts < max_attempts) : (attempts += 1) {            if (coprimeNonzero(candidate, size)) return candidate;            candidate +%= increment;            if (candidate == 0) candidate = 1;        }        unreachable;    }};pub const Caller = struct {    index: u8 = 0,    pub fn id(self: Caller) usize {        return self.index;    }};pub const CallerStats = struct {    runs: u64 = 0,    tasks: u64 = 0,    workers: u64 = 0,    elapsed_ns: u64 = 0,};pub const PoolStats = struct {    runs: u64 = 0,    serial_runs: u64 = 0,    threaded_runs: u64 = 0,    tasks: u64 = 0,    stolen_tasks: u64 = 0,    elapsed_ns: u64 = 0,};const CallerRegistry = struct {    guard: std.atomic.Mutex = .unlocked,    count: usize = 1,    lengths: [max_callers]u8 = @splat(0),    names: [max_callers][caller_name_capacity]u8 = @splat(@splat(0)),    fn add(        self: *CallerRegistry,        caller_name: []const u8,    ) error{ EmptyCallerName, CallerNameTooLong, TooManyCallers }!Caller {        if (caller_name.len == 0) return error.EmptyCallerName;        if (caller_name.len > caller_name_capacity) {            return error.CallerNameTooLong;        }        while (!self.guard.tryLock()) std.atomic.spinLoopHint();        defer self.guard.unlock();        for (self.names[1..self.count], 1..) |stored, index| {            const length = self.lengths[index];            if (std.mem.eql(u8, stored[0..length], caller_name)) {                return .{ .index = @intCast(index) };            }        }        if (self.count == max_callers) return error.TooManyCallers;        const index = self.count;        @memcpy(self.names[index][0..caller_name.len], caller_name);        self.lengths[index] = @intCast(caller_name.len);        self.count += 1;        return .{ .index = @intCast(index) };    }    fn name(self: *CallerRegistry, caller: Caller) ?[]const u8 {        while (!self.guard.tryLock()) std.atomic.spinLoopHint();        defer self.guard.unlock();        if (caller.index == 0 or caller.index >= self.count) return null;        const length = self.lengths[caller.index];        return self.names[caller.index][0..length];    }};var caller_registry = CallerRegistry{};pub fn addCaller(    name: []const u8,) error{ EmptyCallerName, CallerNameTooLong, TooManyCallers }!Caller {    return caller_registry.add(name);}pub fn callerName(caller: Caller) ?[]const u8 {    return caller_registry.name(caller);}pub fn workerRange(    begin: u64,    end: u64,    worker_count: usize,    worker_index: usize,) error{ InvalidRange, EmptyWorkers, WorkerOutOfBounds, TooManyTasks }!struct {    begin: u64,    end: u64,} {    if (begin > end) return error.InvalidRange;    if (worker_count == 0) return error.EmptyWorkers;    if (worker_index >= worker_count) return error.WorkerOutOfBounds;    const task_count_u64 = end - begin;    if (task_count_u64 > std.math.maxInt(usize)) return error.TooManyTasks;    const task_count: usize = @intCast(task_count_u64);    const minimum = task_count / worker_count;    const remainder = task_count % worker_count;    const local_begin = worker_index * minimum + @min(worker_index, remainder);    const local_count = minimum + @intFromBool(worker_index < remainder);    return .{        .begin = begin + local_begin,        .end = begin + local_begin + local_count,    };}const Worker = struct {    begin: std.atomic.Value(usize) = std.atomic.Value(usize).init(0),    end: usize = 0,    last_tasks: usize = 0,    last_stolen: usize = 0,    wait_epoch: std.atomic.Value(u32) = std.atomic.Value(u32).init(0),    barrier_epoch: std.atomic.Value(u32) = std.atomic.Value(u32).init(0),    victims: [max_victims]u8 = @splat(0),    victim_count: u8 = 0,    padding: [if (@sizeOf(usize) == 8) 19 else 35]u8 = @splat(0),};comptime {    if (@sizeOf(Worker) != 64) @compileError("thread worker must occupy one cache line");}const Callback = *const fn (*anyopaque, u64, usize) void;const AutoTuner = autotune.AutoTune(Config, max_configs, 30);pub const ThreadPool = struct {    pub const State = enum(u8) {        initialization,        steady,        teardown,    };    pub const Capacity = struct {        requested_threads: usize,        threads: u8,        workers: u8,        pub fn derive(requested_threads: usize) Capacity {            const supported = topology.haveThreadingSupport() and                sys.thread.threadsSupported();            const thread_count: usize = if (supported)                @min(requested_threads, max_threads)            else                0;            const worker_count: usize = thread_count + 1;            return .{                .requested_threads = requested_threads,                .threads = @intCast(thread_count),                .workers = @intCast(worker_count),            };        }        pub fn wasClamped(self: Capacity) bool {            return self.requested_threads != self.threads;        }    };    pub const StartError = sys.thread.SpawnError || error{        AlreadyDeinitialized,        PoolMoved,    };    pub const RunError = StartError || error{        Busy,        InvalidRange,        TooManyTasks,    };    state: State = .initialization,    capacity: Capacity,    mapping: PoolWorkerMapping,    address: usize = 0,    wait_mode: std.atomic.Value(u8) =        std.atomic.Value(u8).init(@backingInt(PoolWaitMode.block)),    config_bits: std.atomic.Value(u16) =        std.atomic.Value(u16).init((Config{ .wait_type = .block }).encode()),    epoch: std.atomic.Value(u32) = std.atomic.Value(u32).init(0),    work_available: std.atomic.Value(bool) =        std.atomic.Value(bool).init(false),    exiting: std.atomic.Value(bool) = std.atomic.Value(bool).init(false),    busy: std.atomic.Value(bool) = std.atomic.Value(bool).init(false),    synchronization: sys.thread.Mutex = .{},    done: sys.thread.Condition = .{},    started_threads: usize = 0,    handles: [max_threads]sys.thread.JoinHandle = undefined,    workers: [max_workers]Worker align(64) = @splat(.{}),    task_context: *anyopaque = undefined,    task_callback: Callback = undefined,    current_begin: u64 = 0,    current_tasks: usize = 0,    stats: PoolStats = .{},    caller_stats: [max_callers]CallerStats = @splat(.{}),    block_tuner: AutoTuner = .{},    spin_tuner: AutoTuner = .{},    pub fn init(        requested_threads: usize,        mapping: PoolWorkerMapping,    ) ThreadPool {        var self = ThreadPool{            .capacity = Capacity.derive(requested_threads),            .mapping = mapping,        };        var candidates: [max_configs]Config = undefined;        self.block_tuner.setCandidates(            Config.candidates(.block, &candidates),        ) catch unreachable;        self.spin_tuner.setCandidates(            Config.candidates(.spin, &candidates),        ) catch unreachable;        self.config_bits.store(            self.block_tuner.nextConfig().encode(),            .monotonic,        );        self.configureWorkers();        return self;    }    pub fn maxThreads() usize {        if (!topology.haveThreadingSupport() or            !sys.thread.threadsSupported())        {            return 0;        }        var affinity = topology.LogicalProcessorSet{};        if (topology.getThreadAffinity(&affinity)) {            const count = affinity.count();            if (count != 0) return @min(count - 1, max_threads);        }        return @min(            topology.totalLogicalProcessors() -| 1,            max_threads,        );    }    pub fn numThreadsFromCores(        allocator: std.mem.Allocator,    ) std.mem.Allocator.Error!usize {        if (!topology.haveThreadingSupport() or            !sys.thread.threadsSupported())        {            return 0;        }        var detected = try topology.init(allocator);        defer detected.deinit(allocator);        if (detected.packages.len == 0) return maxThreads();        return @min(detected.packages[0].cores.len -| 1, max_threads);    }    pub fn start(self: *ThreadPool) StartError!void {        if (self.state == .teardown) return error.AlreadyDeinitialized;        if (self.state == .steady) {            try self.requireStableAddress();            return;        }        self.address = @intFromPtr(self);        if (self.capacity.threads == 0) {            self.state = .steady;            return;        }        var spawned: usize = 0;        while (spawned < self.capacity.threads) : (spawned += 1) {            self.handles[spawned] = sys.thread.spawn(workerMain, .{                WorkerContext{                    .pool = self,                    .worker_index = spawned + 1,                },            }) catch |err| {                self.stopSpawned(spawned);                return err;            };            nameWorker(self.handles[spawned], spawned);        }        self.synchronization.lock();        while (self.started_threads != self.capacity.threads) {            self.done.wait(&self.synchronization);        }        self.synchronization.unlock();        self.state = .steady;    }    pub fn deinit(self: *ThreadPool) void {        if (self.state == .teardown) return;        if (self.state == .initialization) {            self.state = .teardown;            self.address = 0;            return;        }        self.requireStableAddress() catch unreachable;        std.debug.assert(!self.busy.load(.acquire));        if (self.capacity.threads != 0) {            self.signalExit();            for (self.handles[0..self.capacity.threads]) |handle| {                handle.join();            }        }        self.state = .teardown;        self.address = 0;    }    pub fn numWorkers(self: *const ThreadPool) usize {        return self.capacity.workers;    }    pub fn wasClamped(self: *const ThreadPool) bool {        return self.capacity.wasClamped();    }    pub fn setWaitMode(        self: *ThreadPool,        mode: PoolWaitMode,    ) (error{Busy} || StartError)!void {        if (self.state == .teardown) return error.AlreadyDeinitialized;        if (self.busy.cmpxchgStrong(            false,            true,            .acquire,            .monotonic,        ) != null) return error.Busy;        defer self.busy.store(false, .release);        self.wait_mode.store(@backingInt(mode), .release);        const mode_tuner = self.tuner(mode);        const selected = mode_tuner.best() orelse mode_tuner.nextConfig();        self.config_bits.store(selected.encode(), .release);        if (self.state == .steady and self.capacity.threads != 0) {            try self.requireStableAddress();            const epoch = self.wakeWorkers(false, true);            self.waitForWorkers(epoch);        }    }    pub fn waitMode(self: *const ThreadPool) PoolWaitMode {        return @fromBackingInt(@intCast(self.wait_mode.load(.acquire)));    }    pub fn config(self: *const ThreadPool) Config {        return Config.decode(self.config_bits.load(.acquire));    }    pub fn autoTuneComplete(self: *const ThreadPool) bool {        return self.tunerConst(self.waitMode()).best() != null;    }    pub fn autoTuneCosts(self: *ThreadPool) []autotune.CostDistribution {        return self.tuner(self.waitMode()).costs();    }    pub fn statsSnapshot(self: *ThreadPool) PoolStats {        self.synchronization.lock();        defer self.synchronization.unlock();        return self.stats;    }    pub fn callerStats(        self: *ThreadPool,        caller: Caller,    ) ?CallerStats {        if (caller.index >= max_callers) return null;        self.synchronization.lock();        defer self.synchronization.unlock();        return self.caller_stats[caller.index];    }    pub fn resetStats(self: *ThreadPool) error{Busy}!void {        if (self.busy.cmpxchgStrong(            false,            true,            .acquire,            .monotonic,        ) != null) return error.Busy;        defer self.busy.store(false, .release);        self.synchronization.lock();        defer self.synchronization.unlock();        self.stats = .{};        self.caller_stats = @splat(.{});    }    pub fn globalWorkerIndex(        self: *const ThreadPool,        local_worker_index: usize,    ) error{WorkerOutOfBounds}!usize {        if (local_worker_index >= self.numWorkers()) {            return error.WorkerOutOfBounds;        }        return self.mapping.globalIndex(local_worker_index);    }    pub fn run(        self: *ThreadPool,        begin: u64,        end: u64,        context: anytype,        comptime body: fn (@TypeOf(context), u64, usize) void,    ) RunError!void {        return self.runWithCaller(begin, end, .{}, context, body);    }    pub fn runWithCaller(        self: *ThreadPool,        begin: u64,        end: u64,        caller: Caller,        context: anytype,        comptime body: fn (@TypeOf(context), u64, usize) void,    ) RunError!void {        if (begin > end) return error.InvalidRange;        const task_count_u64 = end - begin;        if (task_count_u64 > std.math.maxInt(usize)) {            return error.TooManyTasks;        }        const task_count: usize = @intCast(task_count_u64);        const Context = @TypeOf(context);        const Adapter = CallbackAdapter(Context, body);        var context_storage = context;        const started_ns = nowNanoseconds();        if (task_count <= 1 or self.numWorkers() == 1) {            var task = begin;            while (task < end) : (task += 1) body(context, task, 0);            self.recordRun(                caller,                task_count,                false,                0,                elapsedNanoseconds(started_ns),            );            return;        }        if (self.busy.cmpxchgStrong(            false,            true,            .acquire,            .monotonic,        ) != null) return error.Busy;        defer self.busy.store(false, .release);        try self.start();        try self.requireStableAddress();        self.task_context = @ptrCast(&context_storage);        self.task_callback = Adapter.call;        self.current_begin = begin;        self.current_tasks = task_count;        self.divideRange(begin, end);        for (self.workers[0..self.numWorkers()]) |*worker| {            worker.last_tasks = 0;            worker.last_stolen = 0;        }        const epoch = self.wakeWorkers(            true,            self.config().wait_type == .block,        );        self.runWorker(0);        self.waitForWorkers(epoch);        var stolen_tasks: usize = 0;        for (self.workers[0..self.numWorkers()]) |worker| {            stolen_tasks += worker.last_stolen;        }        const elapsed_ns = elapsedNanoseconds(started_ns);        self.recordRun(            caller,            task_count,            true,            stolen_tasks,            elapsed_ns,        );        self.notifyAutotune(elapsed_ns);    }    fn configureWorkers(self: *ThreadPool) void {        const worker_count: u32 = self.capacity.workers;        for (self.workers[0..worker_count], 0..) |*worker, worker_index| {            worker.* = .{};            worker.victim_count = @intCast(@min(max_victims, worker_count));            const coprime = ShuffledIota.findAnotherCoprime(                worker_count,                @intCast((worker_index + 1) * 257 + worker_index * 13),            );            const shuffled = ShuffledIota.init(coprime);            worker.victims[0] = @intCast(worker_index);            for (1..worker.victim_count) |victim_index| {                worker.victims[victim_index] = @intCast(shuffled.next(                    worker.victims[victim_index - 1],                    worker_count,                ));            }        }    }    fn divideRange(self: *ThreadPool, begin: u64, end: u64) void {        for (self.workers[0..self.numWorkers()], 0..) |*worker, index| {            const range = workerRange(                0,                end - begin,                self.numWorkers(),                index,            ) catch unreachable;            worker.begin.store(@intCast(range.begin), .monotonic);            worker.end = @intCast(range.end);        }    }    fn runWorker(self: *ThreadPool, worker_index: usize) void {        var task_count: usize = 0;        var stolen_count: usize = 0;        if (self.current_tasks <= self.numWorkers()) {            const offset = self.workers[0].begin.load(.monotonic) +                worker_index;            const end = self.workers[self.numWorkers() - 1].end;            if (offset < end) {                self.task_callback(                    self.task_context,                    self.current_begin + offset,                    worker_index,                );                task_count = 1;            }        } else {            const worker = &self.workers[worker_index];            for (worker.victims[0..worker.victim_count]) |victim_index| {                const victim = &self.workers[victim_index];                while (true) {                    const offset = victim.begin.fetchAdd(1, .monotonic);                    if (offset >= victim.end) break;                    self.task_callback(                        self.task_context,                        self.current_begin + offset,                        worker_index,                    );                    task_count += 1;                    stolen_count += @intFromBool(victim_index != worker_index);                }            }        }        self.workers[worker_index].last_tasks = task_count;        self.workers[worker_index].last_stolen = stolen_count;    }    fn wakeWorkers(        self: *ThreadPool,        work_available: bool,        wake_blocked: bool,    ) u32 {        self.work_available.store(work_available, .monotonic);        const epoch = self.epoch.load(.monotonic) +% 1;        self.epoch.store(epoch, .release);        for (self.workers[1..self.numWorkers()]) |*worker| {            worker.wait_epoch.store(epoch, .release);        }        if (wake_blocked and self.capacity.threads != 0) {            wait.wakeAll(&self.workers[1].wait_epoch);        }        return epoch;    }    fn waitForWorkers(self: *ThreadPool, epoch: u32) void {        spin.callWithSpin(            self.config().spin_type,            BarrierContext{ .pool = self, .epoch = epoch },            waitAtBarrier,        );    }    fn waitForWork(        self: *ThreadPool,        worker_index: usize,        observed_epoch: u32,    ) ?u32 {        while (true) {            if (self.exiting.load(.acquire)) return null;            const config_value = self.config();            switch (config_value.wait_type) {                .block => {                    const next_epoch = wait.blockUntilDifferent(                        observed_epoch,                        &self.workers[1].wait_epoch,                    );                    if (self.exiting.load(.acquire)) return null;                    return next_epoch;                },                .spin_shared, .spin_separate => {                    const watched = if (config_value.wait_type ==                        .spin_separate)                        &self.workers[worker_index].wait_epoch                    else                        &self.epoch;                    var next_epoch: u32 = observed_epoch;                    spin.callWithSpin(                        config_value.spin_type,                        SpinWaitContext{                            .previous = observed_epoch,                            .watched = watched,                            .next = &next_epoch,                        },                        waitUntilDifferent,                    );                    if (self.exiting.load(.acquire)) return null;                    return next_epoch;                },            }        }    }    fn workerStarted(self: *ThreadPool) void {        self.synchronization.lock();        self.started_threads += 1;        self.done.broadcast();        self.synchronization.unlock();    }    fn workerReached(self: *ThreadPool, worker_index: usize, epoch: u32) void {        self.workers[worker_index].barrier_epoch.store(epoch, .release);    }    fn stopSpawned(self: *ThreadPool, spawned: usize) void {        self.signalExit();        for (self.handles[0..spawned]) |handle| handle.join();        self.exiting.store(false, .monotonic);        self.epoch.store(0, .monotonic);        self.started_threads = 0;        self.address = 0;        self.work_available.store(false, .monotonic);        self.configureWorkers();    }    fn signalExit(self: *ThreadPool) void {        self.exiting.store(true, .release);        const epoch = self.epoch.load(.monotonic) +% 1;        self.epoch.store(epoch, .release);        for (self.workers[1..self.numWorkers()]) |*worker| {            worker.wait_epoch.store(epoch, .release);        }        if (self.capacity.threads != 0) {            wait.wakeAll(&self.workers[1].wait_epoch);        }    }    fn requireStableAddress(self: *const ThreadPool) error{PoolMoved}!void {        if (self.address != @intFromPtr(self)) return error.PoolMoved;    }    fn recordRun(        self: *ThreadPool,        caller: Caller,        tasks: usize,        threaded: bool,        stolen_tasks: usize,        elapsed_ns: u64,    ) void {        self.synchronization.lock();        defer self.synchronization.unlock();        self.stats.runs +|= 1;        self.stats.serial_runs +|= @intFromBool(!threaded);        self.stats.threaded_runs +|= @intFromBool(threaded);        self.stats.tasks +|= tasks;        self.stats.stolen_tasks +|= stolen_tasks;        self.stats.elapsed_ns +|= elapsed_ns;        const caller_index = if (caller.index < max_callers)            caller.index        else            0;        const stats = &self.caller_stats[caller_index];        stats.runs +|= 1;        stats.tasks +|= tasks;        stats.workers +|= if (threaded) self.numWorkers() else 1;        stats.elapsed_ns +|= elapsed_ns;    }    fn notifyAutotune(self: *ThreadPool, elapsed_ns: u64) void {        const mode = self.waitMode();        const selected = blk: {            const tuner_value = self.tuner(mode);            if (tuner_value.best()) |best| break :blk best.*;            tuner_value.notifyCost(@max(elapsed_ns, 1));            break :blk if (tuner_value.best()) |best|                best.*            else                tuner_value.nextConfig().*;        };        self.config_bits.store(selected.encode(), .release);    }    fn tuner(self: *ThreadPool, mode: PoolWaitMode) *AutoTuner {        return switch (mode) {            .block => &self.block_tuner,            .spin => &self.spin_tuner,        };    }    fn tunerConst(self: *const ThreadPool, mode: PoolWaitMode) *const AutoTuner {        return switch (mode) {            .block => &self.block_tuner,            .spin => &self.spin_tuner,        };    }};const WorkerContext = struct {    pool: *ThreadPool,    worker_index: usize,};const SpinWaitContext = struct {    previous: u32,    watched: *const std.atomic.Value(u32),    next: *u32,};fn waitUntilDifferent(context: SpinWaitContext, policy: anytype) void {    context.next.* = policy.untilDifferent(        context.previous,        context.watched,    ).value;}const BarrierContext = struct {    pool: *ThreadPool,    epoch: u32,};fn waitAtBarrier(context: BarrierContext, policy: anytype) void {    for (context.pool.workers[1..context.pool.numWorkers()]) |*worker| {        _ = policy.untilEqual(context.epoch, &worker.barrier_epoch);    }}fn workerMain(context: WorkerContext) void {    const pool = context.pool;    pool.workerStarted();    var observed_epoch: u32 = 0;    while (pool.waitForWork(context.worker_index, observed_epoch)) |epoch| {        if (pool.work_available.load(.acquire)) {            pool.runWorker(context.worker_index);        }        pool.workerReached(context.worker_index, epoch);        observed_epoch = epoch;    }}fn nameWorker(handle: sys.thread.JoinHandle, index: usize) void {    if (comptime std.Thread.max_name_len < 9) return;    var storage: [std.Thread.max_name_len]u8 = undefined;    const name_value = std.fmt.bufPrint(        &storage,        "worker{d:0>3}",        .{index},    ) catch return;    handle.setName(name_value) catch {};}fn CallbackAdapter(    comptime Context: type,    comptime body: fn (Context, u64, usize) void,) type {    return struct {        fn call(            opaque_context: *anyopaque,            task: u64,            worker: usize,        ) void {            const context: *Context = @ptrCast(@alignCast(opaque_context));            body(context.*, task, worker);        }    };}fn nowNanoseconds() u64 {    return @intCast(@max(sys.time.nanoTimestamp(), 0));}fn elapsedNanoseconds(started: u64) u64 {    return nowNanoseconds() -| started;}test "Highway shuffled iota detects coprimes and exact permutations" {    for (1..40) |size_usize| {        const size: u32 = @intCast(size_usize);        const coprime = ShuffledIota.findAnotherCoprime(size, 1);        try std.testing.expect(ShuffledIota.coprimeNonzero(coprime, size));        const shuffled = ShuffledIota.init(coprime);        for (0..size) |start_usize| {            var visited: [40]u8 = @splat(0);            var current: u32 = @intCast(start_usize);            for (0..size) |_| {                visited[current] += 1;                current = shuffled.next(current, size);            }            for (visited[0..size]) |count| {                try std.testing.expectEqual(@as(u8, 1), count);            }        }    }    for (1..500) |value| {        try std.testing.expect(ShuffledIota.coprimeNonzero(1, @intCast(value)));        try std.testing.expect(ShuffledIota.coprimeNonzero(@intCast(value), 1));    }}test "Highway binary coprime agrees for powers products and primes" {    for (1..20) |i| {        const a = @as(u32, 1) << @intCast(i);        for (1..20) |j| {            const b = @as(u32, 1) << @intCast(j);            try std.testing.expect(!ShuffledIota.coprimeNonzero(a, b));        }    }    for (1..30) |i| {        const power = @as(u32, 1) << @intCast(i);        try std.testing.expect(ShuffledIota.coprimeNonzero(power, power + 1));        try std.testing.expect(ShuffledIota.coprimeNonzero(power, power - 1));        try std.testing.expect(ShuffledIota.coprimeNonzero(power + 1, power));        try std.testing.expect(ShuffledIota.coprimeNonzero(power - 1, power));    }    var random_state: u32 = 0x4f1b_c3d9;    for (0..5_000) |_| {        random_state = random_state *% 1_664_525 +% 1_013_904_223;        const x = (random_state & 0xfff7) + 2;        random_state = random_state *% 1_664_525 +% 1_013_904_223;        const y = (random_state & 0xfff7) + 2;        const product = x * y;        try std.testing.expect(!ShuffledIota.coprimeNonzero(product, x));        try std.testing.expect(!ShuffledIota.coprimeNonzero(product, y));        try std.testing.expect(!ShuffledIota.coprimeNonzero(x, product));        try std.testing.expect(!ShuffledIota.coprimeNonzero(y, product));    }    const primes = [_]u32{        2,   3,   5,   7,   11,  13,  17,  19,  23,  29,  31,  37,        41,  43,  47,  53,  59,  61,  67,  71,  73,  79,  83,  89,        97,  101, 103, 107, 109, 113, 127, 131, 137, 139, 149, 151,        157, 163, 167, 173, 179, 181, 191, 193, 197, 199, 211, 223,        227, 229, 233, 239, 241, 251, 257, 263, 269, 271,    };    for (primes, 0..) |a, i| {        for (primes[i + 1 ..]) |b| {            try std.testing.expect(ShuffledIota.coprimeNonzero(a, b));            try std.testing.expect(ShuffledIota.coprimeNonzero(b, a));        }    }}test "Highway independent shuffles retain coverage with bounded contention" {    var shuffles: [40]ShuffledIota = @splat(.{});    var current: [40]u32 = @splat(0);    var visited_all: [40]u8 = @splat(0);    for (1..40) |size_usize| {        const size: u32 = @intCast(size_usize);        @memset(visited_all[0..size], 0);        for (0..size) |index_usize| {            const index: u32 = @intCast(index_usize);            shuffles[index] = .{ .coprime = ShuffledIota.findAnotherCoprime(                size,                (index + 1) * 257 + index * 13,            ) };            current[index] = index;        }        var bad_steps: usize = 0;        for (0..size) |_| {            var visited: [40]u8 = @splat(0);            for (current[0..size]) |value| {                visited[value] += 1;                visited_all[value] = 1;            }            var contended: usize = 0;            var maximum: u8 = 0;            for (visited[0..size]) |count| {                contended += @intFromBool(count > 1);                maximum = @max(maximum, count);            }            const expected: usize = @intFromFloat(                std.math.sqrt(@as(f32, @floatFromInt(size))) * 2.0,            );            bad_steps += @intFromBool(contended > expected and maximum > 3);            for (current[0..size], 0..) |*value, index| {                value.* = shuffles[index].next(value.*, size);            }        }        for (visited_all[0..size]) |visited| {            try std.testing.expectEqual(@as(u8, 1), visited);        }        try std.testing.expect(bad_steps < 4);    }}test "Highway pool configurations preserve layout and candidate order" {    try std.testing.expectEqual(@as(usize, 4), @sizeOf(Config));    var storage: [max_configs]Config = undefined;    const block = Config.candidates(.block, &storage);    try std.testing.expectEqual(@as(usize, 1), block.len);    try std.testing.expectEqual(WaitType.block, block[0].wait_type);    try std.testing.expectEqual(spin.SpinType.pause, block[0].spin_type);    const block_config = block[0];    const detected = spin.detectSpin(0);    const candidates = Config.candidates(.spin, &storage);    const expected_count: usize = if (detected == .pause) 2 else 4;    try std.testing.expectEqual(expected_count, candidates.len);    for (candidates, 0..) |candidate, index| {        try std.testing.expect(candidate.wait_type != .block);        try std.testing.expectEqual(            if (index & 1 == 0)                WaitType.spin_shared            else                WaitType.spin_separate,            candidate.wait_type,        );        try std.testing.expectEqual(            if (index < 2) detected else spin.SpinType.pause,            candidate.spin_type,        );    }    var name_storage: [64]u8 = undefined;    const formatted = try block_config.formatName(&name_storage);    try std.testing.expect(std.mem.startsWith(u8, formatted, "Pause"));    try std.testing.expect(std.mem.endsWith(u8, formatted, "Block    "));}test "Highway worker range division is balanced and contiguous" {    for (1..40) |worker_count| {        for (0..80) |task_count| {            var cursor: u64 = 11;            var minimum: usize = std.math.maxInt(usize);            var maximum: usize = 0;            for (0..worker_count) |worker_index| {                const range = try workerRange(                    11,                    11 + task_count,                    worker_count,                    worker_index,                );                try std.testing.expectEqual(cursor, range.begin);                const count: usize = @intCast(range.end - range.begin);                minimum = @min(minimum, count);                maximum = @max(maximum, count);                cursor = range.end;            }            try std.testing.expectEqual(11 + task_count, cursor);            try std.testing.expect(maximum - minimum <= 1);        }    }}test "Highway pool worker mapping preserves local and cluster indices" {    const local = PoolWorkerMapping{};    try std.testing.expectEqual(@as(usize, 7), local.globalIndex(7));    const cluster = try PoolWorkerMapping.init(3, 8);    try std.testing.expectEqual(@as(usize, 27), cluster.globalIndex(3));    const across = try PoolWorkerMapping.init(all_clusters, 8);    try std.testing.expectEqual(@as(usize, 24), across.globalIndex(3));    try std.testing.expectError(        error.InvalidCluster,        PoolWorkerMapping.init(all_clusters + 1, 1),    );    try std.testing.expectError(        error.EmptyCluster,        PoolWorkerMapping.init(0, 0),    );    var pool = ThreadPool.init(3, cluster);    defer pool.deinit();    for (0..pool.numWorkers()) |worker| {        try std.testing.expectEqual(            cluster.globalIndex(worker),            try pool.globalWorkerIndex(worker),        );    }    try std.testing.expectError(        error.WorkerOutOfBounds,        pool.globalWorkerIndex(pool.numWorkers()),    );}const HitContext = struct {    begin: u64,    hits: []std.atomic.Value(u32),    worker_bits: *std.atomic.Value(usize),};fn recordHit(context: *HitContext, task: u64, worker: usize) void {    _ = context.hits[@intCast(task - context.begin)].fetchAdd(1, .monotonic);    _ = context.worker_bits.fetchOr(        @as(usize, 1) << @intCast(worker),        .monotonic,    );}test "Highway thread pool runs every task once in block and spin modes" {    if (!topology.haveThreadingSupport() or        !sys.thread.threadsSupported())    {        return error.SkipZigTest;    }    var pool = ThreadPool.init(@min(ThreadPool.maxThreads(), 6), .{});    defer pool.deinit();    var hits: [97]std.atomic.Value(u32) =        @splat(std.atomic.Value(u32).init(0));    var worker_bits = std.atomic.Value(usize).init(0);    var context = HitContext{        .begin = 23,        .hits = &hits,        .worker_bits = &worker_bits,    };    for ([_]PoolWaitMode{ .spin, .block }) |mode| {        try pool.setWaitMode(mode);        for (&hits) |*hit| hit.store(0, .monotonic);        worker_bits.store(0, .monotonic);        try pool.run(23, 120, &context, recordHit);        for (&hits) |*hit| {            try std.testing.expectEqual(                @as(u32, 1),                hit.load(.monotonic),            );        }        try std.testing.expect(worker_bits.load(.monotonic) != 0);    }}test "Highway thread pool preserves ranges near the u64 limit" {    if (!topology.haveThreadingSupport() or        !sys.thread.threadsSupported() or ThreadPool.maxThreads() == 0)    {        return error.SkipZigTest;    }    const begin = std.math.maxInt(u64) - 17;    var pool = ThreadPool.init(@min(ThreadPool.maxThreads(), 2), .{});    defer pool.deinit();    var hits: [17]std.atomic.Value(u32) =        @splat(std.atomic.Value(u32).init(0));    var worker_bits = std.atomic.Value(usize).init(0);    var context = HitContext{        .begin = begin,        .hits = &hits,        .worker_bits = &worker_bits,    };    try pool.run(begin, std.math.maxInt(u64), &context, recordHit);    for (&hits) |*hit| {        try std.testing.expectEqual(@as(u32, 1), hit.load(.monotonic));    }}const AssignmentContext = struct {    expected_tasks: usize,    expected_workers: usize,    calls: *std.atomic.Value(usize),    worker_bits: *std.atomic.Value(usize),};fn recordAssignment(    context: *AssignmentContext,    task: u64,    worker: usize,) void {    std.debug.assert(task < context.expected_tasks);    std.debug.assert(worker < context.expected_workers);    _ = context.calls.fetchAdd(1, .monotonic);    _ = context.worker_bits.fetchOr(        @as(usize, 1) << @intCast(worker),        .monotonic,    );}test "Highway small assignments retain valid worker identities" {    for ([_]usize{ 0, 1, 3, 5, 8 }) |requested| {        var pool = ThreadPool.init(@min(ThreadPool.maxThreads(), requested), .{});        defer pool.deinit();        for (1..3) |multiplier| {            const tasks = pool.numWorkers() * multiplier;            var calls = std.atomic.Value(usize).init(0);            var worker_bits = std.atomic.Value(usize).init(0);            var context = AssignmentContext{                .expected_tasks = tasks,                .expected_workers = pool.numWorkers(),                .calls = &calls,                .worker_bits = &worker_bits,            };            try pool.run(0, tasks, &context, recordAssignment);            try std.testing.expectEqual(tasks, calls.load(.monotonic));            try std.testing.expect(                @popCount(worker_bits.load(.monotonic)) <= pool.numWorkers(),            );        }    }}const SumContext = struct {    counters: []std.atomic.Value(usize),};fn addTask(context: *SumContext, task: u64, worker: usize) void {    _ = context.counters[worker].fetchAdd(@intCast(task), .monotonic);}test "Highway pool wait modes preserve the task sum and caller statistics" {    if (!topology.haveThreadingSupport() or        !sys.thread.threadsSupported())    {        return error.SkipZigTest;    }    const caller = try addCaller("thread-pool-test");    try std.testing.expectEqual(        caller.id(),        (try addCaller("thread-pool-test")).id(),    );    try std.testing.expectEqualStrings(        "thread-pool-test",        callerName(caller).?,    );    var pool = ThreadPool.init(@min(ThreadPool.maxThreads(), 9), .{});    defer pool.deinit();    var counters: [max_workers]std.atomic.Value(usize) =        @splat(std.atomic.Value(usize).init(0));    var context = SumContext{ .counters = &counters };    const task_count = pool.numWorkers() * 19;    for ([_]PoolWaitMode{ .spin, .block }) |mode| {        try pool.setWaitMode(mode);        for (counters[0..pool.numWorkers()]) |*counter| {            counter.store(0, .monotonic);        }        try pool.runWithCaller(            0,            task_count,            caller,            &context,            addTask,        );        var actual: usize = 0;        for (counters[0..pool.numWorkers()]) |*counter| {            actual += counter.load(.monotonic);        }        try std.testing.expectEqual(            task_count * (task_count - 1) / 2,            actual,        );    }    const caller_stats = pool.callerStats(caller).?;    try std.testing.expectEqual(@as(u64, 2), caller_stats.runs);    try std.testing.expectEqual(@as(u64, task_count * 2), caller_stats.tasks);    try std.testing.expectEqual(@as(u64, 2), pool.statsSnapshot().threaded_runs);    try pool.resetStats();    try std.testing.expectEqual(PoolStats{}, pool.statsSnapshot());    try std.testing.expectEqual(CallerStats{}, pool.callerStats(caller).?);}const NestedContext = struct {    begin: u64,    end: u64,    hits: []std.atomic.Value(u32),    inner_calls: *std.atomic.Value(usize),};fn innerTask(context: *NestedContext, task: u64, worker: usize) void {    std.debug.assert(worker == 0);    std.debug.assert(context.begin <= task);    std.debug.assert(task < context.end);    _ = context.inner_calls.fetchAdd(1, .monotonic);}fn outerTask(context: *NestedContext, task: u64, _: usize) void {    std.debug.assert(context.begin <= task);    std.debug.assert(task < context.end);    _ = context.hits[@intCast(task - context.begin)].fetchAdd(1, .monotonic);    var inner = ThreadPool.init(0, .{});    defer inner.deinit();    inner.run(context.begin, context.end, context, innerTask) catch unreachable;}test "Highway pools reuse shifted ranges and allow nested serial runs" {    var hits: [20]std.atomic.Value(u32) =        @splat(std.atomic.Value(u32).init(0));    var inner_calls = std.atomic.Value(usize).init(0);    for ([_]usize{ 0, 3, 6 }) |requested| {        var pool = ThreadPool.init(@min(ThreadPool.maxThreads(), requested), .{});        defer pool.deinit();        for ([_]PoolWaitMode{ .spin, .block }) |mode| {            try pool.setWaitMode(mode);            for (0..20) |task_count| {                for (0..8) |begin| {                    for (&hits) |*hit| hit.store(0, .monotonic);                    inner_calls.store(0, .monotonic);                    var context = NestedContext{                        .begin = begin,                        .end = begin + task_count,                        .hits = &hits,                        .inner_calls = &inner_calls,                    };                    try pool.run(                        context.begin,                        context.end,                        &context,                        outerTask,                    );                    for (hits[0..task_count]) |*hit| {                        try std.testing.expectEqual(                            @as(u32, 1),                            hit.load(.monotonic),                        );                    }                    try std.testing.expectEqual(                        task_count * task_count,                        inner_calls.load(.monotonic),                    );                }            }        }    }}fn emptyTask(_: void, _: u64, _: usize) void {}test "Highway live pools switch wait modes repeatedly" {    if (!topology.haveThreadingSupport() or        !sys.thread.threadsSupported() or ThreadPool.maxThreads() == 0)    {        return error.SkipZigTest;    }    var pool = ThreadPool.init(@min(ThreadPool.maxThreads(), 9), .{});    defer pool.deinit();    try pool.run(0, 2, {}, emptyTask);    for (0..100) |iteration| {        const mode: PoolWaitMode = if ((iteration * 17 + 5) & 1 == 0)            .spin        else            .block;        try pool.setWaitMode(mode);        try std.testing.expectEqual(mode, pool.waitMode());        try std.testing.expectEqual(            mode == .block,            pool.config().wait_type == .block,        );    }    try pool.setWaitMode(.block);}test "Highway pool autotuning visits candidates and converges" {    if (!topology.haveThreadingSupport() or        !sys.thread.threadsSupported() or ThreadPool.maxThreads() == 0)    {        return error.SkipZigTest;    }    var pool = ThreadPool.init(@min(ThreadPool.maxThreads(), 2), .{});    defer pool.deinit();    try pool.setWaitMode(.spin);    var runs: usize = 0;    while (!pool.autoTuneComplete() and runs < max_configs * 64) : (runs += 1) {        try pool.run(0, 2, {}, emptyTask);    }    try std.testing.expect(pool.autoTuneComplete());    try std.testing.expect(runs >= 30);    const costs = pool.autoTuneCosts();    try std.testing.expect(costs.len == 2 or costs.len == 4);    for (costs) |*cost| {        try std.testing.expect(            cost.bufferedCount() != 0 or cost.onlineCount() != 0,        );    }}test "Highway pool clamps capacity and rejects invalid ranges" {    var pool = ThreadPool.init(max_threads + 100, .{});    defer pool.deinit();    try std.testing.expect(pool.wasClamped());    try std.testing.expectEqual(max_workers, pool.numWorkers());    try std.testing.expectError(        error.InvalidRange,        pool.run(2, 1, {}, struct {            fn call(_: void, _: u64, _: usize) void {}        }.call),    );    try pool.resetStats();    try std.testing.expectEqual(PoolStats{}, pool.statsSnapshot());}

Source: lib/simd/src/thread/root.zig:4

zig
pub const pool = @import("pool.zig");

Audit

Definitions7
Public names7
Members3
Version26.7.0
Revisiondaab053ee433