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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions miles/backends/experimental/fsdp_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -386,6 +386,7 @@ def _compute_log_prob(
response_lengths=batch["response_lengths"],
with_entropy=(store_prefix == ""),
max_seq_lens=batch.get("max_seq_lens", None),
padded_total_lengths=batch.get("padded_total_lengths", None),
)

batch_result = {
Expand Down
1 change: 1 addition & 0 deletions miles/backends/megatron_utils/parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,7 @@ def get_packed_seq_params(batch: dict[str, torch.Tensor], args: Namespace) -> Pa
max_seqlen_kv=batch["max_seqlen"],
qkv_format="thd",
)
packed_seq_params.miles_allgather_cp = bool(getattr(args, "allgather_cp", False))
batch["packed_seq_params"] = packed_seq_params
return packed_seq_params
else:
Expand Down
5 changes: 5 additions & 0 deletions miles/backends/sglang_utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,6 +136,11 @@ def new_add_argument_wrapper(*name_or_flags, **kwargs):

def validate_args(args):
args.sglang_tp_size = args.rollout_num_gpus_per_engine
args.sglang_dp_size = args.sglang_data_parallel_size
args.sglang_pp_size = args.sglang_pipeline_parallel_size
args.sglang_ep_size = args.sglang_expert_parallel_size
if hasattr(args, "sglang_attention_context_parallel_size"):
args.sglang_attn_cp_size = args.sglang_attention_context_parallel_size

if args.true_on_policy_mode:
args.sglang_enable_deterministic_inference = True
Expand Down
43 changes: 33 additions & 10 deletions miles/backends/training_utils/cp_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ def get_logits_and_tokens_offset_with_cp(
response_length: int,
qkv_format: str = "thd",
max_seq_len: int | None = None,
padded_total_length: int | None = None,
):
"""
All offsets start from the begining of the prompt.
Expand All @@ -31,8 +32,10 @@ def get_logits_and_tokens_offset_with_cp(
assert cp_size > 1

prompt_length = total_length - response_length
effective_total_length = padded_total_length if padded_total_length is not None else total_length

if qkv_format == "thd":
chunk_size = (total_length + 2 * cp_size - 1) // (2 * cp_size)
chunk_size = (effective_total_length + 2 * cp_size - 1) // (2 * cp_size)
else:
assert max_seq_len is not None, "max_seq_len must be provided for qkv_format=bshd"
chunk_size = (max_seq_len + 2 * cp_size - 1) // (2 * cp_size)
Expand Down Expand Up @@ -67,11 +70,12 @@ def _slice_loss_mask_for_local_cp(
loss_mask: torch.Tensor,
qkv_format: str,
max_seq_len: int | None,
padded_total_length: int | None = None,
) -> torch.Tensor:
"""Slice a per-sample response loss mask into this CP rank's local zigzag layout."""
prompt_length = total_length - response_length
_, _, _, tokens_offset = get_logits_and_tokens_offset_with_cp(
total_length, response_length, qkv_format, max_seq_len
total_length, response_length, qkv_format, max_seq_len, padded_total_length
)
mask_0 = loss_mask[tokens_offset[0][0] - prompt_length : tokens_offset[0][1] - prompt_length]
mask_1 = loss_mask[tokens_offset[1][0] - prompt_length : tokens_offset[1][1] - prompt_length]
Expand All @@ -84,9 +88,12 @@ def slice_loss_masks_for_local_cp(
response_lengths: list[int],
qkv_format: str = "thd",
max_seq_lens: list[int] | None = None,
padded_total_lengths: list[int] | None = None,
) -> list[torch.Tensor]:
"""Backward-compatible wrapper for local CP response mask slicing."""
return get_local_response_loss_masks(total_lengths, response_lengths, loss_masks, qkv_format, max_seq_lens)
return get_local_response_loss_masks(
total_lengths, response_lengths, loss_masks, qkv_format, max_seq_lens, padded_total_lengths
)


def get_sum_of_sample_mean(
Expand All @@ -96,6 +103,7 @@ def get_sum_of_sample_mean(
calculate_per_token_loss: bool = False,
qkv_format: str = "thd",
max_seq_lens: list[int] | None = None,
padded_total_lengths: list[int] | None = None,
) -> Callable[[torch.Tensor], torch.Tensor]:
"""
Calculate correct sample mean for CP
Expand Down Expand Up @@ -127,8 +135,11 @@ def sum_of_token(x: torch.Tensor) -> torch.Tensor:
zip(total_lengths, response_lengths, loss_masks, strict=True)
):
max_seq_len = max_seq_lens[i] if max_seq_lens is not None else None
padded_total_length = padded_total_lengths[i] if padded_total_lengths is not None else None
chunked_loss_masks.append(
_slice_loss_mask_for_local_cp(total_length, response_length, loss_mask, qkv_format, max_seq_len)
_slice_loss_mask_for_local_cp(
total_length, response_length, loss_mask, qkv_format, max_seq_len, padded_total_length
)
)
cp_chunk_lengths.append(chunked_loss_masks[i].size(0))

Expand Down Expand Up @@ -161,6 +172,7 @@ def get_local_response_loss_masks(
loss_masks: list[torch.Tensor],
qkv_format: str = "thd",
max_seq_lens: list[int] | None = None,
padded_total_lengths: list[int] | None = None,
) -> list[torch.Tensor]:
"""Return response loss masks aligned with this rank's local log-probs."""
parallel_state = get_parallel_state()
Expand All @@ -172,8 +184,11 @@ def get_local_response_loss_masks(
zip(total_lengths, response_lengths, loss_masks, strict=True)
):
max_seq_len = max_seq_lens[i] if max_seq_lens is not None else None
padded_total_length = padded_total_lengths[i] if padded_total_lengths is not None else None
local_masks.append(
_slice_loss_mask_for_local_cp(total_length, response_length, loss_mask, qkv_format, max_seq_len)
_slice_loss_mask_for_local_cp(
total_length, response_length, loss_mask, qkv_format, max_seq_len, padded_total_length
)
)

return local_masks
Expand All @@ -185,6 +200,7 @@ def all_gather_with_cp(
response_length: int,
qkv_format: str = "thd",
max_seq_len: int | None = None,
padded_total_length: int | None = None,
) -> torch.Tensor:
"""
Gather tensors across all ranks in the context parallel group.
Expand All @@ -198,7 +214,7 @@ def all_gather_with_cp(
return tensor

_, _, logits_offset, _ = get_logits_and_tokens_offset_with_cp(
total_length, response_length, qkv_format, max_seq_len
total_length, response_length, qkv_format, max_seq_len, padded_total_length
)

prompt_length = total_length - response_length
Expand Down Expand Up @@ -322,6 +338,7 @@ def allgather_cp_redistribute(
total_lengths: list[int],
response_lengths: list[int],
max_seq_lens: list[int] | None = None,
padded_total_lengths: list[int] | None = None,
) -> None:
"""Redistribute response tensors from allgather-CP layout to zigzag ring-attn layout.

Expand Down Expand Up @@ -353,7 +370,9 @@ def allgather_cp_redistribute(
# Reconstruct full response tensors with each rank's contiguous contribution
full_resps = []
seq_start = 0
for value, total_length, response_length in zip(values, total_lengths, response_lengths, strict=False):
for idx, (value, total_length, response_length) in enumerate(
zip(values, total_lengths, response_lengths, strict=False)
):
prompt_length = total_length - response_length
logit_global_start = seq_start + prompt_length - 1
logit_global_end = seq_start + total_length - 1
Expand All @@ -376,7 +395,7 @@ def allgather_cp_redistribute(

assert full_resp.size(0) == response_length, f"Expected {response_length}, got {full_resp.size(0)}"
full_resps.append(full_resp)
seq_start += total_length
seq_start += padded_total_lengths[idx] if padded_total_lengths is not None else total_length

# Single differentiable all-reduce to gather full response from all CP ranks
all_cat = torch.cat(full_resps, dim=0)
Expand All @@ -388,8 +407,11 @@ def allgather_cp_redistribute(
zip(all_cat.split(response_lengths, dim=0), total_lengths, response_lengths, strict=False)
):
max_seq_len = max_seq_lens[idx] if max_seq_lens is not None else None
padded_total_length = padded_total_lengths[idx] if padded_total_lengths is not None else None
new_values.append(
slice_log_prob_with_cp(full_resp, total_length, response_length, args.qkv_format, max_seq_len)
slice_log_prob_with_cp(
full_resp, total_length, response_length, args.qkv_format, max_seq_len, padded_total_length
)
)

res[key] = new_values
Expand All @@ -401,6 +423,7 @@ def slice_log_prob_with_cp(
response_length: int,
qkv_format: str = "thd",
max_token_len: int | None = None,
padded_total_length: int | None = None,
) -> list[float] | torch.Tensor:
assert len(log_prob) == response_length

Expand All @@ -412,7 +435,7 @@ def slice_log_prob_with_cp(

prompt_length = total_length - response_length
_, _, logits_offset, _ = get_logits_and_tokens_offset_with_cp(
total_length, response_length, qkv_format, max_token_len
total_length, response_length, qkv_format, max_token_len, padded_total_length
)

chunk_1 = log_prob[logits_offset[0][0] - (prompt_length - 1) : logits_offset[0][1] - (prompt_length - 1)]
Expand Down
97 changes: 93 additions & 4 deletions miles/backends/training_utils/data.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import logging
import math
from argparse import Namespace
from collections.abc import Sequence

Expand All @@ -21,6 +22,41 @@
logger = logging.getLogger(__name__)


def _get_thd_sample_pad_multiple(args: Namespace) -> int | None:
"""Return per-sample padding multiple required by packed THD kernels."""
model_name = (getattr(args, "model_name", "") or "").lower().replace("-", "").replace("_", "")
is_deepseek_v4 = "deepseekv4" in model_name
if not is_deepseek_v4:
return None

compress_ratios = getattr(args, "compress_ratios", None)
active_ratios = [int(ratio) for ratio in (compress_ratios or []) if int(ratio) > 1]
if not active_ratios:
return None
return math.lcm(*active_ratios)


def get_thd_padded_total_lengths(args: Namespace, total_lengths: Sequence[int]) -> list[int] | None:
"""Return model-input lengths after DSV4 THD per-sample alignment."""
if getattr(args, "qkv_format", None) != "thd":
return None
multiple = _get_thd_sample_pad_multiple(args)
if multiple is None:
return None
return [((int(length) + multiple - 1) // multiple) * multiple for length in total_lengths]


def _get_thd_allgather_pad_multiple(cp_size: int, pad_size: int, sample_pad_multiple: int | None) -> int:
"""Return a global THD multiple that keeps every contiguous CP shard kernel-aligned."""
global_pad_multiple = cp_size * pad_size
if sample_pad_multiple is not None:
# DSV4 compressors require each CP-local shard to contain a whole
# pair of compression groups. The global stream must therefore be
# divisible by 2 * cp_size * every active compression ratio.
global_pad_multiple = math.lcm(global_pad_multiple, 2 * cp_size * sample_pad_multiple)
return global_pad_multiple


def _rollout_logprob_dtype(args: Namespace) -> torch.dtype:
if getattr(args, "true_on_policy_mode", False):
if getattr(args, "bf16", False):
Expand Down Expand Up @@ -84,6 +120,8 @@ def get_rollout_data(

rollout_data["max_seq_lens"] = [max_seq_len] * len(rollout_data["tokens"])

padded_total_lengths = get_thd_padded_total_lengths(args, rollout_data["total_lengths"])

# Full-response SGLang OPD fields share rollout CP slicing but retain float32 precision.
for key in ("rollout_log_probs", "teacher_log_probs", "opd_reverse_kl"):
if key in rollout_data:
Expand All @@ -96,6 +134,7 @@ def get_rollout_data(
response_length,
args.qkv_format,
rollout_data["max_seq_lens"][i] if args.qkv_format == "bshd" else None,
padded_total_lengths[i] if padded_total_lengths is not None else None,
),
device=torch.cuda.current_device(),
dtype=dtype,
Expand Down Expand Up @@ -164,6 +203,7 @@ def get_batch(
# use 0 as the pad token id should be fine?
pad_token_id = 0
pad_size = parallel_state.tp.size * pad_multiplier
padded_total_lengths: list[int] | None = None

# for cp, we need all tokens to calculate logprob
batch["unconcat_tokens"] = tokens
Expand All @@ -189,20 +229,36 @@ def get_batch(

elif qkv_format == "thd":
cp_rank = parallel_state.cp.rank
sample_pad_multiple = _get_thd_sample_pad_multiple(data_iterator.args)
padded_total_lengths = [] if sample_pad_multiple is not None else None
if sample_pad_multiple is not None and cp_size > 1 and not allgather_cp:
raise NotImplementedError(
"DeepSeek-V4 THD packing with CP>1 requires --allgather-cp; "
"zigzag CP cannot preserve packed sample boundaries yet."
)

if allgather_cp:
assert batch.get("adapter_slots") is None, "allgather CP is currently not supported with multi-LoRA: "
# DSA mode: concatenate all sequences first, then slice once with CP.
# We also pad the *global* concatenated stream to make per-rank batches equal.
cu_seqlens_list: list[int] = [0]
padded_tokens: list[torch.Tensor] = []
for t in tokens:
if sample_pad_multiple is not None:
sample_pad = (sample_pad_multiple - t.size(0) % sample_pad_multiple) % sample_pad_multiple
if sample_pad:
t = F.pad(t, (0, sample_pad), value=pad_token_id)
assert padded_total_lengths is not None
padded_total_lengths.append(t.size(0))
padded_tokens.append(t)
cu_seqlens_list.append(cu_seqlens_list[-1] + t.size(0))
tokens = padded_tokens

tokens = torch.cat(tokens, dim=0)

# Pad global stream so (1) divisible by cp_size (equal batches),
# (2) divisible by pad_size (reduce fragmentation).
global_pad_size = cp_size * pad_size
global_pad_size = _get_thd_allgather_pad_multiple(cp_size, pad_size, sample_pad_multiple)
pad = (global_pad_size - tokens.size(0) % global_pad_size) % global_pad_size
if pad != 0:
tokens = F.pad(tokens, (0, pad), value=pad_token_id)
Expand All @@ -211,7 +267,16 @@ def get_batch(
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:
tokens = [slice_with_cp(t, pad_token_id, qkv_format) for t in tokens]
padded_tokens: list[torch.Tensor] = []
for t in tokens:
if sample_pad_multiple is not None:
sample_pad = (sample_pad_multiple - t.size(0) % sample_pad_multiple) % sample_pad_multiple
if sample_pad:
t = F.pad(t, (0, sample_pad), value=pad_token_id)
assert padded_total_lengths is not None
padded_total_lengths.append(t.size(0))
padded_tokens.append(t)
tokens = [slice_with_cp(t, pad_token_id, qkv_format) for t in padded_tokens]
sample_token_lengths = [t.size(0) for t in tokens]

cu_seqlens = [0]
Expand All @@ -221,7 +286,8 @@ def get_batch(
tokens = torch.cat(tokens)

# Always pad to reduce memory fragmentation and maybe make the computation faster
pad = (pad_size - tokens.size(0) % pad_size) % pad_size
final_pad_size = math.lcm(pad_size, sample_pad_multiple) if sample_pad_multiple is not None else pad_size
pad = (final_pad_size - tokens.size(0) % final_pad_size) % final_pad_size
if pad != 0:
tokens = F.pad(tokens, (0, pad), value=pad_token_id)
cu_seqlens.append(cu_seqlens[-1] + pad)
Expand All @@ -235,6 +301,8 @@ def get_batch(

batch["cu_seqlens"] = cu_seqlens
batch["max_seqlen"] = max_seqlen
if padded_total_lengths is not None:
batch["padded_total_lengths"] = padded_total_lengths
else:
raise ValueError(f"Unsupported qkv_format: {qkv_format}")

Expand All @@ -261,6 +329,11 @@ def _compute_transform_like_token_ids(ids_list: list):
ids = [slice_with_cp(p, 0, qkv_format, max_seqlen) for p in ids_list]
ids = torch.stack(ids)
elif qkv_format == "thd":
if padded_total_lengths is not None:
ids_list = [
F.pad(p, (0, padded_total_length - p.size(0)), value=0)
for p, padded_total_length in zip(ids_list, padded_total_lengths, strict=True)
]
ids = [slice_with_cp(p, 0, qkv_format) for p in ids_list]
ids = torch.cat(ids)
if pad != 0:
Expand Down Expand Up @@ -293,6 +366,11 @@ def _compute_transform_like_token_ids(ids_list: list):
prompt_length = total_length - response_length
# Align mask to token stream positions (prompt_length-1 left pad, 1 right pad)
loss_mask = F.pad(loss_mask, (prompt_length - 1, 1), value=0)
if padded_total_lengths is not None:
padded_total_length = padded_total_lengths[len(loss_masks)]
sample_pad = padded_total_length - total_length
if sample_pad:
loss_mask = F.pad(loss_mask, (0, sample_pad), value=0)
if allgather_cp:
loss_masks.append(loss_mask)
continue
Expand Down Expand Up @@ -352,6 +430,8 @@ def __init__(
rollout_data: RolloutBatch,
micro_batch_size: int | None = None,
micro_batch_indices: list[list[int]] | None = None,
*,
args: Namespace | None = None,
) -> None:
"""Initialize an iterator over `rollout_data`.

Expand All @@ -360,7 +440,9 @@ def __init__(
micro_batch_size: Fixed contiguous slice size when not using dynamic scheduling.
micro_batch_indices: Explicit indices per micro-batch when using dynamic balancing.
Must be mutually exclusive with `micro_batch_size`.
args: Optional runtime args used for model-specific batch preparation.
"""
self.args = args or Namespace(model_name="", compress_ratios=None)
self.rollout_data = rollout_data
self.micro_batch_size = micro_batch_size
self.micro_batch_indices = micro_batch_indices
Expand Down Expand Up @@ -448,7 +530,14 @@ def get_data_iterator(
def _generate_data_iterator(rollout_data, micro_batch_size, micro_batch_indices=None):
data_iterator = []
for _ in range(vpp_size):
data_iterator.append(DataIterator(rollout_data, micro_batch_size, micro_batch_indices))
data_iterator.append(
DataIterator(
rollout_data,
micro_batch_size,
micro_batch_indices,
args=args,
)
)
return data_iterator

if not args.use_dynamic_batch_size:
Expand Down
Loading