|
|
|
@@ -4,6 +4,7 @@ import pytest
|
|
|
|
|
import torch
|
|
|
|
|
|
|
|
|
|
from sglang.srt.debug_utils.comparator.aligner.unsharder.executor import (
|
|
|
|
|
UnsharderResult,
|
|
|
|
|
_apply_unshard,
|
|
|
|
|
_verify_replicated_group,
|
|
|
|
|
execute_unsharder_plan,
|
|
|
|
@@ -23,7 +24,7 @@ from sglang.srt.debug_utils.comparator.dims import (
|
|
|
|
|
ParallelAxis,
|
|
|
|
|
parse_dims,
|
|
|
|
|
)
|
|
|
|
|
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
|
|
|
|
|
from sglang.srt.debug_utils.comparator.output_types import ReplicatedCheckResult
|
|
|
|
|
from sglang.test.ci.ci_register import register_cpu_ci
|
|
|
|
|
|
|
|
|
|
register_cpu_ci(est_time=10, suite="default", nightly=True)
|
|
|
|
@@ -49,11 +50,12 @@ class TestExecuteUnsharderPlan:
|
|
|
|
|
assert len(plans) == 1
|
|
|
|
|
|
|
|
|
|
named_shards: list[torch.Tensor] = _name_tensors(shards, dim_specs)
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
result = execute_unsharder_plan(plans[0], named_shards)
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
assert torch.allclose(result[0].rename(None), full_tensor)
|
|
|
|
|
assert warnings == []
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
plans[0], named_shards
|
|
|
|
|
)
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
assert torch.allclose(unsharder_result.tensors[0].rename(None), full_tensor)
|
|
|
|
|
assert unsharder_result.replicated_checks == []
|
|
|
|
|
|
|
|
|
|
def test_scrambled_world_ranks_correct_result(self) -> None:
|
|
|
|
|
full_tensor = torch.randn(4, 8)
|
|
|
|
@@ -79,11 +81,12 @@ class TestExecuteUnsharderPlan:
|
|
|
|
|
dim_specs,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
result = execute_unsharder_plan(plans[0], tensors_ordered_by_world_rank)
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
assert torch.allclose(result[0].rename(None), full_tensor)
|
|
|
|
|
assert warnings == []
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
plans[0], tensors_ordered_by_world_rank
|
|
|
|
|
)
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
assert torch.allclose(unsharder_result.tensors[0].rename(None), full_tensor)
|
|
|
|
|
assert unsharder_result.replicated_checks == []
|
|
|
|
|
|
|
|
|
|
def test_single_step_reduces_tensor_count(self) -> None:
|
|
|
|
|
"""8 tensors with 2 groups of 4 produce 2 output tensors."""
|
|
|
|
@@ -113,13 +116,15 @@ class TestExecuteUnsharderPlan:
|
|
|
|
|
tensors.append(source[tp_rank])
|
|
|
|
|
|
|
|
|
|
named_tensors: list[torch.Tensor] = _name_tensors(tensors, dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
intermediate = execute_unsharder_plan(plans[0], named_tensors)
|
|
|
|
|
assert len(intermediate) == 4
|
|
|
|
|
intermediate_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
plans[0], named_tensors
|
|
|
|
|
)
|
|
|
|
|
assert len(intermediate_result.tensors) == 4
|
|
|
|
|
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
final = execute_unsharder_plan(plans[1], intermediate)
|
|
|
|
|
assert len(final) == 1
|
|
|
|
|
final_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
plans[1], intermediate_result.tensors
|
|
|
|
|
)
|
|
|
|
|
assert len(final_result.tensors) == 1
|
|
|
|
|
|
|
|
|
|
def test_cp_tp_concat(self) -> None:
|
|
|
|
|
"""CP=2 + TP=2: multi-step unshard reconstructs original tensor."""
|
|
|
|
@@ -145,9 +150,9 @@ class TestExecuteUnsharderPlan:
|
|
|
|
|
assert len(plans) == 2
|
|
|
|
|
|
|
|
|
|
current: list[torch.Tensor] = _name_tensors(tensors, dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
for plan in plans:
|
|
|
|
|
current = execute_unsharder_plan(plan, current)
|
|
|
|
|
for plan in plans:
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, current)
|
|
|
|
|
current = unsharder_result.tensors
|
|
|
|
|
|
|
|
|
|
assert len(current) == 1
|
|
|
|
|
assert torch.allclose(current[0].rename(None), full_tensor)
|
|
|
|
@@ -187,9 +192,9 @@ class TestExecuteUnsharderPlan:
|
|
|
|
|
assert len(plans) == 2
|
|
|
|
|
|
|
|
|
|
current: list[torch.Tensor] = _name_tensors(tensors, dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
for plan in plans:
|
|
|
|
|
current = execute_unsharder_plan(plan, current)
|
|
|
|
|
for plan in plans:
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, current)
|
|
|
|
|
current = unsharder_result.tensors
|
|
|
|
|
|
|
|
|
|
assert len(current) == 1
|
|
|
|
|
assert torch.allclose(current[0].rename(None), full_tensor)
|
|
|
|
@@ -241,9 +246,9 @@ class TestExecuteUnsharderPlan:
|
|
|
|
|
assert len(plans) == 3
|
|
|
|
|
|
|
|
|
|
current: list[torch.Tensor] = _name_tensors(tensors, dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
for plan in plans:
|
|
|
|
|
current = execute_unsharder_plan(plan, current)
|
|
|
|
|
for plan in plans:
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, current)
|
|
|
|
|
current = unsharder_result.tensors
|
|
|
|
|
|
|
|
|
|
assert len(current) == 1
|
|
|
|
|
assert torch.allclose(current[0].rename(None), full_tensor)
|
|
|
|
@@ -290,9 +295,9 @@ class TestExecuteUnsharderPlan:
|
|
|
|
|
assert len(plans) == 3
|
|
|
|
|
|
|
|
|
|
current: list[torch.Tensor] = _name_tensors(tensors, dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
for plan in plans:
|
|
|
|
|
current = execute_unsharder_plan(plan, current)
|
|
|
|
|
for plan in plans:
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, current)
|
|
|
|
|
current = unsharder_result.tensors
|
|
|
|
|
|
|
|
|
|
assert len(current) == 1
|
|
|
|
|
assert torch.allclose(current[0].rename(None), full_tensor)
|
|
|
|
@@ -312,11 +317,12 @@ class TestPickOperation:
|
|
|
|
|
assert len(plans) == 1
|
|
|
|
|
assert isinstance(plans[0].params, PickParams)
|
|
|
|
|
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
result = execute_unsharder_plan(plans[0], [tensor, tensor.clone()])
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
assert torch.allclose(result[0].rename(None), tensor)
|
|
|
|
|
assert warnings == []
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
plans[0], [tensor, tensor.clone()]
|
|
|
|
|
)
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
assert torch.allclose(unsharder_result.tensors[0].rename(None), tensor)
|
|
|
|
|
assert all(c.passed for c in unsharder_result.replicated_checks)
|
|
|
|
|
|
|
|
|
|
def test_pick_multiple_groups(self) -> None:
|
|
|
|
|
"""PickParams with multiple groups picks one from each."""
|
|
|
|
@@ -348,10 +354,11 @@ class TestPickOperation:
|
|
|
|
|
tensor = torch.randn(4)
|
|
|
|
|
tensors = [tensor.clone() for _ in range(4)]
|
|
|
|
|
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
result = execute_unsharder_plan(pick_plans[0], tensors)
|
|
|
|
|
assert len(result) == 2
|
|
|
|
|
assert warnings == []
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
pick_plans[0], tensors
|
|
|
|
|
)
|
|
|
|
|
assert len(unsharder_result.tensors) == 2
|
|
|
|
|
assert all(c.passed for c in unsharder_result.replicated_checks)
|
|
|
|
|
|
|
|
|
|
def test_replicated_tp_sharded_cp_e2e(self) -> None:
|
|
|
|
|
"""CP2 TP2, dims='b s(cp) d': replicated TP pick + sharded CP concat round-trip."""
|
|
|
|
@@ -376,9 +383,9 @@ class TestPickOperation:
|
|
|
|
|
assert len(plans) == 2
|
|
|
|
|
|
|
|
|
|
current: list[torch.Tensor] = _name_tensors(tensors, dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
for plan in plans:
|
|
|
|
|
current = execute_unsharder_plan(plan, current)
|
|
|
|
|
for plan in plans:
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, current)
|
|
|
|
|
current = unsharder_result.tensors
|
|
|
|
|
|
|
|
|
|
assert len(current) == 1
|
|
|
|
|
assert torch.allclose(current[0].rename(None), full_tensor)
|
|
|
|
@@ -406,44 +413,44 @@ class TestPickOperation:
|
|
|
|
|
assert all(isinstance(p.params, PickParams) for p in plans)
|
|
|
|
|
|
|
|
|
|
current: list[torch.Tensor] = _name_tensors(tensors, dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
for plan in plans:
|
|
|
|
|
current = execute_unsharder_plan(plan, current)
|
|
|
|
|
for plan in plans:
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, current)
|
|
|
|
|
current = unsharder_result.tensors
|
|
|
|
|
|
|
|
|
|
assert len(current) == 1
|
|
|
|
|
assert torch.allclose(current[0].rename(None), full_tensor)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class TestVerifyReplicatedGroup:
|
|
|
|
|
def test_warns_on_mismatch(self) -> None:
|
|
|
|
|
"""_verify_replicated_group produces warning when replicas differ."""
|
|
|
|
|
def test_fails_on_mismatch(self) -> None:
|
|
|
|
|
"""_verify_replicated_group returns failed check when replicas differ."""
|
|
|
|
|
tensor_a = torch.ones(4)
|
|
|
|
|
tensor_b = torch.ones(4) + 0.1
|
|
|
|
|
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
_verify_replicated_group(
|
|
|
|
|
[tensor_a, tensor_b],
|
|
|
|
|
axis=ParallelAxis.TP,
|
|
|
|
|
group_index=0,
|
|
|
|
|
)
|
|
|
|
|
assert len(warnings) == 1
|
|
|
|
|
assert warnings[0].axis == "tp"
|
|
|
|
|
assert warnings[0].group_index == 0
|
|
|
|
|
assert warnings[0].differing_index == 1
|
|
|
|
|
assert warnings[0].baseline_index == 0
|
|
|
|
|
assert warnings[0].max_abs_diff == pytest.approx(0.1, abs=1e-5)
|
|
|
|
|
checks: list[ReplicatedCheckResult] = _verify_replicated_group(
|
|
|
|
|
[tensor_a, tensor_b],
|
|
|
|
|
axis=ParallelAxis.TP,
|
|
|
|
|
group_index=0,
|
|
|
|
|
)
|
|
|
|
|
assert len(checks) == 1
|
|
|
|
|
assert checks[0].axis == "tp"
|
|
|
|
|
assert checks[0].group_index == 0
|
|
|
|
|
assert checks[0].compared_index == 1
|
|
|
|
|
assert checks[0].baseline_index == 0
|
|
|
|
|
assert not checks[0].passed
|
|
|
|
|
assert checks[0].diff.max_abs_diff == pytest.approx(0.1, abs=1e-5)
|
|
|
|
|
|
|
|
|
|
def test_no_warn_when_identical(self) -> None:
|
|
|
|
|
"""_verify_replicated_group produces no warning for identical replicas."""
|
|
|
|
|
def test_passes_when_identical(self) -> None:
|
|
|
|
|
"""_verify_replicated_group returns passed check for identical replicas."""
|
|
|
|
|
tensor = torch.randn(4, 8)
|
|
|
|
|
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
_verify_replicated_group(
|
|
|
|
|
[tensor, tensor.clone()],
|
|
|
|
|
axis=ParallelAxis.TP,
|
|
|
|
|
group_index=0,
|
|
|
|
|
)
|
|
|
|
|
assert warnings == []
|
|
|
|
|
checks: list[ReplicatedCheckResult] = _verify_replicated_group(
|
|
|
|
|
[tensor, tensor.clone()],
|
|
|
|
|
axis=ParallelAxis.TP,
|
|
|
|
|
group_index=0,
|
|
|
|
|
)
|
|
|
|
|
assert len(checks) == 1
|
|
|
|
|
assert checks[0].passed
|
|
|
|
|
|
|
|
|
|
def test_multiple_mismatches(self) -> None:
|
|
|
|
|
"""_verify_replicated_group reports each differing replica."""
|
|
|
|
@@ -451,19 +458,20 @@ class TestVerifyReplicatedGroup:
|
|
|
|
|
other_a = torch.ones(4)
|
|
|
|
|
other_b = torch.ones(4) * 2
|
|
|
|
|
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
_verify_replicated_group(
|
|
|
|
|
[baseline, other_a, other_b],
|
|
|
|
|
axis=ParallelAxis.CP,
|
|
|
|
|
group_index=1,
|
|
|
|
|
)
|
|
|
|
|
assert len(warnings) == 2
|
|
|
|
|
assert warnings[0].differing_index == 1
|
|
|
|
|
assert warnings[1].differing_index == 2
|
|
|
|
|
assert warnings[1].max_abs_diff == pytest.approx(2.0, abs=1e-5)
|
|
|
|
|
checks: list[ReplicatedCheckResult] = _verify_replicated_group(
|
|
|
|
|
[baseline, other_a, other_b],
|
|
|
|
|
axis=ParallelAxis.CP,
|
|
|
|
|
group_index=1,
|
|
|
|
|
)
|
|
|
|
|
assert len(checks) == 2
|
|
|
|
|
assert checks[0].compared_index == 1
|
|
|
|
|
assert not checks[0].passed
|
|
|
|
|
assert checks[1].compared_index == 2
|
|
|
|
|
assert not checks[1].passed
|
|
|
|
|
assert checks[1].diff.max_abs_diff == pytest.approx(2.0, abs=1e-5)
|
|
|
|
|
|
|
|
|
|
def test_execute_returns_warnings(self) -> None:
|
|
|
|
|
"""execute_unsharder_plan emits warnings for replicated mismatch."""
|
|
|
|
|
def test_execute_returns_replicated_checks(self) -> None:
|
|
|
|
|
"""execute_unsharder_plan returns replicated checks for mismatch."""
|
|
|
|
|
dim_specs = parse_dims("h d")
|
|
|
|
|
parallel_infos = [
|
|
|
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2)},
|
|
|
|
@@ -474,56 +482,58 @@ class TestVerifyReplicatedGroup:
|
|
|
|
|
tensor_a = torch.zeros(4)
|
|
|
|
|
tensor_b = torch.ones(4)
|
|
|
|
|
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
result = execute_unsharder_plan(plans[0], [tensor_a, tensor_b])
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
assert len(warnings) == 1
|
|
|
|
|
assert torch.allclose(result[0].rename(None), tensor_a)
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
plans[0], [tensor_a, tensor_b]
|
|
|
|
|
)
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
assert len(unsharder_result.replicated_checks) == 1
|
|
|
|
|
assert not unsharder_result.replicated_checks[0].passed
|
|
|
|
|
assert torch.allclose(unsharder_result.tensors[0].rename(None), tensor_a)
|
|
|
|
|
|
|
|
|
|
def test_atol_boundary_within(self) -> None:
|
|
|
|
|
"""Difference exactly at atol (1e-6) -> torch.allclose passes -> no warning."""
|
|
|
|
|
"""Difference exactly at atol (1e-6) -> passed."""
|
|
|
|
|
baseline = torch.zeros(4)
|
|
|
|
|
other = torch.full((4,), 1e-6)
|
|
|
|
|
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
_verify_replicated_group(
|
|
|
|
|
[baseline, other],
|
|
|
|
|
axis=ParallelAxis.TP,
|
|
|
|
|
group_index=0,
|
|
|
|
|
)
|
|
|
|
|
assert warnings == []
|
|
|
|
|
checks: list[ReplicatedCheckResult] = _verify_replicated_group(
|
|
|
|
|
[baseline, other],
|
|
|
|
|
axis=ParallelAxis.TP,
|
|
|
|
|
group_index=0,
|
|
|
|
|
)
|
|
|
|
|
assert len(checks) == 1
|
|
|
|
|
assert checks[0].passed
|
|
|
|
|
|
|
|
|
|
def test_atol_boundary_exceeded(self) -> None:
|
|
|
|
|
"""Difference just above atol (1e-6 + 1e-9) -> torch.allclose fails -> warning."""
|
|
|
|
|
"""Difference just above atol (1e-6 + 1e-9) -> failed."""
|
|
|
|
|
baseline = torch.zeros(4)
|
|
|
|
|
other = torch.full((4,), 1e-6 + 1e-9)
|
|
|
|
|
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
_verify_replicated_group(
|
|
|
|
|
[baseline, other],
|
|
|
|
|
axis=ParallelAxis.TP,
|
|
|
|
|
group_index=0,
|
|
|
|
|
)
|
|
|
|
|
assert len(warnings) == 1
|
|
|
|
|
assert warnings[0].differing_index == 1
|
|
|
|
|
checks: list[ReplicatedCheckResult] = _verify_replicated_group(
|
|
|
|
|
[baseline, other],
|
|
|
|
|
axis=ParallelAxis.TP,
|
|
|
|
|
group_index=0,
|
|
|
|
|
)
|
|
|
|
|
assert len(checks) == 1
|
|
|
|
|
assert not checks[0].passed
|
|
|
|
|
assert checks[0].compared_index == 1
|
|
|
|
|
|
|
|
|
|
def test_recompute_pseudo_mismatch_warns(self) -> None:
|
|
|
|
|
"""_verify_replicated_group produces warning for RECOMPUTE_PSEUDO axis mismatch."""
|
|
|
|
|
def test_recompute_pseudo_mismatch(self) -> None:
|
|
|
|
|
"""_verify_replicated_group returns failed check for RECOMPUTE_PSEUDO axis mismatch."""
|
|
|
|
|
tensor_a = torch.ones(4)
|
|
|
|
|
tensor_b = torch.ones(4) + 0.1
|
|
|
|
|
|
|
|
|
|
with warning_sink.context() as warnings:
|
|
|
|
|
_verify_replicated_group(
|
|
|
|
|
[tensor_a, tensor_b],
|
|
|
|
|
axis=ParallelAxis.RECOMPUTE_PSEUDO,
|
|
|
|
|
group_index=0,
|
|
|
|
|
)
|
|
|
|
|
assert len(warnings) == 1
|
|
|
|
|
assert warnings[0].axis == "recompute_pseudo"
|
|
|
|
|
assert warnings[0].group_index == 0
|
|
|
|
|
assert warnings[0].differing_index == 1
|
|
|
|
|
assert warnings[0].baseline_index == 0
|
|
|
|
|
assert warnings[0].max_abs_diff == pytest.approx(0.1, abs=1e-5)
|
|
|
|
|
checks: list[ReplicatedCheckResult] = _verify_replicated_group(
|
|
|
|
|
[tensor_a, tensor_b],
|
|
|
|
|
axis=ParallelAxis.RECOMPUTE_PSEUDO,
|
|
|
|
|
group_index=0,
|
|
|
|
|
)
|
|
|
|
|
assert len(checks) == 1
|
|
|
|
|
assert checks[0].axis == "recompute_pseudo"
|
|
|
|
|
assert checks[0].group_index == 0
|
|
|
|
|
assert checks[0].compared_index == 1
|
|
|
|
|
assert checks[0].baseline_index == 0
|
|
|
|
|
assert not checks[0].passed
|
|
|
|
|
assert checks[0].diff.max_abs_diff == pytest.approx(0.1, abs=1e-5)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class TestThdCpConcat:
|
|
|
|
@@ -537,12 +547,11 @@ class TestThdCpConcat:
|
|
|
|
|
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3]),
|
|
|
|
|
groups=[[0, 1]],
|
|
|
|
|
)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
result = execute_unsharder_plan(plan, [rank0, rank1])
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1])
|
|
|
|
|
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
expected = torch.tensor([1, 2, 3, 4, 5, 6])
|
|
|
|
|
assert torch.equal(result[0].rename(None), expected)
|
|
|
|
|
assert torch.equal(unsharder_result.tensors[0].rename(None), expected)
|
|
|
|
|
|
|
|
|
|
def test_multi_seq(self) -> None:
|
|
|
|
|
"""Multi-seq THD unshard: 2 ranks, seq_lens=[50, 32, 46]."""
|
|
|
|
@@ -563,11 +572,10 @@ class TestThdCpConcat:
|
|
|
|
|
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[50, 32, 46]),
|
|
|
|
|
groups=[[0, 1]],
|
|
|
|
|
)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
result = execute_unsharder_plan(plan, [rank0, rank1])
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1])
|
|
|
|
|
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
unsharded: torch.Tensor = result[0].rename(None)
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
unsharded: torch.Tensor = unsharder_result.tensors[0].rename(None)
|
|
|
|
|
|
|
|
|
|
# seqA: r0(50) + r1(50) = 100 tokens, values 0..99
|
|
|
|
|
assert torch.equal(unsharded[:100], torch.cat([seq_a_r0, seq_a_r1]))
|
|
|
|
@@ -595,11 +603,10 @@ class TestThdCpConcat:
|
|
|
|
|
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3, 2]),
|
|
|
|
|
groups=[[0, 1]],
|
|
|
|
|
)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
result = execute_unsharder_plan(plan, [rank0, rank1])
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1])
|
|
|
|
|
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
unsharded: torch.Tensor = result[0].rename(None)
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
unsharded: torch.Tensor = unsharder_result.tensors[0].rename(None)
|
|
|
|
|
|
|
|
|
|
assert unsharded.shape == (10, hidden)
|
|
|
|
|
assert torch.equal(unsharded[:6], torch.cat([seq_a_r0, seq_a_r1]))
|
|
|
|
@@ -625,11 +632,10 @@ class TestThdCpConcat:
|
|
|
|
|
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3, 2]),
|
|
|
|
|
groups=[[0, 1]],
|
|
|
|
|
)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
result = execute_unsharder_plan(plan, [rank0, rank1])
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1])
|
|
|
|
|
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
unsharded: torch.Tensor = result[0].rename(None)
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
unsharded: torch.Tensor = unsharder_result.tensors[0].rename(None)
|
|
|
|
|
|
|
|
|
|
assert unsharded.shape == (batch, 10, hidden)
|
|
|
|
|
# seqA: r0(3) + r1(3) = 6 tokens per batch
|
|
|
|
@@ -657,11 +663,12 @@ class TestReduceSum:
|
|
|
|
|
assert isinstance(plans[0].params, ReduceSumParams)
|
|
|
|
|
|
|
|
|
|
named_parts: list[torch.Tensor] = _name_tensors([part_a, part_b], dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
result = execute_unsharder_plan(plans[0], named_parts)
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
plans[0], named_parts
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
assert torch.allclose(result[0].rename(None), full_tensor)
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
assert torch.allclose(unsharder_result.tensors[0].rename(None), full_tensor)
|
|
|
|
|
|
|
|
|
|
def test_tp4_reduce(self) -> None:
|
|
|
|
|
"""4 partial tensors sum to full tensor."""
|
|
|
|
@@ -677,11 +684,12 @@ class TestReduceSum:
|
|
|
|
|
assert len(plans) == 1
|
|
|
|
|
|
|
|
|
|
named_parts: list[torch.Tensor] = _name_tensors(parts, dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
result = execute_unsharder_plan(plans[0], named_parts)
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
plans[0], named_parts
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
assert torch.allclose(result[0].rename(None), full_tensor)
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
assert torch.allclose(unsharder_result.tensors[0].rename(None), full_tensor)
|
|
|
|
|
|
|
|
|
|
def test_multi_axis_concat_then_reduce(self) -> None:
|
|
|
|
|
"""CP concat + TP reduce end-to-end."""
|
|
|
|
@@ -707,9 +715,9 @@ class TestReduceSum:
|
|
|
|
|
assert len(plans) == 2
|
|
|
|
|
|
|
|
|
|
current: list[torch.Tensor] = _name_tensors(tensors, dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
for plan in plans:
|
|
|
|
|
current = execute_unsharder_plan(plan, current)
|
|
|
|
|
for plan in plans:
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, current)
|
|
|
|
|
current = unsharder_result.tensors
|
|
|
|
|
|
|
|
|
|
assert len(current) == 1
|
|
|
|
|
assert torch.allclose(current[0].rename(None), full_tensor)
|
|
|
|
@@ -735,11 +743,12 @@ class TestReduceSum:
|
|
|
|
|
plans = compute_unsharder_plan(dim_specs, parallel_infos)
|
|
|
|
|
|
|
|
|
|
named_parts: list[torch.Tensor] = _name_tensors(parts, dim_specs)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
result = execute_unsharder_plan(plans[0], named_parts)
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
plans[0], named_parts
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
assert torch.allclose(result[0].rename(None), full_tensor)
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
assert torch.allclose(unsharder_result.tensors[0].rename(None), full_tensor)
|
|
|
|
|
|
|
|
|
|
def test_reduce_preserves_named_dims(self) -> None:
|
|
|
|
|
"""Named tensor dimensions are preserved through reduce_sum."""
|
|
|
|
@@ -752,13 +761,16 @@ class TestReduceSum:
|
|
|
|
|
params=ReduceSumParams(),
|
|
|
|
|
groups=[[0, 1]],
|
|
|
|
|
)
|
|
|
|
|
with warning_sink.context():
|
|
|
|
|
result = execute_unsharder_plan(plan, [part_a, part_b])
|
|
|
|
|
unsharder_result: UnsharderResult = execute_unsharder_plan(
|
|
|
|
|
plan, [part_a, part_b]
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
assert len(result) == 1
|
|
|
|
|
assert result[0].names == ("h", "d")
|
|
|
|
|
assert len(unsharder_result.tensors) == 1
|
|
|
|
|
assert unsharder_result.tensors[0].names == ("h", "d")
|
|
|
|
|
expected = (part_a.rename(None) + part_b.rename(None)).refine_names("h", "d")
|
|
|
|
|
assert torch.allclose(result[0].rename(None), expected.rename(None))
|
|
|
|
|
assert torch.allclose(
|
|
|
|
|
unsharder_result.tensors[0].rename(None), expected.rename(None)
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__":
|
|
|
|
|