Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
75 changes: 75 additions & 0 deletions configs/full_configs/mlm_baseline.yaml
Original file line number Diff line number Diff line change
@@ -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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

change this to stlm

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

sorry, can I clarify what should be changed to stlm?

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

the dataset_name, but its okay, this will all be reworked later

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
2 changes: 2 additions & 0 deletions trainers/build_trainers.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
BytePoolingDataloader,
NextTokenMLMDataloader,
ConversationalDataloader,
MLMDataloader
)
from trainers.loss_fn import (
cross_entropy_loss_fn,
Expand Down Expand Up @@ -86,6 +87,7 @@ def build_dropout_scheduler(trainer_cfg):
"byte_pooling": BytePoolingDataloader,
"next_token_mlm": NextTokenMLMDataloader,
"conversational": ConversationalDataloader,
"mlm": MLMDataloader
}


Expand Down
56 changes: 56 additions & 0 deletions trainers/dataloader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 != <pad_token_id>
# mask &= X != <mask_token_id>
# mask &= X != <unk_token_id>
## 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
2 changes: 1 addition & 1 deletion trainers/loss_fn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

we are 100% sure this has no impacts elsewhere? should be okay but..

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes. Ignore_index's default value is -100. So, changing it from -100 to -1 would only impact existing label (y) assignments of -1 or -100.

  1. I searched our code base for the assignment of -100, there was none.
  2. I searched our code base for the assignment of -1. Apart from the new MLMDataLoader using label[~mask]=-1, other assignments were used in arguments such as 'dim=-1', or just as index values.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

image

# B, S, 1
# unflatten
loss = loss.view(B, seq_len)
Expand Down