diff --git a/configs/full_configs/mlm_baseline.yaml b/configs/full_configs/mlm_baseline.yaml new file mode 100644 index 00000000..c3db9246 --- /dev/null +++ b/configs/full_configs/mlm_baseline.yaml @@ -0,0 +1,75 @@ +model: + core_model: + core_model_type: generic + num_layers: 10 + ffn: + ffn_type: swiglu + ffn_dim: 1536 + normalization: rms_norm + bias: false + attn: + attn_type: generic + num_heads: 16 + normalization: rms_norm + group_size: 4 + bias: false + is_causal: false + embedder: + tokenizer_type: gpt2 + embedding_model_type: generic + dataset_name: simple_en_wiki + lm_head: + normalization: rms_norm + bias: false + lm_head_type: generic + hidden_dim: 512 + context_window: 512 + vocab_size: 50257 + model_shell_type: standard + embedding_weight_tying: true + positional_encoding_type: rope +trainer: + dropout_scheduler: + dropout_type: constant + dropout: 0.1 + dataset: simple_en_wiki + training: + trainer_type: base_trainer + batch_size: 24 + gradient_accumulation_steps: 20 + max_iters: 25000 + lr_decay_iters: 25000 + warmup_iters: 5000 + eval_interval: 5000 + log_interval: 100 + eval_iters: 200 + checkpoint_interval: 1000000000.0 + run_profiler: false + + optimizer: + name: nanoGPTadamW + lr: 0.0006 + min_lr: 6.0e-05 + weight_decay: 0.1 + beta1: 0.9 + beta2: 0.95 + grad_clip: 1.0 + decay_lr: true + warmup_iters: 5000 + lr_scheduler: + name: cosine + dataloader: + name: mlm + loss_fn: + name: cross_entropy + +general: + logging: + wandb_log: true + wandb_project: SuperTinyLanguageModels + paths: + output_dir: outputs + data_dir: data + checkpoint_dir: checkpoints + seed: 489 + device: cuda diff --git a/trainers/build_trainers.py b/trainers/build_trainers.py index ff8ff23c..8092852f 100644 --- a/trainers/build_trainers.py +++ b/trainers/build_trainers.py @@ -10,6 +10,7 @@ BytePoolingDataloader, NextTokenMLMDataloader, ConversationalDataloader, + MLMDataloader ) from trainers.loss_fn import ( cross_entropy_loss_fn, @@ -86,6 +87,7 @@ def build_dropout_scheduler(trainer_cfg): "byte_pooling": BytePoolingDataloader, "next_token_mlm": NextTokenMLMDataloader, "conversational": ConversationalDataloader, + "mlm": MLMDataloader } diff --git a/trainers/dataloader.py b/trainers/dataloader.py index 9e26d638..f6de1c71 100644 --- a/trainers/dataloader.py +++ b/trainers/dataloader.py @@ -389,3 +389,59 @@ def get_batch(self, split="train", masking_pct=0.15): return X, (y, mask) +class MLMDataloader(BaseDataloader): + """ + Similarly to the generic dataloader, but mask out some tokens and + return the mask used. + """ + def get_batch(self, split="train", masking_pct=0.15): + """ + Get a train/val batch + """ + data = np.memmap( + os.path.join(self.tokenized_data_path, f"{split}.bin"), + dtype=np.uint16, + mode="r", + ) + + ## generate the index_ids for X + idxs = torch.randint(len(data) - self.context_window, (self.batch_size,)) + X = torch.stack( + [ + torch.from_numpy((data[i : i + self.context_window]).astype(np.int64)) + for i in idxs + ] + ) + + ## create a mask of True and False, where True will be Indicies to mask + mask = torch.rand(X.size()) < masking_pct + # mask &= X != + # mask &= X != + # mask &= X != + ## typically, there won't be any BOS or EOS tokens in the input_ids. + + ## create clones of the input_ids + mlm_data = X.clone() + labels = X.clone() + + ## get the indices of the mask + mask_idx = mask.nonzero(as_tuple=True) + + ## randomise the mask tokens + mask_idx_shuffle = torch.randperm(mask_idx[0].size(0)) + + ## get the indices of the mask tokens + tomask_idx = mask_idx_shuffle[:int(mask_idx[0].shape[0] * 0.8)] + torandom_idx = mask_idx_shuffle[int(mask_idx[0].shape[0] * 0.9):] + + ## mask the tokens + mlm_data[mask_idx[0][tomask_idx], mask_idx[1][tomask_idx]] = 0 + mlm_data[mask_idx[0][torandom_idx], mask_idx[1][torandom_idx]] = torch.randint(1, self.vocab_size, (torandom_idx.shape[0],)) + + ## create the labels + labels[~mask] = -1 ## -1 is the ignore_index + + X = mlm_data.pin_memory().to(self.device, non_blocking=True) + y = labels.pin_memory().to(self.device, non_blocking=True) + + return X, y \ No newline at end of file diff --git a/trainers/loss_fn.py b/trainers/loss_fn.py index a956e47f..4889482b 100644 --- a/trainers/loss_fn.py +++ b/trainers/loss_fn.py @@ -75,7 +75,7 @@ def compute_perplexity(logits, y, char_lengths, mask=None): # flatten both logits = logits.view(-1, logits.size(-1)) y = y.view(-1) - loss = torch.nn.functional.cross_entropy(logits, y, reduction="none") + loss = torch.nn.functional.cross_entropy(logits, y, reduction="none", ignore_index=-1) # B, S, 1 # unflatten loss = loss.view(B, seq_len)