Views
No views yet
1import torch
2from transformer_encoder import CustomConfig, TransformerEncoder
3
4config = CustomConfig()
5model = TransformerEncoder(config)
6
7input_ids = torch.randint(0, config.vocab_size, (2, 10))
8attention_mask = torch.ones_like(input_ids)
9
10output = model(input_ids, attention_mask)
11print(output.shape)torch.Size([2, 10, 256])