diff --git a/src/mcore_bridge/model/gpts/nemotron_h.py b/src/mcore_bridge/model/gpts/nemotron_h.py index 744602e..176871a 100644 --- a/src/mcore_bridge/model/gpts/nemotron_h.py +++ b/src/mcore_bridge/model/gpts/nemotron_h.py @@ -9,8 +9,6 @@ import torch from megatron.core.extensions.transformer_engine import TEColumnParallelLinear, TENorm, TERowParallelLinear from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add -from megatron.core.models.hybrid.hybrid_block import HybridStack, HybridStackSubmodules -from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec from megatron.core.ssm.mamba_layer import MambaLayer, MambaLayerSubmodules from megatron.core.ssm.mamba_mixer import MambaMixer, MambaMixerSubmodules from megatron.core.transformer.spec_utils import ModuleSpec @@ -21,9 +19,21 @@ from mcore_bridge.utils import get_logger from ..constant import ModelType -from ..hybrid_model import HybridModel from ..register import ModelLoader, ModelMeta, register_model +try: + from megatron.core.models.hybrid.hybrid_block import HybridStack, HybridStackSubmodules + from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec + + from ..hybrid_model import HybridModel + + _HYBRID_MODEL_AVAILABLE = True +except ImportError as error: + if not (error.name or '').startswith('megatron.core.models.hybrid'): + raise + HybridModel = HybridStack = HybridStackSubmodules = hybrid_stack_spec = None + _HYBRID_MODEL_AVAILABLE = False + logger = get_logger() @@ -410,9 +420,11 @@ def build_model(self, pre_process=True, post_process=True, vp_stage: Optional[in return model -register_model(ModelMeta( - ModelType.nemotron_h, - ['nemotron_h'], - bridge_cls=NemotronHBridge, - loader=NemotronHLoader, -)) +if _HYBRID_MODEL_AVAILABLE: + register_model( + ModelMeta( + ModelType.nemotron_h, + ['nemotron_h'], + bridge_cls=NemotronHBridge, + loader=NemotronHLoader, + )) diff --git a/src/mcore_bridge/model/gpts/qwen3_next.py b/src/mcore_bridge/model/gpts/qwen3_next.py index 0d4e290..f91f037 100644 --- a/src/mcore_bridge/model/gpts/qwen3_next.py +++ b/src/mcore_bridge/model/gpts/qwen3_next.py @@ -471,10 +471,10 @@ def forward(self, hidden_states: torch.Tensor, **kwargs): # Note: for packed inputs, we do not perform padding_free unpadding. # Doing so would allow different sequences to see each other; for efficiency we keep this implementation. if thd_format: + max_seqlen_q = int(packed_seq_params.max_seqlen_q) new_hidden_states = hidden_states.new_zeros( - (packed_seq_params.num_samples, packed_seq_params.max_seqlen_q.item(), hidden_states.shape[-1])) - attention_mask = hidden_states.new_zeros( - (packed_seq_params.num_samples, packed_seq_params.max_seqlen_q.item()), dtype=torch.bool) + (packed_seq_params.num_samples, max_seqlen_q, hidden_states.shape[-1])) + attention_mask = hidden_states.new_zeros((packed_seq_params.num_samples, max_seqlen_q), dtype=torch.bool) cu_seqlens_q = packed_seq_params.cu_seqlens_q for i in range(packed_seq_params.num_samples): start, end = cu_seqlens_q[i], cu_seqlens_q[i + 1] diff --git a/src/mcore_bridge/model/mm_gpts/qwen3_5.py b/src/mcore_bridge/model/mm_gpts/qwen3_5.py index 8ba5378..b673bf7 100644 --- a/src/mcore_bridge/model/mm_gpts/qwen3_5.py +++ b/src/mcore_bridge/model/mm_gpts/qwen3_5.py @@ -40,10 +40,10 @@ def forward(self, hidden_states: torch.Tensor, **kwargs): # Note: for packed inputs, we do not perform padding_free unpadding. # Doing so would allow different sequences to see each other; for efficiency we keep this implementation. if thd_format: + max_seqlen_q = int(packed_seq_params.max_seqlen_q) new_hidden_states = hidden_states.new_zeros( - (packed_seq_params.num_samples, packed_seq_params.max_seqlen_q.item(), hidden_states.shape[-1])) - attention_mask = hidden_states.new_zeros( - (packed_seq_params.num_samples, packed_seq_params.max_seqlen_q.item()), dtype=torch.bool) + (packed_seq_params.num_samples, max_seqlen_q, hidden_states.shape[-1])) + attention_mask = hidden_states.new_zeros((packed_seq_params.num_samples, max_seqlen_q), dtype=torch.bool) cu_seqlens_q = packed_seq_params.cu_seqlens_q for i in range(packed_seq_params.num_samples): start, end = cu_seqlens_q[i], cu_seqlens_q[i + 1]