n_layers = 1 class BERT(nn.Module): def __init__(self): super(BERT, self).__init__() #for converting tokens into vector embeddings self.embedding = Embedding() #encoder blocks self.encoder_blocks = nn.ModuleList([EncoderBlock() for _ in range(n_layers)]) #for decoding a word vector (or tensor of them) into token predictions self.decoder = nn.Linear(d_model, tokenizer.vocab_size, bias=False) #for converting the first output token into a binary classification self.classifier = nn.Linear(d_model, 1, bias=False) def forward(self, x, seg, masked_token_locations): #x of shape [batch x seq_len x model_dim] embeddings = self.embedding(x, seg) x = embeddings for block in self.encoder_blocks: x = block(x) #passing first token through classifier clsf_logits = self.classifier(x[:,0,:]) #passing masked tokens through decoder masked_token_embeddings = embeddings[masked_token_locations.bool()] token_logits = self.decoder(masked_token_embeddings) return clsf_logits, token_logits