mirror of
https://github.com/radixark/miles.git
synced 2026-10-01 23:06:14 +08:00
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:
co-authored by
Claude Opus 5.5
parent
9e4260de04
commit
3439ec7513
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user