[CPU] Fix NUMA/core binding for DP ranks (#32856)
This commit is contained in:
@@ -1,10 +1,13 @@
|
||||
import os
|
||||
import re
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
kernel = torch.ops.sgl_kernel
|
||||
|
||||
from sglang.srt.utils.numa_utils import init_threads_binding
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -26,5 +29,47 @@ class TestBinding(CustomTestCase):
|
||||
self.assertEqual(bindings, expected_cores)
|
||||
|
||||
|
||||
class TestInitThreadsBinding(unittest.TestCase):
|
||||
"""Tests for init_threads_binding: NUMA/core selection by global rank.
|
||||
|
||||
numa_index is the worker's global device id (gpu_id) across all DP
|
||||
replicas; world_size is dp_size * tp_size * pp_size.
|
||||
"""
|
||||
|
||||
@patch(
|
||||
"sglang.srt.utils.numa_utils.get_cpu_ids_by_node",
|
||||
return_value=["0,1,2,3", "4,5,6,7", "8,9,10,11", "12,13,14,15"],
|
||||
)
|
||||
@patch.dict(os.environ, {"SGLANG_CPU_OMP_THREADS_BIND": "all"})
|
||||
def test_dp_ranks_get_distinct_numa_nodes(self, _mock_nodes):
|
||||
# dp_size=2, tp_size=2 -> world_size=4, numa_index 0..3 global rank.
|
||||
results = [init_threads_binding(numa_index=i, world_size=4) for i in range(4)]
|
||||
self.assertEqual(results, ["0,1,2,3", "4,5,6,7", "8,9,10,11", "12,13,14,15"])
|
||||
|
||||
@patch(
|
||||
"sglang.srt.utils.numa_utils.get_cpu_ids_by_node",
|
||||
return_value=["0,1,2,3", "4,5,6,7"],
|
||||
)
|
||||
@patch.dict(os.environ, {"SGLANG_CPU_OMP_THREADS_BIND": "0-1|4-5"})
|
||||
def test_explicit_bind_list_indexed_by_numa_index(self, _mock_nodes):
|
||||
self.assertEqual(init_threads_binding(numa_index=0, world_size=2), "0-1")
|
||||
self.assertEqual(init_threads_binding(numa_index=1, world_size=2), "4-5")
|
||||
|
||||
@patch(
|
||||
"sglang.srt.utils.numa_utils.get_cpu_ids_by_node",
|
||||
return_value=["0,1", "2,3", "4,5"],
|
||||
)
|
||||
@patch.dict(os.environ, {"SGLANG_CPU_OMP_THREADS_BIND": "0-1|2-3|4-5"})
|
||||
def test_router_worker_uses_global_numa_index(self, _mock_nodes):
|
||||
# Router mode: each worker is an independent dp_size=1 server, so
|
||||
# world_size is locally 1 even though the bind string has multiple
|
||||
# groups. numa_index is still the global gpu_id and must be able to
|
||||
# select any group, not just index 0.
|
||||
self.assertEqual(
|
||||
init_threads_binding(numa_index=2, world_size=1),
|
||||
"4-5",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user