mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
178 lines
8.0 KiB
Python
178 lines
8.0 KiB
Python
"""Distributed test for the load-balanced DeepSeek-V4 CSA indexer under contiguous CP.
|
|
|
|
Run with:
|
|
torchrun --nproc_per_node=2 tests/fast-gpu/test_dsv4_indexer_cp_balance.py
|
|
torchrun --nproc_per_node=4 tests/fast-gpu/test_dsv4_indexer_cp_balance.py
|
|
|
|
Each rank runs V4Indexer.forward's two steps around the key gather, start_row_exchange and
|
|
topk_for_local_rows, once keeping its rows and once balancing them with a collective on the CP group
|
|
while the rows are in flight. The balanced picks must equal the local ones, bit for bit for the torch
|
|
top-k and as sets for flashinfer; unpacked, the local ones must also equal the pre-balancing
|
|
indexer's (tests/fast/test_dsv4_thd.py pins the THD bounds to running each sample alone). Cases:
|
|
unpacked batch 1 and 2, THD packs with a long document and with odd scored-row counts, a pack of
|
|
tiny documents that leaves a CP4 rank nothing to score, and a pack of equal short documents, which
|
|
must skip the exchange.
|
|
"""
|
|
|
|
import os
|
|
import sys
|
|
|
|
import torch
|
|
import torch.distributed as dist
|
|
from tests.ci.ci_register import register_cuda_ci
|
|
|
|
from miles_plugins.models.deepseek_v4.ops.cp_row_balance import RowExchange
|
|
from miles_plugins.models.deepseek_v4.ops.kernel.tilelang_indexer_fwd import (
|
|
_make_causal_cu_seqlens,
|
|
batched_indexer_fwd,
|
|
)
|
|
from miles_plugins.models.deepseek_v4.ops.thd_utils import ThdLayout, compressed_cu_seqlens
|
|
from miles_plugins.models.deepseek_v4.ops.v4_indexer import start_row_exchange, topk_for_local_rows
|
|
from miles_plugins.models.dsa_topk import get_dsa_topk_fn
|
|
|
|
register_cuda_ci(
|
|
est_time=60,
|
|
suite="stage-c-4-gpu-h200",
|
|
labels=["precision", "megatron"],
|
|
hardware=["hopper", "blackwell"],
|
|
)
|
|
|
|
SEQLEN_GLOBAL = 16384
|
|
RATIO = 4
|
|
HEADS, INDEX_DIM, TOPK = 64, 128, 512
|
|
|
|
# flashinfer's default top-k breaks exact score ties differently from call to call, so two runs of
|
|
# the same rows can disagree now and then; its deterministic mode (read at call time) does not.
|
|
os.environ["SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC"] = "1"
|
|
|
|
|
|
def setup_dist():
|
|
rank = int(os.environ["RANK"])
|
|
world_size = int(os.environ["WORLD_SIZE"])
|
|
torch.cuda.set_device(int(os.environ["LOCAL_RANK"]))
|
|
dist.init_process_group(backend="nccl", rank=rank, world_size=world_size)
|
|
return rank, world_size
|
|
|
|
|
|
def _inputs(rank, rank_rows, n_kv, bsz):
|
|
"""Keys are all-gathered in the real layer, so every rank builds the same global key block."""
|
|
generator = torch.Generator(device="cuda").manual_seed(1234)
|
|
q = torch.randn(SEQLEN_GLOBAL, bsz, HEADS, INDEX_DIM, device="cuda", dtype=torch.bfloat16, generator=generator)
|
|
k = torch.randn(n_kv, bsz, INDEX_DIM, device="cuda", dtype=torch.bfloat16, generator=generator)
|
|
weights = torch.randn(SEQLEN_GLOBAL, bsz, HEADS, device="cuda", dtype=torch.float32, generator=generator)
|
|
local = slice(rank * rank_rows, (rank + 1) * rank_rows)
|
|
return q[local].contiguous(), k, weights[local].contiguous()
|
|
|
|
|
|
def _thd_layout(seq_lens, rank, rank_rows):
|
|
cu_seqlens = torch.tensor([0, *torch.tensor(seq_lens).cumsum(0).tolist()], device="cuda", dtype=torch.int32)
|
|
layout = ThdLayout(
|
|
cu_seqlens=cu_seqlens, seq_lens=tuple(seq_lens), global_start=rank * rank_rows, max_seqlen=max(seq_lens)
|
|
)
|
|
layout.cu_seqlens_compressed = compressed_cu_seqlens(cu_seqlens, RATIO)
|
|
return layout
|
|
|
|
|
|
def _unbalanced_unpacked_topk(q, k, weights, rank, topk_fn):
|
|
"""The unpacked indexer before balancing: this rank's own rows, bounds sliced from the contiguous run."""
|
|
rank_rows = q.shape[0]
|
|
cu_ks, cu_ke = _make_causal_cu_seqlens(SEQLEN_GLOBAL, k.shape[0], RATIO, q.device)
|
|
cu_ks, cu_ke = (
|
|
cu_ks[rank * rank_rows : (rank + 1) * rank_rows],
|
|
cu_ke[rank * rank_rows : (rank + 1) * rank_rows],
|
|
)
|
|
scores = batched_indexer_fwd(q, k, weights, cu_ks, cu_ke)
|
|
bsz, rows, n_kv = scores.shape
|
|
return topk_fn(scores.reshape(bsz * rows, n_kv), min(TOPK, n_kv)).reshape(bsz, rows, -1)
|
|
|
|
|
|
def _same_picks(got, expected, topk_backend):
|
|
if topk_backend == "torch":
|
|
return torch.equal(got, expected)
|
|
# flashinfer returns each row's picks unsorted
|
|
return torch.equal(got.sort(dim=-1).values, expected.sort(dim=-1).values)
|
|
|
|
|
|
def check_picks(rank, world_size, topk_backend, thd_seq_lens=None, bsz=1):
|
|
"""Returns (the balanced picks equal the local ones, and unpacked the pre-balancing ones; the exchange ran)."""
|
|
rank_rows = SEQLEN_GLOBAL // world_size
|
|
thd_layout = _thd_layout(thd_seq_lens, rank, rank_rows) if thd_seq_lens else None
|
|
n_kv = int(thd_layout.cu_seqlens_compressed[-1]) if thd_layout else SEQLEN_GLOBAL // RATIO
|
|
q, k, weights = _inputs(rank, rank_rows, n_kv, bsz)
|
|
topk_fn = get_dsa_topk_fn(topk_backend)
|
|
options = dict(compress_ratio=RATIO, index_topk=TOPK, topk_fn=topk_fn)
|
|
group = dist.group.WORLD
|
|
|
|
kept = start_row_exchange(q, weights, thd_layout, group, balance=False)
|
|
local = topk_for_local_rows(kept, k, thd_layout, **options)
|
|
exchange = start_row_exchange(q, weights, thd_layout, group, balance=True)
|
|
# in forward the compressor's CP all-gathers queue on this communicator behind the exchange
|
|
dist.all_reduce(torch.ones(1, device="cuda"), group=group)
|
|
balanced = topk_for_local_rows(exchange, k, thd_layout, **options)
|
|
|
|
picks_equal = _same_picks(balanced, local, topk_backend)
|
|
if thd_layout is None:
|
|
expected = _unbalanced_unpacked_topk(q, k, weights, rank, topk_fn)
|
|
picks_equal = picks_equal and _same_picks(local, expected, topk_backend)
|
|
return picks_equal, isinstance(exchange, RowExchange)
|
|
|
|
|
|
def _topk_backends():
|
|
backends = ["torch"]
|
|
try:
|
|
get_dsa_topk_fn("flashinfer")(torch.randn(2, 1024, device="cuda"), 8)
|
|
backends.append("flashinfer")
|
|
except ImportError:
|
|
pass
|
|
return backends
|
|
|
|
|
|
def main():
|
|
rank, world_size = setup_dist()
|
|
try:
|
|
long_doc = SEQLEN_GLOBAL - 4 * 1000 - 116
|
|
cases = {
|
|
"unpacked sequence": dict(thd_seq_lens=None, bsz=1, expect_exchange=True),
|
|
"unpacked batch of 2": dict(thd_seq_lens=None, bsz=2, expect_exchange=True),
|
|
"pack: long doc + short docs + pad": dict(
|
|
thd_seq_lens=[long_doc, 1000, 1000, 1000, 1000, 116], bsz=1, expect_exchange=True
|
|
),
|
|
# ranks score unequal, odd row counts, so the scorer's last 2-row block runs past the end
|
|
"pack: odd scored-row counts": dict(thd_seq_lens=[SEQLEN_GLOBAL - 2, 1, 1], bsz=1, expect_exchange=True),
|
|
"pack: equal short docs": dict(thd_seq_lens=[512] * (SEQLEN_GLOBAL // 512), bsz=1, expect_exchange=False),
|
|
# documents shorter than 2cp: at CP4 rank 1 scores nothing; at CP8 the gate keeps the rows local
|
|
"pack: tiny docs": dict(
|
|
thd_seq_lens=[1] * 7520 + [2] * 2432 + [5] * 800, bsz=1, expect_exchange=world_size in (2, 4)
|
|
),
|
|
}
|
|
passed = True
|
|
for topk_backend in _topk_backends():
|
|
for name, case in cases.items():
|
|
picks_equal, exchanged = check_picks(
|
|
rank, world_size, topk_backend, thd_seq_lens=case["thd_seq_lens"], bsz=case["bsz"]
|
|
)
|
|
flags = torch.tensor(
|
|
[picks_equal, exchanged == case["expect_exchange"]], device="cuda", dtype=torch.int32
|
|
)
|
|
dist.all_reduce(flags, op=dist.ReduceOp.MIN)
|
|
ok = bool(flags.all())
|
|
passed = passed and ok
|
|
if rank == 0:
|
|
print(
|
|
f"CP={world_size} {topk_backend:10s} {name:36s} picks equal: {bool(flags[0])} "
|
|
f"exchange as expected: {bool(flags[1])}"
|
|
)
|
|
if rank == 0:
|
|
print(f"\nCP={world_size} test PASSED!" if passed else "FAILED!")
|
|
if not passed:
|
|
sys.exit(1)
|
|
finally:
|
|
dist.destroy_process_group()
|
|
|
|
|
|
if __name__ == "__main__":
|
|
# Self-bootstrap under torchrun when run as `python3 file.py`, the CUDA CI runner's mode.
|
|
if "RANK" not in os.environ:
|
|
os.execvp("torchrun", ["torchrun", "--nproc_per_node=4", __file__])
|
|
main()
|