diff --git a/megatron/core/models/multimodal/llava_model.py b/megatron/core/models/multimodal/llava_model.py index af0bcf6e9fd..d3e5d5e26f8 100644 --- a/megatron/core/models/multimodal/llava_model.py +++ b/megatron/core/models/multimodal/llava_model.py @@ -17,8 +17,10 @@ from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer import MegatronModule +from megatron.core.transformer.attention import SelfAttentionSubmodules from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.transformer.transformer_layer import TransformerLayerSubmodules from megatron.core.utils import deprecate_inference_params, log_single_rank try: @@ -158,9 +160,18 @@ def __init__( self.context_parallel_lm = language_transformer_config.context_parallel_size if self.sequence_parallel_lm or self.context_parallel_lm > 1: if not language_model_type.startswith('nemotron5-hybrid'): - attn_module = language_transformer_layer_spec.submodules.self_attention + assert isinstance( + language_transformer_layer_spec.submodules, TransformerLayerSubmodules + ) + assert isinstance( + language_transformer_layer_spec.submodules.self_attention.submodules, + SelfAttentionSubmodules, + ) + attn_submodules = ( + language_transformer_layer_spec.submodules.self_attention.submodules + ) assert ( - attn_module.submodules.core_attention == TEDotProductAttention and HAVE_TE + attn_submodules.core_attention == TEDotProductAttention and HAVE_TE ), "Sequence/Context Parallelism is supported only with TE DotProductAttention." if self.context_parallel_lm > 1: self.cp_group = self.pg_collection.cp