[XPU] Adapt device agnostic API usage (#32093)

Co-authored-by: mingfeima <mingfei.ma@intel.com>
This commit is contained in:
ANSHUMAN TRIPATHY
2026-09-11 10:03:12 +08:00
committed by GitHub
co-authored by mingfeima
parent 67d3a2ea57
commit 2adb2e8485
4 changed files with 115 additions and 73 deletions
+85 -48
View File
@@ -39,9 +39,14 @@ from sglang.srt.debug_utils.dumper import (
get_tensor_info,
get_truncated_value,
)
from sglang.srt.utils import kill_process_tree
from sglang.srt.distributed.parallel_state import get_default_distributed_backend
from sglang.srt.utils import get_device, get_device_module, kill_process_tree
from sglang.srt.utils.common import temp_set_env
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.ci.ci_register import (
register_amd_ci,
register_cuda_ci,
register_xpu_ci,
)
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
@@ -52,6 +57,7 @@ from sglang.test.test_utils import (
register_cuda_ci(est_time=30, stage="nightly", runner_config="2-gpu-large")
register_amd_ci(est_time=60, suite="nightly-amd", nightly=True)
register_xpu_ci(est_time=400, suite="nightly-xpu-2-gpu", nightly=True)
@contextmanager
@@ -423,11 +429,15 @@ class TestDumperDistributed:
DUMPER_ENABLE="1",
DUMPER_DIR=str(tmp_path),
):
run_distributed_test(self._test_basic_func, tmpdir=str(tmp_path))
run_distributed_test(
self._test_basic_func,
tmpdir=str(tmp_path),
backend=get_default_distributed_backend(get_device()),
)
@staticmethod
def _test_basic_func(rank, tmpdir):
tensor = torch.randn(10, 10, device=f"cuda:{rank}")
tensor = torch.randn(10, 10, device=get_device(rank))
dumper.dump("tensor_a", tensor, arg=100)
dumper.step()
@@ -442,7 +452,7 @@ class TestDumperDistributed:
dumper.configure(filter=None)
dumper.step()
dumper.dump_dict("obj", {"a": torch.randn(3, device=f"cuda:{rank}"), "b": 42})
dumper.dump_dict("obj", {"a": torch.randn(3, device=get_device(rank)), "b": 42})
dumper.step()
dist.barrier()
@@ -455,7 +465,10 @@ class TestDumperDistributed:
def test_collective_timeout(self):
with temp_set_env(DUMPER_ENABLE="1"):
run_distributed_test(self._test_collective_timeout_func)
run_distributed_test(
self._test_collective_timeout_func,
backend=get_default_distributed_backend(get_device()),
)
@staticmethod
def _test_collective_timeout_func(rank):
@@ -483,11 +496,15 @@ class TestDumperDistributed:
DUMPER_ENABLE="1",
DUMPER_DIR=str(tmp_path),
):
run_distributed_test(self._test_file_content_func, tmpdir=str(tmp_path))
run_distributed_test(
self._test_file_content_func,
tmpdir=str(tmp_path),
backend=get_default_distributed_backend(get_device()),
)
@staticmethod
def _test_file_content_func(rank, tmpdir):
tensor = torch.arange(12, device=f"cuda:{rank}").reshape(3, 4).float()
tensor = torch.arange(12, device=get_device(rank)).reshape(3, 4).float()
dumper.dump("content_check", tensor)
dumper.step()
@@ -509,13 +526,17 @@ class TestDumperFileWriteControl:
DUMPER_DIR=str(tmp_path),
DUMPER_FILTER="name.startswith('keep')",
):
run_distributed_test(self._test_filter_func, tmpdir=str(tmp_path))
run_distributed_test(
self._test_filter_func,
tmpdir=str(tmp_path),
backend=get_default_distributed_backend(get_device()),
)
@staticmethod
def _test_filter_func(rank, tmpdir):
dumper.dump("keep_this", torch.randn(5, device=f"cuda:{rank}"))
dumper.dump("skip_this", torch.randn(5, device=f"cuda:{rank}"))
dumper.dump("not_keep_this", torch.randn(5, device=f"cuda:{rank}"))
dumper.dump("keep_this", torch.randn(5, device=get_device(rank)))
dumper.dump("skip_this", torch.randn(5, device=get_device(rank)))
dumper.dump("not_keep_this", torch.randn(5, device=get_device(rank)))
dumper.step()
dist.barrier()
@@ -531,11 +552,17 @@ class TestDumperFileWriteControl:
DUMPER_ENABLE="1",
DUMPER_DIR=str(tmp_path),
):
run_distributed_test(self._test_save_false_func, tmpdir=str(tmp_path))
run_distributed_test(
self._test_save_false_func,
tmpdir=str(tmp_path),
backend=get_default_distributed_backend(get_device()),
)
@staticmethod
def _test_save_false_func(rank, tmpdir):
dumper.dump("no_save_tensor", torch.randn(5, device=f"cuda:{rank}"), save=False)
dumper.dump(
"no_save_tensor", torch.randn(5, device=get_device(rank)), save=False
)
dumper.step()
dist.barrier()
@@ -1601,7 +1628,11 @@ class TestZmqPortIsolation:
thread = threading.Thread(
target=run_distributed_test,
args=(_dumper_worker,),
kwargs={"http_port": port, "stop_event": stop_event},
kwargs={
"http_port": port,
"stop_event": stop_event,
"backend": get_default_distributed_backend(get_device()),
},
)
thread.start()
threads.append(thread)
@@ -1638,7 +1669,11 @@ class TestDumperHttp:
thread = threading.Thread(
target=run_distributed_test,
args=(_dumper_worker,),
kwargs={"http_port": http_port, "stop_event": stop_event},
kwargs={
"http_port": http_port,
"stop_event": stop_event,
"backend": get_default_distributed_backend(get_device()),
},
)
thread.start()
try:
@@ -3222,9 +3257,9 @@ def _run_graft_test(worker_func, **kwargs):
def _graft_worker_entry(rank, role_port, worker_func, result_queue, kwargs):
import traceback
torch.cuda.set_device(rank)
get_device_module().set_device(rank)
dist.init_process_group(
backend="nccl",
backend=get_default_distributed_backend(get_device()),
init_method=f"tcp://127.0.0.1:{role_port}",
world_size=1,
rank=0,
@@ -3312,9 +3347,9 @@ def _graft_split_worker_entry(
config=_dumper_module.DumperConfig.from_env()
)
torch.cuda.set_device(global_rank)
get_device_module().set_device(global_rank)
dist.init_process_group(
backend="nccl",
backend=get_default_distributed_backend(get_device()),
init_method=f"tcp://127.0.0.1:{role_port}",
world_size=1,
rank=0,
@@ -3432,6 +3467,7 @@ def _make_grafter_test_config(
role = "baseline" if rank == 0 else "target"
return DumperConfig(
grafter_enable=True,
grafter_backend=get_default_distributed_backend(get_device()),
grafter_role=role,
grafter_b2t_filter=b2t_filter,
grafter_t2b_filter=t2b_filter,
@@ -3464,10 +3500,10 @@ class TestGrafterDistributed:
)
try:
if rank == 0:
tensor = torch.tensor([1.0, 2.0, 3.0], device="cuda:0")
tensor = torch.tensor([1.0, 2.0, 3.0], device=get_device(0))
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.zeros(3, device="cuda:1")
target = torch.zeros(3, device=get_device(1))
with _capture_stdout() as captured:
grafter.maybe_intercept(value=target, tags={"name": "x"})
assert target.tolist() == [1.0, 2.0, 3.0], f"got {target.tolist()}"
@@ -3497,10 +3533,10 @@ class TestGrafterDistributed:
)
try:
if rank == 1:
tensor = torch.tensor([4.0, 5.0, 6.0], device="cuda:1")
tensor = torch.tensor([4.0, 5.0, 6.0], device=get_device(1))
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.zeros(3, device="cuda:0")
target = torch.zeros(3, device=get_device(0))
grafter.maybe_intercept(value=target, tags={"name": "x"})
assert target.tolist() == [4.0, 5.0, 6.0], f"got {target.tolist()}"
finally:
@@ -3538,10 +3574,10 @@ class TestGrafterDistributed:
)
try:
if rank == 0:
tensor = torch.tensor([1.0, 2.0, 3.0], device="cuda:0")
tensor = torch.tensor([1.0, 2.0, 3.0], device=get_device(0))
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.zeros(3, device="cuda:1")
target = torch.zeros(3, device=get_device(1))
grafter.maybe_intercept(value=target, tags={"name": "x"})
assert target.tolist() == [2.0, 4.0, 6.0], f"got {target.tolist()}"
finally:
@@ -3564,7 +3600,7 @@ class TestGrafterDistributed:
)
)
try:
target = torch.tensor([7.0, 7.0, 7.0], device=f"cuda:{rank}")
target = torch.tensor([7.0, 7.0, 7.0], device=get_device(rank))
grafter.maybe_intercept(value=target, tags={"name": "other"})
assert target.tolist() == [7.0, 7.0, 7.0], "tensor must not be modified"
assert grafter._pg is None, "group must not init for unmatched name"
@@ -3592,11 +3628,11 @@ class TestGrafterDistributed:
try:
if rank == 0:
# Baseline sends shape=(3,)
tensor = torch.tensor([1.0, 2.0, 3.0], device="cuda:0")
tensor = torch.tensor([1.0, 2.0, 3.0], device=get_device(0))
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
# Target's local target has shape=(4,) — mismatch with sender.
target = torch.tensor([7.0, 7.0, 7.0, 7.0], device="cuda:1")
target = torch.tensor([7.0, 7.0, 7.0, 7.0], device=get_device(1))
# No exception should propagate; tensor must stay unchanged.
grafter.maybe_intercept(value=target, tags={"name": "x"})
assert target.tolist() == [
@@ -3644,10 +3680,10 @@ class TestGrafterDistributed:
)
try:
if rank == 0:
tensor = torch.tensor([1.0, 2.0, 3.0], device="cuda:0")
tensor = torch.tensor([1.0, 2.0, 3.0], device=get_device(0))
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.tensor([9.0, 9.0, 9.0], device="cuda:1")
target = torch.tensor([9.0, 9.0, 9.0], device=get_device(1))
with _capture_stdout() as captured:
grafter.maybe_intercept(value=target, tags={"name": "x"})
assert target.tolist() == [
@@ -3700,14 +3736,14 @@ class TestGrafterDistributed:
try:
if rank == 0:
# Baseline (sender) attaches an extras dict.
tensor = torch.tensor([1.0, 2.0, 3.0], device="cuda:0")
tensor = torch.tensor([1.0, 2.0, 3.0], device=get_device(0))
grafter.maybe_intercept(
value=tensor,
tags={"name": "x"},
extras={"fill_value": 42.0},
)
else:
target = torch.zeros(3, device="cuda:1")
target = torch.zeros(3, device=get_device(1))
grafter.maybe_intercept(value=target, tags={"name": "x"})
assert target.tolist() == [
42.0,
@@ -3737,11 +3773,11 @@ class TestGrafterDistributed:
with _capture_stdout() as captured:
if rank == 1:
time.sleep(4)
tensor = torch.tensor([1.0, 2.0, 3.0], device=f"cuda:{rank}")
tensor = torch.tensor([1.0, 2.0, 3.0], device=get_device(rank))
if rank == 0:
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.zeros(3, device=f"cuda:{rank}")
target = torch.zeros(3, device=get_device(rank))
grafter.maybe_intercept(value=target, tags={"name": "x"})
output = captured.getvalue()
if rank == 0:
@@ -3772,11 +3808,11 @@ class TestGrafterDistributed:
)
try:
if rank == 0:
tensor = torch.tensor([1.0, 2.0, 3.0], device="cuda:0")
tensor = torch.tensor([1.0, 2.0, 3.0], device=get_device(0))
# Note: extras kwarg omitted entirely → None on the wire.
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.zeros(3, device="cuda:1")
target = torch.zeros(3, device=get_device(1))
with _capture_stdout() as captured:
grafter.maybe_intercept(value=target, tags={"name": "x"})
# Default identity transform copies tensor through; recv log
@@ -3808,8 +3844,8 @@ class TestGrafterDistributed:
)
try:
if rank == 0:
t1 = torch.tensor([1.0, 2.0, 3.0], device="cuda:0")
t2 = torch.tensor([4.0, 5.0, 6.0], device="cuda:0")
t1 = torch.tensor([1.0, 2.0, 3.0], device=get_device(0))
t2 = torch.tensor([4.0, 5.0, 6.0], device=get_device(0))
grafter.maybe_intercept(value=t1, tags={"name": "x"})
pg_after_first = grafter._pg
assert pg_after_first is not None
@@ -3818,8 +3854,8 @@ class TestGrafterDistributed:
"_pg must be cached across calls, not re-initialized"
)
else:
target1 = torch.zeros(3, device="cuda:1")
target2 = torch.zeros(3, device="cuda:1")
target1 = torch.zeros(3, device=get_device(1))
target2 = torch.zeros(3, device=get_device(1))
grafter.maybe_intercept(value=target1, tags={"name": "x"})
pg_after_first = grafter._pg
assert pg_after_first is not None
@@ -3867,10 +3903,10 @@ class TestGrafterDistributed:
)
try:
if rank == 0:
tensor = torch.tensor([1.0, 2.0, 3.0], device="cuda:0")
tensor = torch.tensor([1.0, 2.0, 3.0], device=get_device(0))
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.tensor([7.0, 7.0, 7.0], device="cuda:1")
target = torch.tensor([7.0, 7.0, 7.0], device=get_device(1))
with _capture_stdout() as captured:
grafter.maybe_intercept(value=target, tags={"name": "x"})
# target must be unchanged; error must be logged with traceback.
@@ -4210,6 +4246,7 @@ class TestGrafterE2eExample:
DUMPER_GRAFTER_MASTER_PORT=str(graft_port),
DUMPER_GRAFTER_BASELINE_WORLD_SIZE="1",
DUMPER_GRAFTER_TARGET_WORLD_SIZE="1",
DUMPER_GRAFTER_BACKEND=get_default_distributed_backend(get_device()),
DUMPER_GRAFTER_B2T_FILTER="name == 'attn_output'",
DUMPER_GRAFTER_T2B_FILTER="name == 'attn_input'",
DUMPER_GRAFTER_GROUP_NAME="grafter_e2e",
@@ -4249,7 +4286,7 @@ class TestGrafterE2eExample:
# min/max/mean/sample fields wildcard out.
tinfo_f32_4 = (
r"type=<class 'torch\.Tensor'> shape=torch\.Size\(\[4\]\) "
r"dtype=torch\.float32 device=cuda:\d stride=\(1,\) "
r"dtype=torch\.float32 device=\w+:\d stride=\(1,\) "
r"req_grad=False .*"
)
diff = r"rel_diff=[-\d.eE+]+ max_abs=[-\d.eE+]+ mean_abs=[-\d.eE+]+"
@@ -4267,7 +4304,7 @@ class TestGrafterE2eExample:
r"\A"
f"{prefix}\\[Grafter\\] init group: role=baseline "
r"baseline_world=1 target_world=1 rank=0 "
r"init_method=tcp://127\.0\.0\.1:\d+ backend=nccl "
r"init_method=tcp://127\.0\.0\.1:\d+ backend=\w+ "
r"name=grafter_e2e\n"
f"{prefix}\\[Grafter\\] recv role=baseline dir=t2b "
f"tags={attn_input_tags} n_senders=1 "
@@ -4284,7 +4321,7 @@ class TestGrafterE2eExample:
r"\A"
f"{prefix}\\[Grafter\\] init group: role=target "
r"baseline_world=1 target_world=1 rank=1 "
r"init_method=tcp://127\.0\.0\.1:\d+ backend=nccl "
r"init_method=tcp://127\.0\.0\.1:\d+ backend=\w+ "
r"name=grafter_e2e\n"
f"{prefix}\\[Grafter\\] send role=target dir=t2b "
f"tags={attn_input_tags} extras={extras_lit} "
@@ -4320,7 +4357,7 @@ class TestGrafterE2eExample:
# `_e2e_transform` runs on the recv side, asserts the dummy extras
# made it across, then returns target's q so baseline's local
# placeholder is overwritten via .copy_().
q = torch.tensor([99.0, 99.0, 99.0, 99.0], device="cuda:0")
q = torch.tensor([99.0, 99.0, 99.0, 99.0], device=get_device(0))
dumper.dump("attn_input", q)
assert q.tolist() == [1.0, 2.0, 3.0, 4.0], (
f"baseline's q should be overwritten by target's via the t->b graft, "
@@ -4344,7 +4381,7 @@ class TestGrafterE2eExample:
# Step 1: graft input. target sends its real q to baseline along
# with a dummy extras key the recv-side transform will assert on.
q = torch.tensor([1.0, 2.0, 3.0, 4.0], device="cuda:1")
q = torch.tensor([1.0, 2.0, 3.0, 4.0], device=get_device(1))
dumper.dump(
"attn_input",
q,