diff --git a/megatron/model/utils.py b/megatron/model/utils.py index 465e8aa4f..492f1f10b 100644 --- a/megatron/model/utils.py +++ b/megatron/model/utils.py @@ -40,7 +40,16 @@ def init_(tensor): def attention_mask_func(attention_scores, attention_mask): - attention_scores.masked_fill_(attention_mask, -10000.0) + args = get_args() + if args.curriculum_learning: + attention_mask_ = attention_mask + actual_seqlen = attention_scores.size()[2] + if actual_seqlen != attention_mask_.size()[2]: + # attention_mask has size [1, 1, seqlen, seqlen] + attention_mask_ = attention_mask_[:, :, :actual_seqlen, :actual_seqlen].contiguous() + attention_scores.masked_fill_(attention_mask_, -10000.0) + else: + attention_scores.masked_fill_(attention_mask, -10000.0) return attention_scores diff --git a/megatron/training.py b/megatron/training.py index 82919b4a7..dc01f6c4d 100644 --- a/megatron/training.py +++ b/megatron/training.py @@ -729,7 +729,10 @@ def train(forward_step_func, model, optimizer, lr_scheduler, get_num_microbatches() model[0].set_train_batch_size(global_batch_size) - + if args.curriculum_learning and \ + args.pipeline_model_parallel_size >= 1: + args.curriculum_seqlen = args.curriculum_scheduler.update_difficulty( \ + args.iteration + 1) loss_dict, skipped_iter, grad_norm, num_zeros_in_grad = \ train_step(forward_step_func, train_data_iterator, diff --git a/pretrain_gpt.py b/pretrain_gpt.py index 0e34c342e..108fef9b9 100644 --- a/pretrain_gpt.py +++ b/pretrain_gpt.py @@ -64,7 +64,6 @@ def model_provider(pre_process=True, post_process=True): # must be bool or the training crashes expecting bool, but getting Half args.attn_mask = attention_mask.to(torch.bool) - args.attn_mask_original = attention_mask.to(torch.bool) model = GPTModelPipe( num_tokentypes=0, @@ -146,21 +145,13 @@ def get_batch_pipe(data): prefix_indices=None, loss_on_targets_only=args.loss_on_targets_only ) - if args.curriculum_learning: - args.curriculum_seqlen = args.curriculum_scheduler.update_difficulty( \ - args.iteration + 1) - if args.curriculum_seqlen < tokens.size()[1]: - # seqlen-based curriculum learning - # tokens, position_ids, labels, loss_mask have size [batch size, seqlen] - tokens = tokens[:, :args.curriculum_seqlen].contiguous() - position_ids = position_ids[:, :args.curriculum_seqlen].contiguous() - labels = labels[:, :args.curriculum_seqlen].contiguous() - loss_mask = loss_mask[:, :args.curriculum_seqlen].contiguous() - actual_seqlen = tokens.size()[1] - if actual_seqlen != args.attn_mask.size()[2]: - # attention_mask has size [1, 1, seqlen, seqlen] - attention_mask = attention_mask[:, :, :actual_seqlen, :actual_seqlen].contiguous() - args.attn_mask = args.attn_mask_original[:, :, :actual_seqlen, :actual_seqlen].contiguous() + if args.curriculum_learning and args.curriculum_seqlen < tokens.size()[1]: + # seqlen-based curriculum learning + # tokens, position_ids, labels, loss_mask have size [batch size, seqlen] + tokens = tokens[:, :args.curriculum_seqlen].contiguous() + position_ids = position_ids[:, :args.curriculum_seqlen].contiguous() + labels = labels[:, :args.curriculum_seqlen].contiguous() + loss_mask = loss_mask[:, :args.curriculum_seqlen].contiguous() return (tokens, position_ids, attention_mask), (labels, loss_mask)