diff --git a/src/mcore_bridge/model/gpts/deepseek_v4.py b/src/mcore_bridge/model/gpts/deepseek_v4.py index eb4c122..386e70e 100644 --- a/src/mcore_bridge/model/gpts/deepseek_v4.py +++ b/src/mcore_bridge/model/gpts/deepseek_v4.py @@ -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, diff --git a/tests/test_dsv4_mla_rope.py b/tests/test_dsv4_mla_rope.py new file mode 100644 index 0000000..8958bcd --- /dev/null +++ b/tests/test_dsv4_mla_rope.py @@ -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')