(self, pretrained_text_encoder, in_feature = 840, eeg_encoder_nhead=8, eeg_encoder_dim_feedforward = 2048, embed_dim = 768)
| 248 | |
| 249 | class ContrastiveBrainTextEncoder(nn.Module): |
| 250 | def __init__(self, pretrained_text_encoder, in_feature = 840, eeg_encoder_nhead=8, eeg_encoder_dim_feedforward = 2048, embed_dim = 768): |
| 251 | super(ContrastiveBrainTextEncoder, self).__init__() |
| 252 | # EEG Encoder |
| 253 | self.positional_embedding = PositionalEncoding(in_feature) |
| 254 | self.encoder_layer = nn.TransformerEncoderLayer(d_model=in_feature, nhead=eeg_encoder_nhead, dim_feedforward = eeg_encoder_dim_feedforward, batch_first=True) |
| 255 | self.EEG_Encoder = nn.TransformerEncoder(self.encoder_layer, num_layers=6) |
| 256 | self.EEG_pooler = Pooler(in_feature) |
| 257 | self.ln_final = nn.LayerNorm(in_feature) # to be considered |
| 258 | |
| 259 | # project to text embedding |
| 260 | self.EEG_projection = nn.Parameter(torch.empty(in_feature, embed_dim)) |
| 261 | |
| 262 | # Text Encoder |
| 263 | self.TextEncoder = pretrained_text_encoder |
| 264 | |
| 265 | # learned temperature parameter |
| 266 | self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07)) |
| 267 | |
| 268 | def forward(self, input_EEG_features, input_EEG_attn_mask, input_ids, input_text_attention_masks): |
| 269 | # add positional embedding |
nothing calls this directly
no test coverage detected