Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
26 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 10 additions & 1 deletion megatron/model/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
5 changes: 4 additions & 1 deletion megatron/training.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
23 changes: 7 additions & 16 deletions pretrain_gpt.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)

Expand Down