perf(dsv4): balance the CSA indexer across contiguous CP ranks (#3689)

Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
HEJIAN SANG
2026-09-29 09:05:40 -07:00
committed by GitHub
co-authored by Claude Opus 5.5
parent 9e4260de04
commit 3439ec7513
13 changed files with 973 additions and 109 deletions
+10 -1
View File
@@ -1,6 +1,7 @@
import logging
from argparse import Namespace
from collections.abc import Sequence
from dataclasses import dataclass, field
import torch
from megatron.core import mpu
@@ -89,14 +90,22 @@ def verify_megatron_parallel_state(
), f"microbatch_group_size_per_vp_stage mismatch: ParallelState has {actual}, model config has {expected}"
@dataclass
class PackedSeqParamsWithHostCuSeqlens(PackedSeqParams):
"""``PackedSeqParams`` plus the host copy of ``cu_seqlens_q`` that ``get_batch`` already built."""
cu_seqlens_host: tuple[int, ...] = field(kw_only=True)
def get_packed_seq_params(batch: dict[str, torch.Tensor], args: Namespace) -> PackedSeqParams:
if args.qkv_format == "thd":
packed_seq_params = PackedSeqParams(
packed_seq_params = PackedSeqParamsWithHostCuSeqlens(
cu_seqlens_q=batch["cu_seqlens"],
cu_seqlens_kv=batch["cu_seqlens"],
max_seqlen_q=batch["max_seqlen"],
max_seqlen_kv=batch["max_seqlen"],
qkv_format="thd",
cu_seqlens_host=batch["cu_seqlens_host"],
)
batch["packed_seq_params"] = packed_seq_params
return packed_seq_params
+4
View File
@@ -266,6 +266,7 @@ def get_batch(
tokens = F.pad(tokens, (0, pad), value=pad_token_id)
cu_seqlens_list.append(cu_seqlens_list[-1] + pad)
cu_seqlens_host = tuple(cu_seqlens_list)
cu_seqlens = torch.tensor(cu_seqlens_list, dtype=torch.int, device=torch.cuda.current_device())
tokens = tokens.chunk(cp_size, dim=0)[cp_rank]
else:
@@ -285,6 +286,7 @@ def get_batch(
cu_seqlens.append(cu_seqlens[-1] + pad)
# thd requires the cu_seqlens to be of the origin length
cu_seqlens_host = tuple(boundary * cp_size for boundary in cu_seqlens)
cu_seqlens = torch.tensor(cu_seqlens, dtype=torch.int).cuda() * cp_size
max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max().item()
@@ -292,6 +294,8 @@ def get_batch(
tokens = tokens.unsqueeze(0)
batch["cu_seqlens"] = cu_seqlens
# the same boundaries on the host, for consumers that must not sync the device to read them
batch["cu_seqlens_host"] = cu_seqlens_host
batch["max_seqlen"] = max_seqlen
else:
raise ValueError(f"Unsupported qkv_format: {qkv_format}")
@@ -0,0 +1,237 @@
"""Balance causal per-row work across contiguous context-parallel ranks.
Every sequence is cut into 2 * cp chunks and rank r scores chunks 2r and 2 * cp - 1 - 2r, pairing an
early chunk with a late one; for one sequence the exchange is a single pairwise swap. Every rank
derives the same plan from the global sequence lengths. ``send_rows_to_scorers`` starts the
exchange, ``wait()`` hands over the rows at ``scored_positions`` and ``return_to_owners`` sends the
results back; ``LocalRows`` is the same interface without an exchange. Nothing carries autograd.
"""
from dataclasses import dataclass
from functools import lru_cache
import torch
import torch.distributed as dist
from torch import Tensor
@dataclass(frozen=True)
class RowBalancePlan:
"""This rank's side of a balanced exchange; every field is identical in meaning on all ranks.
Attributes:
send_rows: local row indices in send order (grouped by scoring rank, ascending position).
input_splits: rows sent to each rank.
output_splits: rows received from each rank.
scored_positions: global stream positions of the rows this rank scores, in received order.
"""
send_rows: Tensor
input_splits: tuple[int, ...]
output_splits: tuple[int, ...]
scored_positions: Tensor
@property
def num_scored(self) -> int:
return sum(self.output_splits)
def scoring_rank_of_chunk(chunk: int, cp_size: int) -> int:
"""Chunk 2r goes to rank r and chunk 2r + 1 to rank cp - 1 - r: an early chunk pairs with a late one.
This is Miles' zigzag layout with the ranks relabeled so that each rank keeps its own first chunk.
"""
return chunk // 2 if chunk % 2 == 0 else cp_size - 1 - chunk // 2
# enough for every micro-batch in flight under pipeline parallelism; a miss only rebuilds the plan
@lru_cache(maxsize=32)
def plan_causal_row_balance(
seq_lens: tuple[int, ...],
*,
cp_rank: int,
cp_size: int,
device: torch.device | str,
min_gain: float = 0.1,
) -> RowBalancePlan | None:
"""The balanced exchange for a stream of ``seq_lens`` split contiguously over ``cp_size`` ranks.
``seq_lens`` must tile the whole stream, padding included, and the stream must split evenly
over the ranks. Returns None when balancing would not lower the busiest rank's causal cost
(``offset + 1`` per row) by ``min_gain``. A rank may be left with no rows to score. Host work is
proportional to sequences times ranks; the row indices are generated on ``device``. Plans are
memoized, so every layer and every recompute of a micro-batch shares one.
"""
if min(seq_lens) < 0:
raise ValueError(f"sequence lengths must be non-negative, got {min(seq_lens)}")
total = sum(seq_lens)
if total % cp_size:
raise ValueError(f"a stream of {total} rows does not split evenly over {cp_size} ranks")
rank_rows = total // cp_size
pieces = _pieces(seq_lens, cp_size, rank_rows)
if not _worth_balancing(pieces, cp_size, min_gain):
return None
return _plan_for_rank(pieces, cp_rank=cp_rank, cp_size=cp_size, rank_rows=rank_rows, device=device)
@dataclass(frozen=True)
class _Piece:
"""Stream rows ``[start, start + rows)``, held by ``owner`` and scored by ``scorer``.
``offset`` is where the piece starts inside its own sequence.
"""
start: int
rows: int
offset: int
owner: int
scorer: int
def _pieces(seq_lens, cp_size: int, rank_rows: int) -> list[_Piece]:
"""Every sequence's 2 * cp chunks in stream order, cut where a rank's rows end.
Chunk c of a sequence of n rows holds offsets [ceil(c n / 2cp), ceil((c + 1) n / 2cp)), so it is
at most half a rank's rows, rounded up, and crosses at most one rank boundary.
"""
n_chunks = 2 * cp_size
pieces = []
seq_start = 0
for n in seq_lens:
for chunk in range(n_chunks):
lo, hi = _ceil_div(chunk * n, n_chunks), _ceil_div((chunk + 1) * n, n_chunks)
scorer = scoring_rank_of_chunk(chunk, cp_size)
while lo < hi:
owner = (seq_start + lo) // rank_rows
cut = min(hi, (owner + 1) * rank_rows - seq_start)
pieces.append(_Piece(start=seq_start + lo, rows=cut - lo, offset=lo, owner=owner, scorer=scorer))
lo = cut
seq_start += n
return pieces
def _worth_balancing(pieces: list[_Piece], cp_size: int, min_gain: float) -> bool:
"""Whether scoring each piece on its scorer cuts the busiest rank's causal cost by ``min_gain``."""
contiguous, balanced = [0] * cp_size, [0] * cp_size
for piece in pieces:
end = piece.offset + piece.rows
cost = (end * (end + 1) - piece.offset * (piece.offset + 1)) // 2 # sum of offset + 1 over the piece
contiguous[piece.owner] += cost
balanced[piece.scorer] += cost
return max(balanced) <= (1 - min_gain) * max(contiguous)
def _plan_for_rank(
pieces: list[_Piece], *, cp_rank: int, cp_size: int, rank_rows: int, device: torch.device | str
) -> RowBalancePlan:
# sorted() is stable: rows bound for one rank keep ascending position, the order the receiver expects
sent = sorted((piece for piece in pieces if piece.owner == cp_rank), key=lambda piece: piece.scorer)
received = [piece for piece in pieces if piece.scorer == cp_rank]
input_splits, output_splits = [0] * cp_size, [0] * cp_size
for piece in sent:
input_splits[piece.scorer] += piece.rows
for piece in received:
output_splits[piece.owner] += piece.rows
return RowBalancePlan(
send_rows=_concat_ranges([p.start - cp_rank * rank_rows for p in sent], [p.rows for p in sent], device),
input_splits=tuple(input_splits),
output_splits=tuple(output_splits),
scored_positions=_concat_ranges([p.start for p in received], [p.rows for p in received], device),
)
def _ceil_div(numerator: int, denominator: int) -> int:
return -(-numerator // denominator)
def _concat_ranges(starts: list[int], lengths: list[int], device: torch.device | str) -> Tensor:
"""``cat([arange(s, s + n) for s, n in zip(starts, lengths)])``, expanded on ``device``.
Only the per-range table crosses to the device; the rows are generated there.
"""
starts_t, lengths_t = torch.tensor([starts, lengths], dtype=torch.int64)
total = int(lengths_t.sum())
table = torch.stack([starts_t - (torch.cumsum(lengths_t, 0) - lengths_t), lengths_t])
if torch.device(device).type == "cuda":
table = table.pin_memory() # a copy from pageable memory would wait for the stream to drain
base, lengths_t = table.to(device, non_blocking=True)
return base.repeat_interleave(lengths_t, output_size=total) + torch.arange(total, device=device)
class RowExchange:
"""One balanced exchange: rows in flight to this rank's scorer, and the way back for results."""
def __init__(self, plan: RowBalancePlan, cp_group: dist.ProcessGroup, received: list[Tensor], works: list):
self.plan = plan
self._cp_group = cp_group
self._received = received
self._works = works
@property
def scored_positions(self) -> Tensor:
return self.plan.scored_positions
def wait(self) -> list[Tensor]:
"""The sent tensors' rows at ``plan.scored_positions``, in that order."""
for work in self._works:
work.wait()
self._works = []
return self._received
def return_to_owners(self, results: Tensor, *, dim: int = 0) -> Tensor:
"""Per-row results for the scored rows (along ``dim``) back to the local row order."""
if results.requires_grad:
raise ValueError("the row exchange carries no autograd; return detached results")
self.wait()
rows = results.movedim(dim, 0).contiguous()
assert rows.shape[0] == self.plan.num_scored, f"{rows.shape[0]} results for {self.plan.num_scored} rows"
received = rows.new_empty((self.plan.send_rows.numel(), *rows.shape[1:]))
dist.all_to_all_single(
received,
rows,
output_split_sizes=list(self.plan.input_splits),
input_split_sizes=list(self.plan.output_splits),
group=self._cp_group,
)
# rows come back grouped by scoring rank in ascending position, i.e. in send order
local = torch.empty_like(received)
local[self.plan.send_rows] = received
return local.movedim(0, dim).contiguous()
class LocalRows:
"""``RowExchange`` without the exchange: this rank scores its own rows, already in local order."""
def __init__(self, tensors: list[Tensor], scored_positions: Tensor):
self.scored_positions = scored_positions
self._tensors = tensors
def wait(self) -> list[Tensor]:
return self._tensors
def return_to_owners(self, results: Tensor, *, dim: int = 0) -> Tensor:
return results
def send_rows_to_scorers(tensors: list[Tensor], plan: RowBalancePlan, cp_group: dist.ProcessGroup) -> RowExchange:
"""Start sending each tensor's local rows (dim 0) to the ranks that score them."""
for rows in tensors:
if rows.requires_grad:
raise ValueError("the row exchange carries no autograd; send detached tensors")
assert rows.shape[0] == plan.send_rows.numel(), "every tensor needs one row per local position"
received, works = [], []
for rows in tensors:
buffer = rows.new_empty((plan.num_scored, *rows.shape[1:]))
works.append(
dist.all_to_all_single(
buffer,
rows.index_select(0, plan.send_rows),
output_split_sizes=list(plan.output_splits),
input_split_sizes=list(plan.input_splits),
group=cp_group,
async_op=True,
)
)
received.append(buffer)
return RowExchange(plan, cp_group, received, works)
@@ -11,6 +11,7 @@ absolute, so the KV layout is unchanged.
"""
from dataclasses import dataclass
from itertools import pairwise
import torch
import torch.distributed as dist
@@ -21,12 +22,14 @@ from torch import Tensor
class ThdLayout:
"""How this rank's packed stream is laid out; ``None`` stands for the unpacked one.
The first three fields come from the packed sequence parameters. The rest are filled in as
The first four fields come from the packed sequence parameters; ``seq_lens`` holds the
segment lengths on the host, so reading them costs no device sync. The rest are filled in as
the forward runs: ``cu_seqlens_compressed`` before the compressor is called, and the
compaction ones only under CP, where a compressed group can straddle the split.
"""
cu_seqlens: Tensor
seq_lens: tuple[int, ...]
global_start: int
max_seqlen: int
hidden_compact: Tensor | None = None
@@ -39,8 +42,10 @@ class ThdLayout:
"""This rank's layout, or None for any format other than thd."""
if packed_seq_params is None or packed_seq_params.qkv_format != "thd":
return None
cu_seqlens_host = packed_seq_params.cu_seqlens_host
return cls(
cu_seqlens=packed_seq_params.cu_seqlens_q,
seq_lens=tuple(end - start for start, end in pairwise(cu_seqlens_host)),
# CP splits the packed stream contiguously, so this rank's rows start here globally.
global_start=cp_rank * seqlen_local,
max_seqlen=packed_seq_params.max_seqlen_q,
@@ -57,9 +62,14 @@ def batch_of_row(cu_seqlens: Tensor, total_rows: int, global_start: int = 0) ->
Returns:
``[total_rows]`` int64.
"""
n_seg = cu_seqlens.size(0) - 1
row_idx = torch.arange(total_rows, device=cu_seqlens.device, dtype=torch.int64) + global_start
return torch.bucketize(row_idx, cu_seqlens[1:], right=True).clamp(max=max(n_seg - 1, 0))
return segment_of_positions(cu_seqlens, row_idx)
def segment_of_positions(cu_seqlens: Tensor, positions: Tensor) -> Tensor:
"""Segment index owning each global stream position; positions past the end clamp to the last."""
n_seg = cu_seqlens.size(0) - 1
return torch.bucketize(positions, cu_seqlens[1:], right=True).clamp(max=max(n_seg - 1, 0))
def compressed_cu_seqlens(cu_seqlens: Tensor, ratio: int) -> Tensor:
@@ -128,22 +138,32 @@ def get_compress_cu_seqlens_thd(
total_tokens: int,
global_start: int = 0,
) -> tuple[Tensor, Tensor]:
"""Get the compressed rows each packed query may see, as a half-open range.
"""``compress_bounds_at_positions`` for this rank's contiguous rows."""
positions = torch.arange(total_tokens, device=cu_seqlens.device) + global_start
return compress_bounds_at_positions(cu_seqlens, cu_seqlens_compressed, positions, ratio=ratio)
def compress_bounds_at_positions(
cu_seqlens: Tensor,
cu_seqlens_compressed: Tensor,
positions: Tensor,
*,
ratio: int,
) -> tuple[Tensor, Tensor]:
"""Get the compressed rows the packed queries at global stream ``positions`` may see.
A query sees its own segment only, up to ``(pos_in_seg + 1) // ratio`` and never past
what that segment produced; the BSHD ``ks = 0`` convention would let it score entries
of earlier segments once all samples share one flat stream.
Returns:
``(cu_ks, cu_ke)`` int32 ``[total_tokens]``, indices into the compressed keys alone,
not the concatenated KV. The indexer kernel takes them as is;
``(cu_ks, cu_ke)`` int32, one half-open range per position, indices into the compressed
keys alone, not the concatenated KV. The indexer kernel takes them as is;
``get_compress_topk_idxs_thd`` expands them into explicit indices.
"""
device = cu_seqlens.device
batch_ids = batch_of_row(cu_seqlens, total_tokens, global_start)
token_idx = torch.arange(total_tokens, device=device) + global_start
pos_in_seg = token_idx - cu_seqlens[batch_ids]
positions = positions.long()
batch_ids = segment_of_positions(cu_seqlens, positions)
pos_in_seg = positions - cu_seqlens[batch_ids]
cu_ks = cu_seqlens_compressed[batch_ids]
cu_ke = torch.minimum(cu_ks + (pos_in_seg + 1) // ratio, cu_seqlens_compressed[batch_ids + 1])
return cu_ks.int(), cu_ke.int()
@@ -8,16 +8,19 @@ from megatron.core.transformer.module import MegatronModule
from megatron.core.transformer.transformer_config import TransformerConfig
from miles.utils.replay_base import indexer_replay_manager
from miles_plugins.models.deepseek_v4.ops.compressor import DeepSeekV4Compressor
from miles_plugins.models.deepseek_v4.ops.cp_utils import all_gather_cp, get_freqs_cis_for_cp
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.cp_row_balance import (
LocalRows,
RowBalancePlan,
RowExchange,
plan_causal_row_balance,
send_rows_to_scorers,
)
from miles_plugins.models.deepseek_v4.ops.cp_utils import all_gather_cp, get_freqs_cis_for_cp, get_q_positions_for_cp
from miles_plugins.models.deepseek_v4.ops.kernel.tilelang_indexer_fwd import batched_indexer_fwd
from miles_plugins.models.deepseek_v4.ops.qat import fp8_simulate_qat
from miles_plugins.models.deepseek_v4.ops.rope import apply_rotary_emb, wrapped_precompute_freqs_cis
from miles_plugins.models.deepseek_v4.ops.thd_utils import ThdLayout, get_compress_cu_seqlens_thd, get_q_positions_thd
from miles_plugins.models.deepseek_v4.ops.thd_utils import ThdLayout, compress_bounds_at_positions, get_q_positions_thd
from miles_plugins.models.deepseek_v4.ops.utils import rotate_activation
from miles_plugins.models.dsa_topk import get_dsa_topk_fn
@@ -95,7 +98,8 @@ class V4Indexer(MegatronModule):
thd_layout: packed-stream layout, or None when unpacked
Returns:
topk_indices: [batch, seqlen, index_topk] int64
topk_indices: [batch, seqlen, min(index_topk, n_kv)] int32, or index_topk columns of -1
when no segment has a compressed key
"""
# =========================================
@@ -134,50 +138,107 @@ class V4Indexer(MegatronModule):
if self.use_fp8_qat:
q = fp8_simulate_qat(q, 128)
weights, _ = self.linear_weights_proj(x)
softmax_scale = self.index_head_dim**-0.5
weights = (weights * (self.index_n_heads**-0.5) * softmax_scale).float()
# Balance the causal scoring work over contiguous CP (cp_row_balance).
# Replay data holds each rank's own rows, so replay scores them where they are.
balance = cp_size > 1 and cp_group is not None and not indexer_replay_manager.enabled
# started before the compressor to overlap it; unpacked, its CP all-gathers wait for the exchange
exchange = start_row_exchange(q, weights, thd_layout, cp_group, balance=balance)
del q, weights # scored from exchange.wait()
pre_grouped = thd_layout is not None and thd_layout.compressed_group_ids is not None
k = self.compressor(thd_layout.hidden_compact if pre_grouped else x, thd_layout)
if k is None:
# Nothing to score when no segment reaches compress_ratio; -1 leaves each query
# on its sliding window.
# Nothing to score when no segment reaches compress_ratio; -1 leaves each query on its
# sliding window. The compressor returns None only without CP, so no rows are in flight.
assert isinstance(exchange, LocalRows), "the compressor returned no keys while rows were in flight"
return torch.full((bsz, seqlen, self.index_topk), -1, dtype=torch.int32, device=x.device)
weights, _ = self.linear_weights_proj(x)
softmax_scale = self.index_head_dim**-0.5
weights = weights * (self.index_n_heads**-0.5) * softmax_scale
if cp_size > 1 and cp_group is not None:
k = all_gather_cp(k, dim=0, cp_group=cp_group)
if thd_layout is not None and thd_layout.seq_to_rank_row is not None:
# Per-row bounds are sequence-major, so reorder the rank-major gather first.
k = k.index_select(0, thd_layout.seq_to_rank_row.clamp(min=0).long())
seqlen_global = seqlen * cp_size
seqlen_kv = k.shape[0]
if thd_layout is None:
cu_ks, cu_ke = _make_causal_cu_seqlens(seqlen_global, seqlen_kv, self.compress_ratio, q.device)
# cu_seqlens are for global positions; slice to local query positions
if cp_size > 1 and cp_group is not None:
cp_rank = cp_group.rank()
cu_ks = cu_ks[cp_rank * seqlen : (cp_rank + 1) * seqlen]
cu_ke = cu_ke[cp_rank * seqlen : (cp_rank + 1) * seqlen]
else:
cu_ks, cu_ke = get_compress_cu_seqlens_thd(
thd_layout.cu_seqlens,
thd_layout.cu_seqlens_compressed,
ratio=self.compress_ratio,
total_tokens=seqlen,
global_start=thd_layout.global_start,
)
index_scores = batched_indexer_fwd(q, k, weights.float(), cu_ks, cu_ke)
# index_scores: [batch, seqlen, n_kv]; topk over the KV dim. Route through the indexer
# replay manager (flattened to [n_tokens, n_kv], matching the record/replay convention) so
# RL replay can pin the rollout's top-k picks. get_topk_fn is transparent when disabled.
topk_count = min(self.index_topk, index_scores.size(-1))
# The manager records and replays flattened [tokens, n_kv] scores, matching the MoE seam.
bsz, seqlen, n_kv = index_scores.shape
# RL replay can pin the rollout's top-k picks here; get_topk_fn is transparent when disabled.
topk_fn = indexer_replay_manager.get_topk_fn(get_dsa_topk_fn(self.topk_backend), return_probs=False)
topk_indices = topk_fn(index_scores.reshape(bsz * seqlen, n_kv), topk_count)
topk_indices = topk_indices.reshape(bsz, seqlen, topk_count)
return topk_for_local_rows(
exchange,
k,
thd_layout,
compress_ratio=self.compress_ratio,
index_topk=self.index_topk,
topk_fn=topk_fn,
)
return topk_indices
def start_row_exchange(q, weights, thd_layout, cp_group, *, balance: bool) -> RowExchange | LocalRows:
"""Start sending this rank's indexer rows to their scoring ranks, or keep them if balancing does not pay."""
tensors = [q.detach(), weights.detach()]
plan = _row_balance_plan(q.shape[0], thd_layout, cp_group, q.device) if balance else None
if plan is None:
cp_size = cp_group.size() if cp_group is not None else 1
positions = get_q_positions_for_cp(q.shape[0], cp_size=cp_size, cp_group=cp_group, device=q.device)
return LocalRows(tensors, positions)
return send_rows_to_scorers(tensors, plan, cp_group)
def topk_for_local_rows(exchange, k, thd_layout, *, compress_ratio, index_topk, topk_fn):
"""The top-k picks for this rank's rows, in local order, scored on the rank ``exchange`` sent them to."""
q, weights = exchange.wait()
topk_indices = indexer_topk(
q,
k,
weights,
exchange.scored_positions,
thd_layout,
compress_ratio=compress_ratio,
index_topk=index_topk,
topk_fn=topk_fn,
)
# [batch, rows, topk]: the picks go back along the row dim
return exchange.return_to_owners(topk_indices, dim=1)
def indexer_topk(q, k, weights, positions, thd_layout, *, compress_ratio, index_topk, topk_fn):
"""Score the query rows at global stream ``positions`` against their visible compressed keys.
Args:
q: [rows, batch, heads, head_dim] index queries of those rows
k: [n_kv, batch, head_dim] every compressed key of the stream, sequence-major under THD
weights: [rows, batch, heads] fp32 head weights
positions: [rows] global stream positions of the rows
thd_layout: packed-stream layout, or None when unpacked
Returns:
[batch, rows, min(index_topk, n_kv)] int32 compressed-key indices
"""
if q.shape[0] == 0:
# a balanced plan can leave a rank nothing to score, and TileLang cannot launch an empty grid
return torch.empty(q.shape[1], 0, min(index_topk, k.shape[0]), dtype=torch.int32, device=q.device)
if thd_layout is None:
cu_ks = torch.zeros_like(positions, dtype=torch.int32)
cu_ke = ((positions + 1) // compress_ratio).int()
else:
cu_ks, cu_ke = compress_bounds_at_positions(
thd_layout.cu_seqlens, thd_layout.cu_seqlens_compressed, positions, ratio=compress_ratio
)
index_scores = batched_indexer_fwd(q, k, weights, cu_ks, cu_ke)
bsz, rows, n_kv = index_scores.shape
topk_count = min(index_topk, n_kv)
# flattened to [n_tokens, n_kv], the record/replay convention shared with the MoE seam
topk_indices = topk_fn(index_scores.reshape(bsz * rows, n_kv), topk_count)
return topk_indices.reshape(bsz, rows, topk_count)
def _row_balance_plan(seqlen_local, thd_layout, cp_group, device) -> RowBalancePlan | None:
"""This micro-batch's balanced exchange; every CP rank derives the same one."""
cp_size = cp_group.size()
total_rows = seqlen_local * cp_size
# each batch row of an unpacked sample is one sequence
seq_lens = (total_rows,) if thd_layout is None else thd_layout.seq_lens
assert sum(seq_lens) == total_rows, f"segment lengths cover {sum(seq_lens)} rows of a {total_rows}-row stream"
return plan_causal_row_balance(seq_lens, cp_rank=cp_group.rank(), cp_size=cp_size, device=device)
+37 -56
View File
@@ -115,7 +115,9 @@ class ScriptArgs(command_utils.ExecuteTrainConfig):
dsa_kernel_backend: Literal["none", "tilelang", "cudnn"] | None = None
optimizer_offload: bool = True
use_fault_tolerance: bool = True
cp_size: int = 1
# None runs each recipe's own CP. Only the single-node miles impl lets it vary (TP takes the GPUs
# CP leaves, CP split with --allgather-cp); every other recipe accepts only its own CP size.
cp_size: int | None = None
# debug configs
dump_details: bool = False
@@ -153,6 +155,7 @@ class ScriptArgs(command_utils.ExecuteTrainConfig):
assert not (self.train_mxfp8 or self.rollout_mxfp8), "train_mxfp8/rollout_mxfp8 require Blackwell"
assert self.rollout_num_nodes >= 0
assert self.rollout_num_nodes < self.num_nodes
assert self.cp_size is None or self.cp_size >= 1, f"cp_size must be at least 1, got {self.cp_size}"
self.colocate = self.rollout_num_nodes == 0
self.actor_num_nodes = self.num_nodes - self.rollout_num_nodes
self.actor_num_gpus_per_node = self.num_gpus_per_node
@@ -399,22 +402,13 @@ def _get_parallel_config(args: ScriptArgs) -> str:
# Single-node smoke-test configs
if actor_num_nodes == 1:
if args.dsv4_impl == "megatron":
# dsv4_hybrid needs cp_partition_mode='contiguous' for CP>1, which miles does not set
# The plugin rejects TP>1; the TP ranks go to DP instead.
return (
"--tensor-model-parallel-size 1 "
"--pipeline-model-parallel-size 1 "
"--context-parallel-size 1 "
f"--expert-model-parallel-size {actor_num_gpus_per_node} "
"--expert-tensor-parallel-size 1 "
)
return (
f"--tensor-model-parallel-size {actor_num_gpus_per_node} "
"--sequence-parallel "
"--pipeline-model-parallel-size 1 "
"--context-parallel-size 1 "
f"--expert-model-parallel-size {actor_num_gpus_per_node} "
"--expert-tensor-parallel-size 1 "
)
return _parallel_flags(args, tp=1, cp=1, ep=actor_num_gpus_per_node)
cp_size = args.cp_size or 1
if actor_num_gpus_per_node % cp_size:
raise NotImplementedError(f"cp_size={cp_size} does not divide {actor_num_gpus_per_node} GPUs")
return _parallel_flags(args, tp=actor_num_gpus_per_node // cp_size, cp=cp_size, ep=actor_num_gpus_per_node)
if actor_num_gpus_per_node == 4:
if total_gpus == 32: # 8 nodes x 4 GPUs
@@ -423,50 +417,14 @@ def _get_parallel_config(args: ScriptArgs) -> str:
# CP>1, which no launcher exercises yet -- so the TP and CP ranks both go
# to DP. max-tokens-per-gpu below doubles to keep the per-micro-batch
# budget (max_tokens_per_gpu * cp_size) equal to the miles recipe's.
return (
"--tensor-model-parallel-size 1 "
"--pipeline-model-parallel-size 8 "
"--decoder-first-pipeline-num-layers 4 "
"--decoder-last-pipeline-num-layers 3 "
"--context-parallel-size 1 "
"--expert-model-parallel-size 4 "
"--expert-tensor-parallel-size 1 "
)
return (
"--tensor-model-parallel-size 2 "
"--sequence-parallel "
"--pipeline-model-parallel-size 8 "
"--decoder-first-pipeline-num-layers 4 "
"--decoder-last-pipeline-num-layers 3 "
"--context-parallel-size 2 "
"--allgather-cp "
"--expert-model-parallel-size 4 "
"--expert-tensor-parallel-size 1 "
)
return _parallel_flags(args, tp=1, pp=8, pp_edge_layers=(4, 3), cp=1, ep=4)
return _parallel_flags(args, tp=2, pp=8, pp_edge_layers=(4, 3), cp=2, ep=4)
if actor_num_gpus_per_node == 8:
if total_gpus == 64: # 8 nodes x 8 GPUs
return (
"--tensor-model-parallel-size 8 "
"--sequence-parallel "
"--pipeline-model-parallel-size 8 "
"--decoder-first-pipeline-num-layers 4 "
"--decoder-last-pipeline-num-layers 3 "
"--context-parallel-size 1 "
"--expert-model-parallel-size 8 "
"--expert-tensor-parallel-size 1 "
)
return _parallel_flags(args, tp=8, pp=8, pp_edge_layers=(4, 3), cp=1, ep=8)
elif total_gpus == 256: # 32 nodes x 8 GPUs (Pro)
return (
"--tensor-model-parallel-size 8 "
"--sequence-parallel "
"--pipeline-model-parallel-size 8 "
"--decoder-first-pipeline-num-layers 7 "
"--decoder-last-pipeline-num-layers 6 "
"--context-parallel-size 1 "
"--expert-model-parallel-size 32 "
"--expert-tensor-parallel-size 1 "
)
return _parallel_flags(args, tp=8, pp=8, pp_edge_layers=(7, 6), cp=1, ep=32)
raise NotImplementedError(
f"No pre-set parallel config for {total_gpus} GPUs. "
@@ -474,6 +432,29 @@ def _get_parallel_config(args: ScriptArgs) -> str:
)
def _parallel_flags(
args: ScriptArgs, *, tp: int, cp: int, ep: int, pp: int = 1, pp_edge_layers: tuple[int, int] | None = None
) -> str:
"""One recipe's parallel flags; a recipe runs its own CP size, so ``cp_size`` must be unset or equal."""
if args.cp_size not in (None, cp):
raise NotImplementedError(
f"cp_size={args.cp_size} is untested here: this recipe (--dsv4-impl {args.dsv4_impl}, "
f"{args.actor_num_nodes}x{args.actor_num_gpus_per_node} GPUs) runs CP{cp}"
)
flags = [f"--tensor-model-parallel-size {tp}"]
if tp > 1:
flags.append("--sequence-parallel")
flags.append(f"--pipeline-model-parallel-size {pp}")
if pp_edge_layers is not None:
first, last = pp_edge_layers
flags += [f"--decoder-first-pipeline-num-layers {first}", f"--decoder-last-pipeline-num-layers {last}"]
flags.append(f"--context-parallel-size {cp}")
if cp > 1:
flags.append("--allgather-cp") # DeepSeek V4 rejects the zigzag CP split
flags += [f"--expert-model-parallel-size {ep}", "--expert-tensor-parallel-size 1"]
return "".join(f"{flag} " for flag in flags)
def _train(args: ScriptArgs):
U = args.create_backend()
if args.train_mxfp8 or args.rollout_mxfp8:
@@ -0,0 +1,40 @@
"""Smoke run of DeepSeek-V4-Flash's CP path: the 4-layer RL case on four GPUs (miles impl).
The base case (test_deepseek_v4_flash_4layer_ci.py) at TP2 with sequence parallelism, CP2 with the
all-gather CP split, EP4. Each micro-batch holds one unpacked sample, so the CSA layer always takes
the load-balanced indexer path. The run catches crashes, hangs and train-side nondeterminism, not
wrong picks: its train-rollout metrics are tracked but not gated, and short GSM8K samples have fewer
compressed keys than the top-k keeps. tests/fast-gpu/test_dsv4_indexer_cp_balance.py checks the picks.
"""
import dataclasses
import os
from tests.ci.ci_register import register_cuda_ci
from tests.ci.metric_history import register_ci_gate
from tests.e2e.megatron.model_scripts import test_deepseek_v4_flash_4layer_ci as base
register_cuda_ci(
est_time=1900, suite="stage-c-4-gpu-h200", labels=["megatron", "model-scripts"], hardware=["hopper", "blackwell"]
)
register_ci_gate(metric_key="train/grad_norm")
register_ci_gate(metric_key="train/ppo_kl")
register_ci_gate(metric_key="train/train_rollout_logprob_abs_diff")
register_ci_gate(metric_key="train/train_rollout_kl")
register_ci_gate(metric_key="rollout/raw_reward")
prepare = base.prepare
execute = base.execute
def _args():
return dataclasses.replace(base._args(), cp_size=2)
if __name__ == "__main__":
args = _args()
prepare(args)
for proxy_var in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"):
os.environ.pop(proxy_var, None)
execute(args)
@@ -0,0 +1,177 @@
"""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()
@@ -0,0 +1,38 @@
"""`get_batch`'s host copy of `cu_seqlens` must describe the same packed stream as the device tensor."""
from types import SimpleNamespace
import pytest
import torch
from miles.backends.training_utils import cp_utils
from miles.backends.training_utils import data as data_utils
@pytest.mark.parametrize("allgather_cp", [False, True])
@pytest.mark.parametrize("cp_size", [1, 2, 4])
def test_the_host_cu_seqlens_match_the_device_tensor(
monkeypatch: pytest.MonkeyPatch, allgather_cp: bool, cp_size: int
):
"""The DSv4 row balancer plans from the host copy to avoid a device sync, so a drift would misroute rows."""
lengths = [13, 11, 5, 1]
rollout = {
"tokens": [torch.arange(1, n + 1) for n in lengths],
"loss_masks": [torch.ones(n // 2, dtype=torch.int) for n in lengths],
"total_lengths": lengths,
"response_lengths": [n // 2 for n in lengths],
}
monkeypatch.setattr(torch.cuda, "current_device", lambda: torch.device("cpu"))
monkeypatch.setattr(torch.Tensor, "cuda", lambda self, *args, **kwargs: self)
for cp_rank in range(cp_size):
state = SimpleNamespace(cp=SimpleNamespace(rank=cp_rank, size=cp_size), tp=SimpleNamespace(rank=0, size=1))
monkeypatch.setattr(data_utils, "get_parallel_state", lambda state=state: state)
monkeypatch.setattr(cp_utils, "get_parallel_state", lambda state=state: state)
batch = data_utils.get_batch(
data_utils.DataIterator(rollout, micro_batch_size=len(lengths)),
list(rollout),
pad_multiplier=7,
qkv_format="thd",
allgather_cp=allgather_cp,
)
assert batch["cu_seqlens_host"] == tuple(batch["cu_seqlens"].tolist())
@@ -9,6 +9,72 @@ from tests.fast.launch_scripts.py_harness import (
)
_SINGLE_NODE_4LAYER = {
"hardware": "H200",
"model_name": "DeepSeek-V4-Flash-FP8-4layer",
"num_nodes": 1,
"num_gpus_per_node": 4,
}
_EIGHT_NODES = {"hardware": "H200", "num_nodes": 8, "num_gpus_per_node": 4}
_EIGHT_NODES_OF_8 = {"hardware": "H200", "num_nodes": 8, "num_gpus_per_node": 8}
_THIRTY_TWO_NODES_OF_8 = {"hardware": "H200", "num_nodes": 32, "num_gpus_per_node": 8}
def _train_command(monkeypatch, tmp_path, overrides):
freeze_environment(monkeypatch)
recording = install_command_recorder(monkeypatch)
module = import_launch_script(REPO_ROOT / "scripts/run_deepseek_v4.py")
call_entrypoint(module, "train", overrides, sandbox=tmp_path)
return recording.commands[-1]
@pytest.mark.parametrize(
("cp_size", "expected"),
[
(None, ["--tensor-model-parallel-size 4", "--sequence-parallel", "--context-parallel-size 1"]),
(2, ["--tensor-model-parallel-size 2", "--sequence-parallel", "--context-parallel-size 2", "--allgather-cp"]),
(4, ["--tensor-model-parallel-size 1", "--context-parallel-size 4", "--allgather-cp"]),
],
)
def test_single_node_miles_impl_splits_gpus_between_tp_and_cp(monkeypatch, tmp_path, cp_size, expected):
"""DSV4 CP must use the all-gather split; TP takes the GPUs CP leaves, with SP only above TP1."""
overrides = _SINGLE_NODE_4LAYER | {"dsv4_impl": "miles"} | ({"cp_size": cp_size} if cp_size else {})
command = _train_command(monkeypatch, tmp_path, overrides)
for flag in expected:
assert flag in command
tp_size = 4 // (cp_size or 1)
assert ("--allgather-cp" in command) == (tp_size < 4)
assert ("--sequence-parallel" in command) == (tp_size > 1)
# Every recipe that pins its CP size, with that size.
_FIXED_CP_RECIPES = [
(_SINGLE_NODE_4LAYER | {"dsv4_impl": "megatron"}, 1),
(_EIGHT_NODES | {"dsv4_impl": "megatron"}, 1),
(_EIGHT_NODES | {"dsv4_impl": "miles"}, 2),
(_EIGHT_NODES_OF_8 | {"dsv4_impl": "megatron"}, 1),
(_EIGHT_NODES_OF_8 | {"dsv4_impl": "miles"}, 1),
(_THIRTY_TWO_NODES_OF_8 | {"dsv4_impl": "miles"}, 1),
]
@pytest.mark.parametrize(("recipe", "recipe_cp_size"), _FIXED_CP_RECIPES)
def test_recipes_with_a_fixed_cp_size_accept_it_and_reject_another(monkeypatch, tmp_path, recipe, recipe_cp_size):
command = _train_command(monkeypatch, tmp_path, recipe | {"cp_size": recipe_cp_size})
assert f"--context-parallel-size {recipe_cp_size}" in command
with pytest.raises(NotImplementedError, match="is untested here"):
_train_command(monkeypatch, tmp_path, recipe | {"cp_size": recipe_cp_size * 2})
def test_a_node_count_without_a_recipe_reports_the_missing_recipe(monkeypatch, tmp_path):
with pytest.raises(NotImplementedError, match="No pre-set parallel config"):
_train_command(
monkeypatch, tmp_path, {"hardware": "H200", "num_nodes": 2, "num_gpus_per_node": 8, "cp_size": 2}
)
@pytest.mark.parametrize(
("overrides", "expected_size"),
[
@@ -0,0 +1,209 @@
"""CPU tests for the causal row balancer shared by CP indexers.
Every rank's plan is built on the host and the two all-to-alls are replayed in Python, so the
routing, the balance and the gate are checked without GPUs; the exchange then runs over a gloo
process group. Rows carry their own global position, which makes a row routed to the wrong rank
or returned out of order visible.
"""
import random
from functools import partial
import pytest
import torch
import torch.distributed as dist
from tests.ci.ci_register import register_cpu_ci
from tests.fast.dist_utils import init_gloo, run_multiprocess
from miles_plugins.models.deepseek_v4.ops.cp_row_balance import (
plan_causal_row_balance,
scoring_rank_of_chunk,
send_rows_to_scorers,
)
register_cpu_ci(est_time=15, suite="stage-a-cpu", labels=[])
def _plans(seq_lens, cp_size, min_gain=0.1):
return [
plan_causal_row_balance(tuple(seq_lens), cp_rank=rank, cp_size=cp_size, device="cpu", min_gain=min_gain)
for rank in range(cp_size)
]
def _reference_plan(seq_lens, cp_rank, cp_size, min_gain):
"""The balancer spelled out one row at a time: each row's chunk, scoring rank and owner."""
total = sum(seq_lens)
rank_rows = total // cp_size
lens = torch.tensor(seq_lens)
positions = torch.arange(total)
seq = torch.repeat_interleave(torch.arange(len(seq_lens)), lens)
offset = positions - (torch.cumsum(lens, 0) - lens)[seq]
chunk = offset * 2 * cp_size // lens[seq]
scorer = torch.where(chunk % 2 == 0, chunk // 2, cp_size - 1 - chunk // 2)
owner = positions // rank_rows
cost = (offset + 1).double()
contiguous = torch.zeros(cp_size, dtype=torch.float64).index_add_(0, owner, cost)
balanced = torch.zeros(cp_size, dtype=torch.float64).index_add_(0, scorer, cost)
if balanced.max() > (1 - min_gain) * contiguous.max():
return None
local_scorer = scorer[cp_rank * rank_rows : (cp_rank + 1) * rank_rows]
mine = scorer == cp_rank
return dict(
send_rows=torch.argsort(local_scorer, stable=True),
input_splits=tuple(torch.bincount(local_scorer, minlength=cp_size).tolist()),
output_splits=tuple(torch.bincount(owner[mine], minlength=cp_size).tolist()),
scored_positions=positions[mine],
)
def _all_to_all(send_buffers, send_splits, recv_splits):
"""What all_to_all_single delivers: rank r gets each peer's slice for r, in peer order."""
cp_size = len(send_buffers)
chunks = [list(buffer.split(list(splits))) for buffer, splits in zip(send_buffers, send_splits, strict=True)]
received = []
for dst in range(cp_size):
pieces = [chunks[src][dst] for src in range(cp_size)]
assert [piece.shape[0] for piece in pieces] == list(recv_splits[dst]), "sender and receiver disagree"
received.append(torch.cat(pieces))
return received
def _costs(seq_lens):
return torch.cat([torch.arange(1, n + 1) for n in seq_lens]).double()
@pytest.mark.parametrize("cp_size", [2, 4, 8])
@pytest.mark.parametrize("seed", range(12))
def test_matches_the_row_by_row_reference(cp_size, seed):
rng = random.Random(seed)
# empty, shorter-than-2cp and long sequences; the last one pads the stream to a multiple of cp
seq_lens = [rng.choice([0, 1, 3, rng.randint(1, 64), rng.randint(1, 3000)]) for _ in range(rng.randint(1, 10))]
seq_lens.append(cp_size - sum(seq_lens) % cp_size)
for min_gain in (-1.0, 0.1):
for rank in range(cp_size):
plan = plan_causal_row_balance(
tuple(seq_lens), cp_rank=rank, cp_size=cp_size, device="cpu", min_gain=min_gain
)
reference = _reference_plan(seq_lens, rank, cp_size, min_gain)
assert (plan is None) == (reference is None), f"gate differs for {seq_lens}"
if plan is not None:
assert torch.equal(plan.send_rows, reference["send_rows"])
assert plan.input_splits == reference["input_splits"]
assert plan.output_splits == reference["output_splits"]
assert torch.equal(plan.scored_positions, reference["scored_positions"])
@pytest.mark.parametrize("cp_size", [2, 4, 8])
@pytest.mark.parametrize("seq_lens", [[4096], [3000, 1096], [2048, 1024, 512, 512], [4000, 90, 6]])
def test_rows_reach_their_scorer_and_come_back_in_order(cp_size, seq_lens):
plans = _plans(seq_lens, cp_size, min_gain=-1.0)
rank_rows = sum(seq_lens) // cp_size
local_rows = [torch.arange(r * rank_rows, (r + 1) * rank_rows) for r in range(cp_size)]
sent = [rows[plan.send_rows] for rows, plan in zip(local_rows, plans, strict=True)]
scored = _all_to_all(sent, [p.input_splits for p in plans], [p.output_splits for p in plans])
for plan, rows in zip(plans, scored, strict=True):
assert torch.equal(rows, plan.scored_positions)
results = [rows * 10 for rows in scored]
returned = _all_to_all(results, [p.output_splits for p in plans], [p.input_splits for p in plans])
for plan, rows, back in zip(plans, local_rows, returned, strict=True):
local = torch.empty_like(back)
local[plan.send_rows] = back
assert torch.equal(local, rows * 10)
@pytest.mark.parametrize("cp_size", [2, 4, 8])
def test_one_sequence_balances_to_within_a_percent(cp_size):
seq_lens = [cp_size * 8192]
costs = _costs(seq_lens)
per_rank = torch.stack([costs[plan.scored_positions].sum() for plan in _plans(seq_lens, cp_size)])
assert per_rank.max() / per_rank.min() < 1.01
@pytest.mark.parametrize("cp_size", [2, 4, 8])
def test_one_sequence_is_a_single_pairwise_swap(cp_size):
"""Rank r keeps its first half and trades its second half with rank cp - 1 - r only."""
rank_rows = 4096
for rank, plan in enumerate(_plans([cp_size * rank_rows], cp_size)):
peers = {dst for dst, rows in enumerate(plan.input_splits) if rows}
assert peers <= {rank, cp_size - 1 - rank}
assert plan.input_splits[rank] == rank_rows // 2
def test_chunk_pairing_pairs_early_with_late():
"""Chunk 2r goes to rank r and chunk 2r + 1 to rank cp - 1 - r, so costs sum to 2cp - 1 per rank."""
cp_size = 4
ranks = [scoring_rank_of_chunk(chunk, cp_size) for chunk in range(2 * cp_size)]
by_rank = [[chunk for chunk, rank in enumerate(ranks) if rank == r] for r in range(cp_size)]
assert by_rank == [[0, 7], [2, 5], [3, 4], [1, 6]]
assert {sum(chunks) for chunks in by_rank} == {2 * cp_size - 1}
def test_a_pack_of_equal_short_documents_keeps_the_contiguous_layout():
"""Each rank already holds the same mix of positions, so moving rows would only add traffic."""
assert all(plan is None for plan in _plans([512] * 64, 4))
def test_a_long_document_in_a_pack_is_balanced():
seq_lens = [24576, 2048, 2048, 2048, 2048]
rank_rows = sum(seq_lens) // 4
costs = _costs(seq_lens)
contiguous = torch.stack([costs[r * rank_rows : (r + 1) * rank_rows].sum() for r in range(4)])
balanced = torch.stack([costs[plan.scored_positions].sum() for plan in _plans(seq_lens, 4)])
assert balanced.max() < 0.75 * contiguous.max()
def test_the_stream_must_split_evenly_over_the_ranks():
with pytest.raises(ValueError, match="split evenly"):
plan_causal_row_balance((100, 21), cp_rank=0, cp_size=2, device="cpu")
def test_differentiable_inputs_are_rejected():
"""The exchange carries no autograd, so a gradient would be dropped without a word."""
plan = plan_causal_row_balance((64,), cp_rank=0, cp_size=2, device="cpu")
with pytest.raises(ValueError, match="no autograd"):
send_rows_to_scorers([torch.zeros(32, 4, requires_grad=True)], plan, cp_group=None)
def _exchange_worker(rank, world_size, port, seq_lens):
init_gloo(rank, world_size, port=port)
try:
rank_rows = sum(seq_lens) // world_size
plan = plan_causal_row_balance(seq_lens, cp_rank=rank, cp_size=world_size, device="cpu")
assert plan is not None
positions = torch.arange(rank * rank_rows, (rank + 1) * rank_rows)
# mixed dtypes and trailing shapes, each row tagged with its global position
tagged = positions.view(-1, 1, 1).expand(-1, 2, 3).to(torch.bfloat16).contiguous()
weights = positions.double().view(-1, 1) * 0.5
exchange = send_rows_to_scorers([positions, tagged, weights], plan, dist.group.WORLD)
received = exchange.wait()
scored = exchange.plan.scored_positions
assert torch.equal(received[0], scored)
assert torch.equal(received[1], scored.view(-1, 1, 1).expand(-1, 2, 3).to(torch.bfloat16))
assert torch.equal(received[2], scored.double().view(-1, 1) * 0.5)
# results with rows along dim 1, as the indexer's [batch, rows, topk] picks
picks = (scored * 10).view(1, -1, 1).expand(2, -1, 4).int()
back = exchange.return_to_owners(picks, dim=1)
assert torch.equal(back, (positions * 10).view(1, -1, 1).expand(2, -1, 4).int())
finally:
dist.destroy_process_group()
@pytest.mark.parametrize("world_size", [2, 4])
@pytest.mark.parametrize("seq_lens", [(512,), (320, 96, 64, 32)])
def test_exchange_round_trips_over_a_process_group(world_size, seq_lens):
run_multiprocess(partial(_exchange_worker, seq_lens=seq_lens), world_size=world_size)
def test_exchange_round_trips_when_a_rank_scores_nothing():
"""Documents shorter than 2cp leave chunks empty, so the gate can accept a plan with an idle rank."""
seq_lens = (1,) * 235 + (2,) * 76 + (5,) * 25
assert [plan.num_scored for plan in _plans(seq_lens, 4)] == [336, 0, 126, 50]
run_multiprocess(partial(_exchange_worker, seq_lens=seq_lens), world_size=4)
+21
View File
@@ -19,9 +19,11 @@ from miles_plugins.models.deepseek_v4.ops.thd_utils import (
CompressorInputCompact,
compact_gather_index,
compact_group_capacity,
compress_bounds_at_positions,
compressed_cu_seqlens,
compressed_rank_layout,
compressor_boundary_width,
get_compress_cu_seqlens_thd,
get_compress_topk_idxs_thd,
get_window_topk_idxs_thd,
to_rank_major_rows,
@@ -234,3 +236,22 @@ def test_an_empty_compressed_stream_yields_no_rows():
)
assert (rows == -1).all()
assert not valid.any()
@pytest.mark.parametrize("ratio", [4, 128])
@pytest.mark.parametrize("shape", sorted(SHAPES))
def test_compress_bounds_follow_the_position_not_the_row(ratio, shape):
"""The CP-balanced indexer scores rows on another rank, so a row's bounds may depend only on
its global position; they must match what that token sees when its sample runs alone."""
lens = SHAPES[shape]
total = sum(lens)
cu = _cu(lens)
cu_comp = compressed_cu_seqlens(cu, ratio)
ks, ke = get_compress_cu_seqlens_thd(cu, cu_comp, ratio=ratio, total_tokens=total)
for token in _probe_tokens(lens):
_, want = _reference_rows(lens, ratio, token)
assert set(range(int(ks[token]), int(ke[token]))) == want, f"{shape} ratio={ratio} token={token}"
positions = torch.randperm(total, generator=torch.Generator().manual_seed(0))
got_ks, got_ke = compress_bounds_at_positions(cu, cu_comp, positions, ratio=ratio)
assert torch.equal(got_ks, ks[positions]) and torch.equal(got_ke, ke[positions])
@@ -48,6 +48,7 @@ def _thd(cu, *, max_seqlen=0, compressed_group_ids=None):
"""ThdLayout for a single-rank packed stream."""
return ThdLayout(
cu_seqlens=cu,
seq_lens=tuple(torch.diff(cu).tolist()),
global_start=0,
max_seqlen=max_seqlen,
compressed_group_ids=compressed_group_ids,