Skip to content
Merged
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
6 changes: 5 additions & 1 deletion src/mcore_bridge/model/gpts/deepseek_v4.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,11 +80,15 @@ def _apply_mla_rope(t, freqs, *, config, cu_seqlens, cp_group, inverse=False):
f'freqs.shape[0]={freqs.shape[0]} vs tokens={t.shape[0]}. `GPTModel` must pre-index the '
'rotary table by `position_ids` (requires `apply_rope_fusion=False`), and under CP the '
'`position_ids` must be split with the same partition mode as the hidden states.')
# Do not forward cu_seqlens. The packed helper re-derives positions from it and, when
# cp_size > 1, assumes Megatron's zigzag split. DSv4 uses a contiguous split and already
# stores the absolute position in row i of `freqs`, so the elementwise bshd path is the
# one that matches this tensor. The argument stays so existing callers keep working.
return apply_rotary_pos_emb(
t,
freqs,
config=config,
cu_seqlens=cu_seqlens,
cu_seqlens=None,
cp_group=cp_group,
mla_rotary_interleaved=True,
mla_output_remove_interleaving=True,
Expand Down
46 changes: 46 additions & 0 deletions tests/test_dsv4_mla_rope.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
"""DSv4 RoPE must not hand packed cu_seqlens to Megatron's zigzag helper."""
import ast
import torch
from pathlib import Path


def _load_apply():
path = Path(__file__).resolve().parents[1] / 'src/mcore_bridge/model/gpts/deepseek_v4.py'
tree = ast.parse(path.read_text())
fn = next(node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == '_apply_mla_rope')
module = ast.Module(body=[fn], type_ignores=[])
ast.fix_missing_locations(module)
captured = {}

def apply_rotary_pos_emb(t, freqs, **kwargs):
captured['kwargs'] = kwargs
captured['t'] = t
captured['freqs'] = freqs
return t

namespace = {'apply_rotary_pos_emb': apply_rotary_pos_emb}
exec(compile(module, str(path), 'exec'), namespace)
return namespace['_apply_mla_rope'], captured


def test_packed_cu_seqlens_are_not_forwarded():
apply, captured = _load_apply()
tokens = torch.zeros(4, 2, 8)
freqs = torch.zeros(4, 1, 1, 4)
cu = torch.tensor([0, 4])
apply(tokens, freqs, config=object(), cu_seqlens=cu, cp_group=object())
assert captured['kwargs']['cu_seqlens'] is None
assert captured['kwargs']['mla_rotary_interleaved'] is True
assert captured['t'] is tokens
assert captured['freqs'] is freqs


def test_misaligned_frequencies_still_fail():
apply, _ = _load_apply()
try:
apply(torch.zeros(4, 1, 4), torch.zeros(3, 1, 1, 2), config=object(), cu_seqlens=None, cp_group=None)
except AssertionError as error:
assert 'row-aligned' in str(error)
else:
raise AssertionError('expected a row-alignment failure')
Loading