You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

C#自定义ThreadPool性能异常:运行测试时笔记本卡顿严重

自定义任务窃取线程池导致设备卡顿的问题分析

我为学习需求实现了一个带任务窃取机制的C#自定义ThreadPool,给定了不可修改的IThreadPool接口、WorkStealingQueue类、DotNetThreadPoolWrapper类及测试用例。运行测试时笔记本出现严重卡顿甚至死机,请求排查代码问题。

要求的算法逻辑:

  • 外部线程将任务放入共享队列;
  • 工作线程将任务放入本地队列;
    工作线程获取任务顺序:
  1. 优先从自身本地队列获取;
  2. 其次从共享队列获取;
  3. 最后尝试窃取其他工作线程的本地队列任务;
  4. 无任务时等待新任务信号。

本人实现的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.29 12:35:55