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
Original file line number Diff line number Diff line change
Expand Up @@ -43,10 +43,17 @@ def _qwen_attention_mask_for_core_attention(
or packed_seq_params is not None
or not isinstance(attention_mask, Tensor)
or attention_mask.ndim != 2
or attention_mask.dtype != torch.bool
):
return attention_mask
return (~attention_mask).unsqueeze(1).unsqueeze(1)

if attention_mask.dtype == torch.bool:
valid_mask = attention_mask
elif attention_mask.dtype in (torch.uint8, torch.int8, torch.int16, torch.int32, torch.int64):
valid_mask = attention_mask != 0
else:
return attention_mask

return (~valid_mask).unsqueeze(1).unsqueeze(1)


class Qwen3VLSelfAttention(SelfAttention):
Expand Down
87 changes: 87 additions & 0 deletions tests/unit_tests/models/qwen_omni/test_qwen3_omni_bridge.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,13 +12,15 @@
# See the License for the specific language governing permissions and
# limitations under the License.

from types import SimpleNamespace
from unittest.mock import Mock, patch

import pytest
import torch
import torch.nn.functional as F

from megatron.bridge.models.conversion.mapping_registry import MegatronMappingRegistry
from megatron.bridge.models.conversion.param_mapping import GatedMLPMapping, QKVMapping
from megatron.bridge.models.hf_pretrained.causal_lm import PreTrainedCausalLM
from megatron.bridge.models.qwen_omni import Qwen3OmniBridge, Qwen3OmniModelProvider

Expand Down Expand Up @@ -201,6 +203,91 @@ def test_mapping_registry_resolves_parallel_parameter_names(self):
assert mapping is not None, megatron_param
assert mapping.hf_param == hf_param

def test_qkv_mapping_preserves_asymmetric_projection_slices(self):
registry = Qwen3OmniBridge().mapping_registry()
mapping = registry.megatron_to_hf_lookup(
"thinker.language_model.decoder.layers.0.self_attention.linear_qkv.weight"
)
assert isinstance(mapping, QKVMapping)

config = SimpleNamespace(num_attention_heads=4, num_query_groups=2, kv_channels=2, hidden_size=4)
module = SimpleNamespace(config=config)
q = torch.arange(8 * 4, dtype=torch.float32).reshape(8, 4) + 10
k = torch.arange(4 * 4, dtype=torch.float32).reshape(4, 4) + 100
v = torch.arange(4 * 4, dtype=torch.float32).reshape(4, 4) + 200

with patch.object(mapping._tp_mapping, "hf_to_megatron", side_effect=lambda weight, _module: weight):
fused = mapping.hf_to_megatron({"q": q, "k": k, "v": v}, module)

expected = torch.cat([q[:4], k[:2], v[:2], q[4:], k[2:], v[2:]], dim=0)
torch.testing.assert_close(fused, expected)

with (
patch.object(mapping, "broadcast_obj_from_pp_rank", side_effect=lambda value, _key: value),
patch.object(mapping._tp_mapping, "megatron_to_hf", return_value={mapping.megatron_param: fused}),
):
restored = mapping.megatron_to_hf(fused, module)

torch.testing.assert_close(restored[mapping.hf_param["q"]], q)
torch.testing.assert_close(restored[mapping.hf_param["k"]], k)
torch.testing.assert_close(restored[mapping.hf_param["v"]], v)

def test_expert_mapping_preserves_asymmetric_gate_up_slices(self):
registry = Qwen3OmniBridge().mapping_registry()
mapping = registry.megatron_to_hf_lookup(
"thinker.language_model.decoder.layers.0.mlp.experts.linear_fc1.weight7"
)
assert isinstance(mapping, GatedMLPMapping)

gate = torch.arange(12, dtype=torch.float32).reshape(4, 3) + 10
up = torch.arange(12, dtype=torch.float32).reshape(4, 3) + 100
fused = mapping.hf_to_megatron({"gate": gate, "up": up}, SimpleNamespace())
torch.testing.assert_close(fused, torch.cat([gate, up], dim=0))

with patch.object(mapping, "broadcast_from_pp_rank", side_effect=lambda value, cache_key=None: value):
restored = mapping.megatron_to_hf(fused, SimpleNamespace())

torch.testing.assert_close(restored[mapping.hf_param["gate"]], gate)
torch.testing.assert_close(restored[mapping.hf_param["up"]], up)

def test_expert_mapping_exports_gate_up_across_two_ep_ranks(self, monkeypatch):
registry = Qwen3OmniBridge().mapping_registry()
mapping = registry.megatron_to_hf_lookup(
"thinker.language_model.decoder.layers.0.mlp.experts.local_experts.1.linear_fc1.weight"
)
assert isinstance(mapping, GatedMLPMapping)

class _FakeGroup:
def size(self):
return 2

mapping.ep_group = _FakeGroup()
monkeypatch.setattr(
"megatron.bridge.models.conversion.param_mapping.get_pg_size",
lambda group: 1 if group is None else group.size(),
)
monkeypatch.setattr(mapping, "broadcast_from_pp_rank", lambda value, cache_key=None: value)
monkeypatch.setattr(mapping, "broadcast_obj_from_pp_rank", lambda value, cache_key=None: value)

def fake_all_gather(outputs, local_weight, group):
assert group is mapping.ep_group
outputs[0].copy_(local_weight)
outputs[1].copy_(local_weight + 1000)

monkeypatch.setattr(torch.distributed, "all_gather", fake_all_gather)

gate = torch.arange(6, dtype=torch.float32).reshape(2, 3) + 10
up = torch.arange(6, dtype=torch.float32).reshape(2, 3) + 100
module = SimpleNamespace(config=SimpleNamespace(num_moe_experts=4))
restored = mapping.megatron_to_hf(torch.cat([gate, up], dim=0), module)

gate_name = "thinker.model.layers.0.mlp.experts.{}.gate_proj.weight"
up_name = "thinker.model.layers.0.mlp.experts.{}.up_proj.weight"
torch.testing.assert_close(restored[gate_name.format(1)], gate)
torch.testing.assert_close(restored[gate_name.format(3)], gate + 1000)
torch.testing.assert_close(restored[up_name.format(1)], up)
torch.testing.assert_close(restored[up_name.format(3)], up + 1000)

def test_provider_bridge_warns_for_audio_output_stack(self, mock_hf_pretrained, caplog):
mock_hf_pretrained.config.enable_audio_output = True
caplog.set_level("WARNING")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -25,22 +25,31 @@ class _FakeTEDotProductAttention:
pass


def test_converts_qwen_valid_mask_for_te(monkeypatch):
@pytest.mark.parametrize("dtype", [torch.bool, torch.int32, torch.int64])
def test_converts_qwen_valid_mask_for_te(monkeypatch, dtype):
monkeypatch.setattr(attention_module, "TEDotProductAttention", _FakeTEDotProductAttention)
valid_mask = torch.tensor([[True, True, False], [False, True, True]])
valid_mask = torch.tensor([[1, 1, 0], [0, 1, 1]], dtype=dtype)

actual = _qwen_attention_mask_for_core_attention(_FakeTEDotProductAttention(), valid_mask, packed_seq_params=None)

expected = torch.tensor([[[[False, False, True]]], [[[True, False, False]]]])
assert actual.dtype == torch.bool
torch.testing.assert_close(actual, expected)

attention_scores = torch.tensor(
[[[[1.0, 2.0, 3.0]]], [[[4.0, 5.0, 6.0]]]],
)
attention_probs = torch.softmax(attention_scores.masked_fill(actual, float("-inf")), dim=-1)
assert torch.isfinite(attention_probs).all()
assert torch.equal(attention_probs.masked_select(actual), torch.zeros(2))


def test_converts_all_valid_qwen_mask_for_te_without_host_sync(monkeypatch):
monkeypatch.setattr(attention_module, "TEDotProductAttention", _FakeTEDotProductAttention)
monkeypatch.setattr(
torch.Tensor, "item", lambda _tensor: pytest.fail("attention mask must not synchronize to host")
)
valid_mask = torch.ones((2, 4), dtype=torch.bool)
valid_mask = torch.ones((2, 4), dtype=torch.long)

actual = _qwen_attention_mask_for_core_attention(_FakeTEDotProductAttention(), valid_mask, packed_seq_params=None)
monkeypatch.undo()
Expand All @@ -51,7 +60,7 @@ def test_converts_all_valid_qwen_mask_for_te_without_host_sync(monkeypatch):
@pytest.mark.parametrize(
("mask", "packed_seq_params"),
[
(torch.ones((2, 4), dtype=torch.long), None),
(torch.zeros((2, 4), dtype=torch.float32), None),
(torch.ones((2, 1, 4, 4), dtype=torch.bool), None),
(torch.ones((2, 4), dtype=torch.bool), object()),
],
Expand Down
Loading