From e6ab124ab02d670e21cd84f57a58adcc6ed41cc6 Mon Sep 17 00:00:00 2001 From: dimapihtar Date: Wed, 8 Apr 2026 06:27:16 -0700 Subject: [PATCH 1/2] remove t5 legacy Signed-off-by: dimapihtar --- megatron/legacy/model/t5_model.py | 195 ------------------------------ pretrain_t5.py | 90 ++++++-------- 2 files changed, 40 insertions(+), 245 deletions(-) delete mode 100644 megatron/legacy/model/t5_model.py diff --git a/megatron/legacy/model/t5_model.py b/megatron/legacy/model/t5_model.py deleted file mode 100644 index 1662188334d..00000000000 --- a/megatron/legacy/model/t5_model.py +++ /dev/null @@ -1,195 +0,0 @@ -# Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved. - -"""T5 model.""" - -import torch - -from megatron.training import get_args -from megatron.core import tensor_parallel -from megatron.legacy.model.enums import AttnMaskType -from megatron.legacy.model.language_model import parallel_lm_logits, get_language_model -from megatron.legacy.model import LayerNorm -from megatron.legacy.model.utils import ( - openai_gelu, - get_linear_layer -) -from .module import MegatronModule - - -def t5_extended_attention_mask(attention_mask_list): - - def attn_mask_postprocess(attn_mask): - # [b, 1, s, s] - extended_attention_mask = attn_mask.unsqueeze(1) - return extended_attention_mask - - return [attn_mask_postprocess(attn_mask) for attn_mask in attention_mask_list] - - -def t5_position_ids(token_ids): - # Create position ids - seq_length = token_ids.size(1) - position_ids = torch.arange(seq_length, dtype=torch.long, - device=token_ids.device) - position_ids = position_ids.unsqueeze(0).expand_as(token_ids) - - return position_ids - - -class T5LMHead(MegatronModule): - """Masked LM head for T5 - - Args: - mpu_vocab_size: model parallel size of vocabulary. - parallel_output: wether output logits being distributed or not. - """ - - def __init__(self, mpu_vocab_size, parallel_output): - super(T5LMHead, self).__init__() - - self.bias = torch.nn.Parameter(torch.zeros(mpu_vocab_size)) - self.bias.model_parallel = True - self.bias.partition_dim = 0 - self.bias.stride = 1 - self.parallel_output = parallel_output - - def forward(self, hidden_states, word_embeddings_weight): - output = parallel_lm_logits(hidden_states, - word_embeddings_weight, - self.parallel_output, - bias=self.bias) - return output - - -class T5Model(MegatronModule): - """T5 Language model.""" - - def __init__(self, - config, - num_tokentypes=0, - parallel_output=True, - pre_process=True, - post_process=True, - add_encoder=True, - add_decoder=True): - super().__init__(config=config) - args = get_args() - - self.fp16_lm_cross_entropy = args.fp16_lm_cross_entropy - self.parallel_output = parallel_output - self.pre_process = pre_process - self.post_process = post_process - self.add_encoder = add_encoder - self.add_decoder = add_decoder - - self.language_model, self._language_model_key = get_language_model( - config=config, - num_tokentypes=num_tokentypes, - add_pooler=False, - add_encoder=add_encoder, - add_decoder=add_decoder, - encoder_attn_mask_type=AttnMaskType.padding, - pre_process=self.pre_process, - post_process=self.post_process) - - self.initialize_word_embeddings() - - if self.pre_process: - self.position_embeddings = self.language_model.embedding.position_embeddings - else: - self.position_embeddings = None - - if self.post_process and self.add_decoder: - self.lm_head = T5LMHead( - self.shared_embedding_or_output_weight().size(0), - parallel_output) - self._lm_head_key = 'lm_head' - - # Tells schedules.py that this model has a skip connection between the encoder's output and the decoder - # (and hence both the encoder and decoder's tensors are required for correct backprop). - self.xattn_needed = True - - def set_input_tensor(self, input_tensor): - """See megatron.legacy.model.transformer.set_input_tensor()""" - self.language_model.set_input_tensor(input_tensor) - - def forward(self, encoder_input_ids, decoder_input_ids, encoder_attn_mask, - decoder_attn_mask, encoder_decoder_attn_mask, - tokentype_ids=None, lm_labels=None, enc_hidden_states=None): - - # Converting the attention masks to proper parameter settings - encoder_attn_mask, decoder_attn_mask, encoder_decoder_attn_mask = t5_extended_attention_mask( - [encoder_attn_mask, decoder_attn_mask, encoder_decoder_attn_mask]) - - encoder_position_ids = t5_position_ids(encoder_input_ids) - decoder_position_ids = t5_position_ids(decoder_input_ids) - - lm_output = self.language_model(encoder_input_ids, - encoder_position_ids, - encoder_attn_mask, - decoder_input_ids, - decoder_position_ids, - decoder_attn_mask, - encoder_decoder_attn_mask, - tokentype_ids=tokentype_ids, - enc_hidden_states=enc_hidden_states) - - if self.post_process and self.add_decoder: - decoder_output, encoder_output = lm_output - # Output. [s, b, h] - lm_logits = self.lm_head(decoder_output, - self.shared_embedding_or_output_weight()) - - if lm_labels is None: - # [s b h] => [b s h] - return lm_logits.transpose(0,1).contiguous() - else: - # [b s] => [s b] - lm_labels = lm_labels.transpose(0,1).contiguous() - if self.fp16_lm_cross_entropy: - assert lm_logits.dtype == torch.half - lm_loss = tensor_parallel.vocab_parallel_cross_entropy(lm_logits, lm_labels) - else: - lm_loss = tensor_parallel.vocab_parallel_cross_entropy(lm_logits.float(), - lm_labels) - # [s b] => [b s] - lm_loss = lm_loss.transpose(0,1).contiguous() - return lm_loss - elif self.add_decoder and not self.add_encoder: - decoder_output, encoder_output = lm_output - return decoder_output - else: - encoder_output = lm_output - return encoder_output - - def state_dict_for_save_checkpoint(self, prefix='', keep_vars=False): - """For easy load when model is combined with other heads, - add an extra key.""" - - state_dict_ = {} - state_dict_[self._language_model_key] \ - = self.language_model.state_dict_for_save_checkpoint(prefix=prefix, - keep_vars=keep_vars) - if self.post_process and self.add_decoder: - state_dict_[self._lm_head_key] \ - = self.lm_head.state_dict_for_save_checkpoint(prefix=prefix, - keep_vars=keep_vars) - # Save word_embeddings. - if self.post_process and not self.pre_process and self.add_decoder: - state_dict_[self._word_embeddings_for_head_key] \ - = self.word_embeddings.state_dict(prefix=prefix, - keep_vars=keep_vars) - return state_dict_ - - def load_state_dict(self, state_dict, strict=True): - """Customized load.""" - - self.language_model.load_state_dict( - state_dict[self._language_model_key], strict=strict) - if self.post_process and self.add_decoder: - self.lm_head.load_state_dict(state_dict[self._lm_head_key], - strict=strict) - # Load word embeddings. - if self.post_process and not self.pre_process and self.add_decoder: - self.word_embeddings.load_state_dict( - state_dict[self._word_embeddings_for_head_key], strict=strict) diff --git a/pretrain_t5.py b/pretrain_t5.py index f849a5098c3..7a88251d0b9 100644 --- a/pretrain_t5.py +++ b/pretrain_t5.py @@ -71,7 +71,7 @@ def model_provider( add_decoder=True, config=None, pg_collection=None, -) -> Union[megatron.legacy.model.T5Model, T5Model]: +) -> T5Model: """Builds the model. Args: @@ -89,60 +89,50 @@ def model_provider( if config is None: config = core_transformer_config_from_args(args) - if args.use_legacy_models: - model = megatron.legacy.model.T5Model( - config=config, - num_tokentypes=0, - parallel_output=True, - pre_process=pre_process, - post_process=post_process, - add_encoder=add_encoder, - add_decoder=add_decoder, - ) - else: - encoder_config = deepcopy(config) - encoder_config.num_layers = args.encoder_num_layers - if args.pipeline_model_parallel_size > 1: - raise ValueError("Pipeline parallelism is not supported for T5.") + encoder_config = deepcopy(config) + encoder_config.num_layers = args.encoder_num_layers + + if args.pipeline_model_parallel_size > 1: + raise ValueError("Pipeline parallelism is not supported for T5.") - encoder_layers_per_pipeline = ( - encoder_config.num_layers // encoder_config.pipeline_model_parallel_size + encoder_layers_per_pipeline = ( + encoder_config.num_layers // encoder_config.pipeline_model_parallel_size + ) + decoder_layers_per_pipeline = config.num_layers // config.pipeline_model_parallel_size + + if args.transformer_impl == "local": + en_block_spec = get_t5_encoder_with_local_block_spec(encoder_layers_per_pipeline) + de_block_spec = get_t5_decoder_with_local_block_spec(decoder_layers_per_pipeline) + elif args.transformer_impl == "transformer_engine": + en_block_spec = get_t5_encoder_with_transformer_engine_block_spec( + encoder_layers_per_pipeline ) - decoder_layers_per_pipeline = config.num_layers // config.pipeline_model_parallel_size - - if args.transformer_impl == "local": - en_block_spec = get_t5_encoder_with_local_block_spec(encoder_layers_per_pipeline) - de_block_spec = get_t5_decoder_with_local_block_spec(decoder_layers_per_pipeline) - elif args.transformer_impl == "transformer_engine": - en_block_spec = get_t5_encoder_with_transformer_engine_block_spec( - encoder_layers_per_pipeline - ) - de_block_spec = get_t5_decoder_with_transformer_engine_block_spec( - decoder_layers_per_pipeline - ) - - print_rank_0('building T5 model ...') - model = T5Model( - config=config, - encoder_config=encoder_config, - transformer_encoder_layer_spec=en_block_spec, - transformer_decoder_layer_spec=de_block_spec, - vocab_size=args.padded_vocab_size, - max_sequence_length=args.max_position_embeddings, - pre_process=pre_process, - post_process=post_process, - fp16_lm_cross_entropy=args.fp16_lm_cross_entropy, - parallel_output=True, - share_embeddings_and_output_weights=not args.untie_embeddings_and_output_weights, - position_embedding_type=args.position_embedding_type, - rotary_percent=args.rotary_percent, - relative_attention_num_buckets=args.relative_attention_num_buckets, - relative_attention_max_distance=args.relative_attention_max_distance, - add_encoder=add_encoder, - add_decoder=add_decoder, + de_block_spec = get_t5_decoder_with_transformer_engine_block_spec( + decoder_layers_per_pipeline ) + print_rank_0('building T5 model ...') + model = T5Model( + config=config, + encoder_config=encoder_config, + transformer_encoder_layer_spec=en_block_spec, + transformer_decoder_layer_spec=de_block_spec, + vocab_size=args.padded_vocab_size, + max_sequence_length=args.max_position_embeddings, + pre_process=pre_process, + post_process=post_process, + fp16_lm_cross_entropy=args.fp16_lm_cross_entropy, + parallel_output=True, + share_embeddings_and_output_weights=not args.untie_embeddings_and_output_weights, + position_embedding_type=args.position_embedding_type, + rotary_percent=args.rotary_percent, + relative_attention_num_buckets=args.relative_attention_num_buckets, + relative_attention_max_distance=args.relative_attention_max_distance, + add_encoder=add_encoder, + add_decoder=add_decoder, + ) + return model From f1074b00d90cec9bdfda417e0b79e413ab595472 Mon Sep 17 00:00:00 2001 From: dimapihtar Date: Wed, 8 Apr 2026 06:35:44 -0700 Subject: [PATCH 2/2] remove t5 import Signed-off-by: dimapihtar --- megatron/legacy/model/__init__.py | 1 - 1 file changed, 1 deletion(-) diff --git a/megatron/legacy/model/__init__.py b/megatron/legacy/model/__init__.py index 1482c11f475..477df5f3a23 100644 --- a/megatron/legacy/model/__init__.py +++ b/megatron/legacy/model/__init__.py @@ -5,5 +5,4 @@ from .bert_model import BertModel from .gpt_model import GPTModel -from .t5_model import T5Model from .language_model import get_language_model