C#自定义ThreadPool性能异常:运行测试时笔记本卡顿严重
我为学习需求实现了一个带任务窃取机制的C#自定义ThreadPool,给定了不可修改的IThreadPool接口、WorkStealingQueue类、DotNetThreadPoolWrapper类及测试用例。运行测试时笔记本出现严重卡顿甚至死机,请求排查代码问题。
要求的算法逻辑:
- 外部线程将任务放入共享队列;
- 工作线程将任务放入本地队列;
工作线程获取任务顺序:
- 优先从自身本地队列获取;
- 其次从共享队列获取;
- 最后尝试窃取其他工作线程的本地队列任务;
- 无任务时等待新任务信号。
本人实现的ThreadPoolCustom代码
public class ThreadPoolCustom : IThreadPool { private readonly ConcurrentQueue<Action> _globalQueue = new ConcurrentQueue<Action>(); private readonly WorkStealingQueue<Action>[] _localQueues; private readonly List<Thread> _workers = new List<Thread>(); private volatile int _counter = 0; private object _lock = new object(); public ThreadPoolCustom(int workerCount) { _localQueues = new WorkStealingQueue<Action>[workerCount]; for (int i = 0; i < workerCount; i++) { _localQueues[i] = new WorkStealingQueue<Action>(); var worker = new Thread(Worker); worker.IsBackground = true; worker.Start(i); _workers.Add(worker); } } public ThreadPoolCustom() : this(Environment.ProcessorCount) { } public void EnqueueAction(Action action) { // tasks in globak queue _globalQueue.Enqueue(action); // local queue foreach (var localQueue in _localQueues) { localQueue.LocalPush(action); } lock (_lock) { Monitor.PulseAll(_lock); } } public long GetTasksProcessedCount() { return _counter; } private void Worker(object data) { var workerId = (int)data; while (true) { Action task = null; // Trying to retrieve a task from the local queue if (_localQueues[workerId].LocalPop(ref task)) { ExecuteTask(task); continue; } // Trying to retrieve a task from the global queue if (TryDequeueFromGlobalQueue(ref task)) { ExecuteTask(task); continue; } // Trying to Steal a task from the other loval queues if (TryStealFromOtherQueues(workerId, ref task)) { ExecuteTask(task); continue; } // If there are no tasks, I wait for a signal to appear in other queues. lock (_lock) { Monitor.Wait(_lock); } } } private bool TryDequeueFromGlobalQueue(ref Action task) { return _globalQueue.TryDequeue(out task); } private bool TryStealFromOtherQueues(int workerId, ref Action task) { // Going to every local queue except my own, trying to steal for (int i = 0; i < _localQueues.Length; i++) { if (i == workerId) continue; if (_localQueues[i]?.TrySteal(ref task) ?? false) { return true; } } return false; } private void ExecuteTask(Action task) { task.Invoke(); Interlocked.Increment(ref _counter); } }
不可修改的接口与类代码
IThreadPool接口
public interface IThreadPool { void EnqueueAction(Action action); long GetTasksProcessedCount(); }
DotNetThreadPoolWrapper类
public class DotNetThreadPoolWrapper : IThreadPool { private long processedTask = 0L; public void EnqueueAction(Action action) { ThreadPool.UnsafeQueueUserWorkItem(delegate { action.Invoke(); Interlocked.Increment(ref processedTask); }, null); } public long GetTasksProcessedCount() => processedTask; }
WorkStealingQueue类
public class WorkStealingQueue<T> { private const int INITIAL_SIZE = 32; private T[] m_array = new T[INITIAL_SIZE]; private int m_mask = INITIAL_SIZE - 1; private volatile int m_headIndex = 0; private volatile int m_tailIndex = 0; private readonly object m_foreignLock = new object(); public bool IsEmpty => m_headIndex >= m_tailIndex; public int Count => m_tailIndex - m_headIndex; public void LocalPush(T obj) { var tail = m_tailIndex; if(tail < m_headIndex + m_mask) { m_array[tail & m_mask] = obj; m_tailIndex = tail + 1; } else { lock (m_foreignLock) { var head = m_headIndex; var count = m_tailIndex - m_headIndex; if(count >= m_mask) { var newArray = new T[m_array.Length << 1]; for(var i = 0; i < m_array.Length; i++) { newArray[i] = m_array[(i + head) & m_mask]; } m_array = newArray; m_headIndex = 0; m_tailIndex = tail = count; m_mask = (m_mask << 1) | 1; } m_array[tail & m_mask] = obj; m_tailIndex = tail + 1; } } } public bool LocalPop(ref T obj) { var tail = m_tailIndex; if(m_headIndex >= tail) { return false; } tail -= 1; Interlocked.Exchange(ref m_tailIndex, tail); if(m_headIndex <= tail) { obj = m_array[tail & m_mask]; return true; } else { lock (m_foreignLock) { if(m_headIndex <= tail) { obj = m_array[tail & m_mask]; return true; } else { m_tailIndex = tail + 1; return false; } } } } public bool TrySteal(ref T obj) { var taken = false; try { taken = Monitor.TryEnter(m_foreignLock); if(taken) { var head = m_headIndex; Interlocked.Exchange(ref m_headIndex, head + 1); if(head < m_tailIndex) { obj = m_array[head & m_mask]; return true; } else { m_headIndex = head; return false; } } } finally { if(taken) { Monitor.Exit(m_foreignLock); } } return false; } }
ThreadPoolTests测试类
public class ThreadPoolTests { public static void Run<TThreadPool>() where TThreadPool : IThreadPool, new() { Run(() => new TThreadPool()); } public static void Run(Func<IThreadPool> threadPoolFactory) { var name = threadPoolFactory().GetType().Name.Replace("ThreadPool", "", StringComparison.OrdinalIgnoreCase); Console.WriteLine($"----------======={name} ThreadPool tests=======----------"); RunTest(LongCalculations); RunTest(ShortCalculations); RunTest(ExtremelyShortCalculations); RunTest(InnerShortCalculations); RunTest(InnerExtremelyShortCalculations); Console.WriteLine("\n"); void RunTest(Action<IThreadPool> test) => test(threadPoolFactory()); } private static void LongCalculations(IThreadPool threadPool) { Console.Write("LongCalculations test: "); var timer = Stopwatch.StartNew(); long enqueueMs; const int actionsCount = 1 * 1000; using(var cev = new CountdownEvent(actionsCount)) { Action sumAction = () => { cev.Signal(); Thread.SpinWait(1000 * 1000); }; for(int i = 0; i < actionsCount; i++) { threadPool.EnqueueAction(sumAction); } enqueueMs = timer.ElapsedMilliseconds; cev.Wait(); } timer.Stop(); Console.WriteLine($" total {timer.ElapsedMilliseconds} ms, enqueue {enqueueMs} ms [tasks processed ~{threadPool.GetTasksProcessedCount()}]"); } private static void ShortCalculations(IThreadPool threadPool) { Console.Write("ShortCalculations test: "); var timer = Stopwatch.StartNew(); long enqueueMs; const int actionsCount = 1 * 1000 * 1000; using(var cev = new CountdownEvent(actionsCount)) { Action sumAction = () => { cev.Signal(); Thread.SpinWait(1000); }; for(var i = 0; i < actionsCount; i++) { threadPool.EnqueueAction(sumAction); } enqueueMs = timer.ElapsedMilliseconds; cev.Wait(); } timer.Stop(); Console.WriteLine($" total {timer.ElapsedMilliseconds} ms, enqueue {enqueueMs} ms [tasks processed ~{threadPool.GetTasksProcessedCount()}]"); } private static void ExtremelyShortCalculations(IThreadPool threadPool) { Console.Write("ExtremelyShortCalculations test: "); var timer = Stopwatch.StartNew(); long enqueueMs; const int actionsCount = 1 * 1000 * 1000; using(var cev = new CountdownEvent(actionsCount)) { Action sumAction = () => { cev.Signal(); }; for(int i = 0; i < actionsCount; i++) { threadPool.EnqueueAction(sumAction); } enqueueMs = timer.ElapsedMilliseconds; cev.Wait(); } timer.Stop(); Console.WriteLine($" total {timer.ElapsedMilliseconds} ms, enqueue {enqueueMs} ms [tasks processed ~{threadPool.GetTasksProcessedCount()}]"); } private static void InnerShortCalculations(IThreadPool threadPool) { Console.Write("InnerCalculations test: "); var timer = Stopwatch.StartNew(); long enqueueMs; const int actionsCount = 1 * 1000; const int subactionsCount = 1 * 1000; using(CountdownEvent outerEvent = new CountdownEvent(actionsCount)) using(CountdownEvent innerEvent = new CountdownEvent(actionsCount * subactionsCount)) { Action innerAction = () => { innerEvent.Signal(); Thread.SpinWait(1000); }; Action outerAction = () => { for(int i = 0; i < subactionsCount; i++) { threadPool.EnqueueAction(innerAction); } outerEvent.Signal(); }; for(int i = 0; i < actionsCount; i++) { threadPool.EnqueueAction(outerAction); } outerEvent.Wait(); enqueueMs = timer.ElapsedMilliseconds; innerEvent.Wait(); } timer.Stop(); Console.WriteLine($" total {timer.ElapsedMilliseconds} ms, enqueue {enqueueMs} ms [tasks processed ~{threadPool.GetTasksProcessedCount()}]"); } private static void InnerExtremelyShortCalculations(IThreadPool threadPool) { Console.Write("InnerExtremelyShortCalculations test: "); var timer = Stopwatch.StartNew(); long enqueueMs; const int actionsCount = 1 * 1000; const int subactionsCount = 1 * 1000; using(CountdownEvent outerEvent = new CountdownEvent(actionsCount)) using(CountdownEvent innerEvent = new CountdownEvent(actionsCount * subactionsCount)) { Action innerAction = () => { innerEvent.Signal(); }; Action outerAction = () => { for(int i = 0; i < subactionsCount; i++) { threadPool.EnqueueAction(innerAction); } outerEvent.Signal(); }; for(int i = 0; i < actionsCount; i++) { threadPool.EnqueueAction(outerAction); } outerEvent.Wait(); enqueueMs = timer.ElapsedMilliseconds; innerEvent.Wait(); } timer.Stop(); Console.WriteLine($" total {timer.ElapsedMilliseconds} ms, enqueue {enqueueMs} ms [tasks processed ~{threadPool.GetTasksProcessedCount()}]"); } }
导致卡顿的核心问题分析
1. 任务数量爆炸式增长
在EnqueueAction方法中,你把同一个任务同时加入全局队列和所有工作线程的本地队列。比如测试中提交100万个任务,实际会被放入N+1个队列(N是工作线程数),总任务量变成100万*(N+1)。以8核CPU为例,总任务量会飙升到900万,直接导致内存占用暴增、CPU被完全占满,系统资源耗尽后出现卡顿甚至死机。
正确逻辑:外部线程提交的任务只放入全局队列;工作线程自身生成的任务(如测试中outerAction里的子任务)才放入当前线程的本地队列,而非所有本地队列。
2. 任务重复执行引发连锁问题
同一个任务被多个队列存储,会导致多个工作线程重复执行同一任务:
_counter计数远大于实际提交的任务数;CountdownEvent被超额触发,测试逻辑出现异常;- 额外的任务执行进一步消耗CPU资源。
3. 唤醒机制低效加剧资源消耗
每次提交任务都调用Monitor.PulseAll唤醒所有等待的工作线程,任务量较大时,会导致大量线程被频繁唤醒又立即进入等待状态,产生大量不必要的上下文切换,进一步加剧CPU资源消耗。
修复建议
修改EnqueueAction方法,仅将外部任务放入全局队列,去掉遍历所有本地队列添加任务的代码,同时将全量唤醒改为唤醒单个线程:
public void EnqueueAction(Action action) { // 仅将外部任务放入全局队列 _globalQueue.Enqueue(action); lock (_lock) { Monitor.Pulse(_lock); // 只唤醒一个线程,避免全量唤醒的资源浪费 } }
另外,若需要让工作线程生成的任务进入本地队列,需在任务执行时(如outerAction内部)调用针对当前线程本地队列的提交逻辑,而非复用EnqueueAction方法。
内容的提问来源于stack exchange,提问作者user20986100

