Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
194 changes: 193 additions & 1 deletion tests/e2e/nightly/single_node/ops/multicard_ops_a3/test_sfa_cp_a2a.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0

import os
import random
import traceback

Expand Down Expand Up @@ -106,7 +107,7 @@ def _worker(rank: int, world_size: int, port: int, result_queue: mp.SimpleQueue)

def test_registered_sfa_dcp_a2a_fused_multi_rank() -> None:
world_size = 2
mp.set_start_method("fork", force=True)
mp.set_start_method("spawn", force=True)
result_queue = mp.SimpleQueue()
port = 29_501 + random.randint(0, 10_000)
processes = [
Expand All @@ -125,3 +126,194 @@ def test_registered_sfa_dcp_a2a_fused_multi_rank() -> None:

assert all(process.exitcode == 0 for process in processes)
assert results == [None] * world_size, "\n".join(result for result in results if result is not None)


@torch.inference_mode()
def _max_sum_graph_worker(rank: int, world_size: int, port: int, result_queue: mp.SimpleQueue) -> None:
dcp_group = None
try:
os.environ["TRITON_CACHE_DIR"] = f"/tmp/sfa_dcp_max_sum_graph_{port}_rank{rank}"
torch_npu.npu.set_device(rank)
init_device_properties_triton()
init_distributed_environment(
world_size=world_size,
rank=rank,
local_rank=rank,
distributed_init_method=f"tcp://127.0.0.1:{port}",
backend="hccl",
)
dcp_group = init_model_parallel_group(
[list(range(world_size))],
local_rank=rank,
backend="hccl",
group_name="sfa_dcp_max_sum_graph_test",
use_device_communicator=False,
)

# Keep input/output storage addresses fixed across capture and both
# replays so the test distinguishes real graph reuse from eager reruns.
for scatter_dim in (0, 1):
num_tokens, num_heads, head_dim = (4, 8, 512) if scatter_dim == 0 else (5, 8, 512)

def make_inputs(
seed: int,
tokens: int,
heads: int,
dimension: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
torch.manual_seed(seed)
outputs = torch.randn(
world_size,
tokens,
heads,
dimension,
dtype=torch.bfloat16,
device="npu",
)
maximum = torch.randn(
world_size,
1,
tokens,
heads,
dtype=torch.float32,
device="npu",
)
summation = torch.rand_like(maximum) + 0.125
return outputs, maximum, summation

sender_outputs, sender_max, sender_sum = make_inputs(
20260918 + scatter_dim,
num_tokens,
num_heads,
head_dim,
)
static_output = sender_outputs[rank].contiguous()
static_max = sender_max[rank].contiguous()
static_sum = sender_sum[rank].contiguous()
input_pointers = (static_output.data_ptr(), static_max.data_ptr(), static_sum.data_ptr())

def expected_output(
outputs: torch.Tensor,
maximum: torch.Tensor,
summation: torch.Tensor,
scatter: int,
tokens: int,
heads: int,
) -> torch.Tensor:
sender_lse = maximum[:, 0] + torch.log(summation[:, 0])
if scatter == 0:
local_tokens = tokens // world_size
token_slice = slice(rank * local_tokens, (rank + 1) * local_tokens)
return _reference_merge(outputs[:, token_slice], sender_lse[:, token_slice])
local_heads = heads // world_size
head_slice = slice(rank * local_heads, (rank + 1) * local_heads)
return _reference_merge(outputs[:, :, head_slice], sender_lse[:, :, head_slice])

eager = torch.ops.vllm.sfa_dcp_a2a_fused_max_sum(
static_output,
static_max,
static_sum,
world_size,
scatter_dim,
dcp_group.unique_name,
)
torch_npu.npu.synchronize()
torch.testing.assert_close(
eager,
expected_output(sender_outputs, sender_max, sender_sum, scatter_dim, num_tokens, num_heads),
atol=2e-2,
rtol=2e-2,
)

for _ in range(3):
torch.ops.vllm.sfa_dcp_a2a_fused_max_sum(
static_output,
static_max,
static_sum,
world_size,
scatter_dim,
dcp_group.unique_name,
)
torch_npu.npu.synchronize()
torch.distributed.barrier(group=dcp_group.device_group)

graph = torch_npu.npu.NPUGraph()
with torch_npu.npu.graph(graph):
graph_output = torch.ops.vllm.sfa_dcp_a2a_fused_max_sum(
static_output,
static_max,
static_sum,
world_size,
scatter_dim,
dcp_group.unique_name,
)
graph.replay()
torch_npu.npu.synchronize()
torch.distributed.barrier(group=dcp_group.device_group)
torch.testing.assert_close(
graph_output,
expected_output(sender_outputs, sender_max, sender_sum, scatter_dim, num_tokens, num_heads),
atol=2e-2,
rtol=2e-2,
)
first_output = graph_output.clone()
graph_output_pointer = graph_output.data_ptr()

changed_outputs, changed_max, changed_sum = make_inputs(
20260919 + scatter_dim,
num_tokens,
num_heads,
head_dim,
)
static_output.copy_(changed_outputs[rank])
static_max.copy_(changed_max[rank])
static_sum.copy_(changed_sum[rank])
sender_outputs = changed_outputs
sender_max = changed_max
sender_sum = changed_sum
torch_npu.npu.synchronize()
assert input_pointers == (static_output.data_ptr(), static_max.data_ptr(), static_sum.data_ptr())

graph.replay()
torch_npu.npu.synchronize()
torch.distributed.barrier(group=dcp_group.device_group)
assert graph_output.data_ptr() == graph_output_pointer
torch.testing.assert_close(
graph_output,
expected_output(sender_outputs, sender_max, sender_sum, scatter_dim, num_tokens, num_heads),
atol=2e-2,
rtol=2e-2,
)
assert not torch.equal(graph_output, first_output)
torch.distributed.barrier(group=dcp_group.device_group)

result_queue.put(None)
except Exception:
result_queue.put(traceback.format_exc())
finally:
if dcp_group is not None:
dcp_group.destroy()
destroy_distributed_environment()


def test_registered_sfa_dcp_a2a_fused_max_sum_aclgraph_multi_rank() -> None:
world_size = 2
mp.set_start_method("spawn", force=True)
result_queue = mp.SimpleQueue()
port = 39_501 + random.randint(0, 10_000)
processes = [
mp.Process(
target=_max_sum_graph_worker,
args=(rank, world_size, port, result_queue),
)
for rank in range(world_size)
]

for process in processes:
process.start()
results = [result_queue.get() for _ in processes]
for process in processes:
process.join()

assert all(process.exitcode == 0 for process in processes)
assert results == [None] * world_size, "\n".join(result for result in results if result is not None)
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from vllm_ascend.ops.triton.sfa_cp import (
fused_sfa_dcp_lse_combine,
pack_sfa_dcp_output_lse,
pack_sfa_dcp_output_max_sum,
)


Expand Down Expand Up @@ -36,6 +37,36 @@ def _simulate_receive(
return torch.stack([send_buffers[source_rank][destination_rank] for source_rank in range(dcp_size)])


def _materialize_pa_bsnd_lse(
softmax_max: torch.Tensor,
softmax_sum: torch.Tensor,
) -> torch.Tensor:
# This is the current-main Python contract. Non-finite results are later
# treated as invalid rank contributions by the existing packed consumer.
return (softmax_max + torch.log(softmax_sum))[:, 0].unsqueeze(-1)


def _simulate_receive_max_sum(
sender_outputs: torch.Tensor,
sender_max: torch.Tensor,
sender_sum: torch.Tensor,
destination_rank: int,
scatter_dim: int,
) -> torch.Tensor:
dcp_size = sender_outputs.shape[0]
send_buffers = [
pack_sfa_dcp_output_max_sum(
sender_outputs[source_rank],
sender_max[source_rank],
sender_sum[source_rank],
dcp_size,
scatter_dim,
)
for source_rank in range(dcp_size)
]
return torch.stack([send_buffers[source_rank][destination_rank] for source_rank in range(dcp_size)])


@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
@pytest.mark.parametrize("scatter_dim", [0, 1])
@pytest.mark.parametrize("head_dim", [96, 128, 160, 256])
Expand Down Expand Up @@ -247,3 +278,123 @@ def test_finite_lse_outside_activation_dtype_range(
tolerance = 2e-2 if dtype == torch.bfloat16 else 1e-2
torch.testing.assert_close(actual, expected, atol=tolerance, rtol=tolerance)
assert torch.count_nonzero(actual).item() > 0


@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
@pytest.mark.parametrize("scatter_dim", [0, 1])
@torch.inference_mode()
def test_pack_pa_bsnd_max_sum_and_fused_lse_combine(
dtype: torch.dtype,
scatter_dim: int,
) -> None:
torch.manual_seed(20260913)
dcp_size = 8
num_tokens, num_heads, head_dim = (16, 8, 512) if scatter_dim == 0 else (5, 64, 512)
sender_outputs = torch.randn(
dcp_size,
num_tokens,
num_heads,
head_dim,
dtype=dtype,
device="npu",
)
sender_max = torch.randn(
dcp_size,
1,
num_tokens,
num_heads,
dtype=torch.float32,
device="npu",
)
sender_sum = torch.rand_like(sender_max) + 0.125
sender_lse = _materialize_pa_bsnd_lse(sender_max, sender_sum)

destination_rank = 3
recv = _simulate_receive_max_sum(
sender_outputs,
sender_max,
sender_sum,
destination_rank,
scatter_dim,
)
if scatter_dim == 0:
local_tokens = num_tokens // dcp_size
token_slice = slice(destination_rank * local_tokens, (destination_rank + 1) * local_tokens)
expected = _reference_merge(sender_outputs[:, token_slice], sender_lse[:, token_slice, :, 0])
else:
local_heads = num_heads // dcp_size
head_slice = slice(destination_rank * local_heads, (destination_rank + 1) * local_heads)
expected = _reference_merge(sender_outputs[:, :, head_slice], sender_lse[:, :, head_slice, 0])

actual = fused_sfa_dcp_lse_combine(recv, head_dim, scatter_dim)
tolerance = 2e-2 if dtype == torch.bfloat16 else 1e-2
torch.testing.assert_close(actual, expected, atol=tolerance, rtol=tolerance)


@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
@pytest.mark.parametrize("scatter_dim", [0, 1])
@torch.inference_mode()
def test_pa_bsnd_max_sum_strides_invalid_statistics_and_all_invalid_rows(
dtype: torch.dtype,
scatter_dim: int,
) -> None:
torch.manual_seed(20260914)
dcp_size, num_tokens, num_heads, head_dim = 8, 16, 64, 128
output_storage = torch.randn(
dcp_size,
num_tokens,
num_heads,
head_dim + 4,
dtype=dtype,
device="npu",
)
max_storage = torch.full(
(dcp_size, 1, num_tokens, num_heads * 2),
70_000.0,
dtype=torch.float32,
device="npu",
)
sum_storage = torch.ones_like(max_storage)
sender_outputs = output_storage[..., :head_dim]
sender_max = max_storage[..., ::2]
sender_sum = sum_storage[..., ::2]
sender_max += torch.arange(dcp_size, dtype=torch.float32, device="npu").view(-1, 1, 1, 1) * 0.25

sender_sum[:, :, 0, :] = 0.0
sender_sum[0, 0, 1, 0] = -1.0
sender_sum[1, 0, 1, 0] = float("nan")
sender_sum[2, 0, 1, 0] = float("inf")
sender_max[3, 0, 1, 0] = float("nan")
sender_max[4, 0, 1, 0] = float("inf")
sender_max[5, 0, 1, 0] = float("-inf")
sender_outputs[:, 0] = float("nan")

assert not sender_outputs.is_contiguous()
assert not sender_max.is_contiguous()
assert not sender_sum.is_contiguous()
sender_lse = _materialize_pa_bsnd_lse(sender_max, sender_sum)

destination_rank = 0
recv = _simulate_receive_max_sum(
sender_outputs,
sender_max,
sender_sum,
destination_rank,
scatter_dim,
)
if scatter_dim == 0:
expected = _reference_merge(
sender_outputs[:, : num_tokens // dcp_size],
sender_lse[:, : num_tokens // dcp_size, :, 0],
)
else:
expected = _reference_merge(
sender_outputs[:, :, : num_heads // dcp_size],
sender_lse[:, :, : num_heads // dcp_size, 0],
)
actual = fused_sfa_dcp_lse_combine(recv, head_dim, scatter_dim)

tolerance = 2e-2 if dtype == torch.bfloat16 else 1e-2
torch.testing.assert_close(actual, expected, atol=tolerance, rtol=tolerance)
assert torch.count_nonzero(actual[0]).item() == 0
assert torch.isfinite(actual[1:]).all()
Loading
Loading