Skip to content

Commit 8a41c4d

Browse files
Merge pull request #4402 from AI-Hypercomputer:hengtaoguo-test
PiperOrigin-RevId: 945359203
2 parents 409f43b + d2eaacf commit 8a41c4d

2 files changed

Lines changed: 14 additions & 21 deletions

File tree

.github/workflows/run_tests_coordinator.yml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -209,7 +209,7 @@ jobs:
209209
"gpu-integration": "--ignore=tests/post_training",
210210
"cpu-unit": "--ignore=tests/post_training",
211211
"cpu-post-training-unit": "",
212-
"cpu-torch-reference": "-o addopts= -rf --import-mode=importlib --strict-markers tests/unit/gemma4_layers_test.py tests/unit/gemma4_small_layers_test.py tests/unit/qwen3_next_vs_reference_test.py tests/unit/qwen3_5_layers_test.py"
212+
"cpu-torch-reference": "-o addopts= -rf --import-mode=importlib --strict-markers tests/unit/gemma4_layers_test.py tests/unit/gemma4_small_layers_test.py tests/unit/qwen3_next_vs_reference_test.py tests/unit/qwen3_5_layers_test.py tests/unit/qwen3_omni_layers_test.py"
213213
}')[inputs.flavor] }}
214214
${{ inputs.additional_pytest_args }}
215215

tests/unit/qwen3_omni_layers_test.py

Lines changed: 13 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
import jax
2626
import jax.numpy as jnp
2727
from jax.sharding import Mesh
28+
import pytest
2829
from maxtext.configs import pyconfig
2930
from maxtext.common import common_types
3031
from maxtext.utils.globals import MAXTEXT_REPO_ROOT
@@ -62,7 +63,6 @@
6263
copy_patch_embed_weights,
6364
copy_patch_merger_weights,
6465
copy_vision_encoder_weights,
65-
create_block_diagonal_attention_mask,
6666
create_random_jax_torch,
6767
)
6868
import numpy as np
@@ -777,6 +777,7 @@ def test_hidden_states_unchanged_without_visual_tokens(self):
777777
np.testing.assert_allclose(np.array(result), hidden_np, rtol=1e-6, atol=1e-6)
778778

779779

780+
@pytest.mark.skip(reason="Requires decord, which may not be installed in remote CI runners.")
780781
class TestQwen3OmniPreprocessing(unittest.TestCase):
781782
"""Test MaxText Qwen3 Omni preprocessor against HuggingFace reference."""
782783

@@ -1003,17 +1004,12 @@ def _test_encoder_layer_with_batch_size(self, batch_size):
10031004

10041005
jax_input, torch_input_3d = create_random_jax_torch(batch_size, seq_len, hidden_size)
10051006

1006-
# PyTorch forward pass - expects 2D input (total_seq_len, hidden_dim) with cu_seqlens
1007-
torch_input_2d = torch_input_3d.reshape(-1, hidden_size)
1008-
1009-
# Create cu_seqlens for PyTorch (cumulative sequence lengths for each batch)
1010-
# For batch_size=2, seq_len=12: [0, 12, 24] indicates two sequences of length 12 each
1011-
cu_seqlens = torch.tensor([i * seq_len for i in range(batch_size + 1)], dtype=torch.int32)
1012-
1013-
attention_mask = create_block_diagonal_attention_mask(cu_seqlens, torch_input_2d.dtype)
1014-
1015-
torch_output_1d = torch_layer(torch_input_2d, cu_seqlens=cu_seqlens, attention_mask=attention_mask)[0]
1016-
torch_output = torch_output_1d.reshape(batch_size, seq_len, hidden_size)
1007+
# PyTorch audio layers take 2D packed input. Run each batch item separately
1008+
# to match MaxText's batched attention without relying on HF's old mask API.
1009+
cu_seqlens = torch.tensor([0, seq_len], dtype=torch.int32)
1010+
torch_output = torch.stack(
1011+
[torch_layer(torch_input_3d[i], cu_seqlens=cu_seqlens)[0] for i in range(batch_size)], dim=0
1012+
)
10171013

10181014
jax_output = maxtext_layer(jax_input, deterministic=True)
10191015

@@ -1144,18 +1140,15 @@ def test_audio_encoder_matches_torch(self):
11441140
torch_after_pos = torch_conv_out + torch_pos_emb
11451141

11461142
# Run through encoder layers + layernorm (but not projector)
1147-
# Process all chunks together
1143+
# Process chunks separately, matching MaxText's (batch * chunks, seq, hidden) attention shape.
11481144
seq_len_per_chunk = torch_after_pos.shape[1]
1149-
cu_seqlens = torch.tensor([i * seq_len_per_chunk for i in range(num_chunks + 1)], dtype=torch.int32)
1150-
attention_mask = create_block_diagonal_attention_mask(cu_seqlens, torch_after_pos.dtype)
1151-
1152-
# Flatten: (num_chunks, seq_len_per_chunk, hidden) -> (num_chunks*seq_len_per_chunk, hidden)
1153-
hidden_state = torch_after_pos.reshape(-1, torch_after_pos.shape[-1])
1145+
cu_seqlens = torch.tensor([0, seq_len_per_chunk], dtype=torch.int32)
1146+
hidden_state = torch_after_pos
11541147
for layer in torch_model.layers:
1155-
hidden_state = layer(hidden_state, cu_seqlens=cu_seqlens, attention_mask=attention_mask)[0]
1148+
hidden_state = torch.stack([layer(chunk, cu_seqlens=cu_seqlens)[0] for chunk in hidden_state], dim=0)
11561149
hidden_state = torch_model.ln_post(hidden_state)
11571150

1158-
# Reshape back: (num_chunks*seq_len_per_chunk, hidden) -> (batch=1, num_chunks*seq_len_per_chunk, hidden)
1151+
# Reshape back: (num_chunks, seq_len_per_chunk, hidden) -> (batch=1, num_chunks*seq_len_per_chunk, hidden)
11591152
torch_output = hidden_state.reshape(1, num_chunks * seq_len_per_chunk, -1)
11601153

11611154
# MaxText forward

0 commit comments

Comments
 (0)