MCPcopy Create free account
hub / github.com/NeuSpeech/EEG-To-Text / __init__

Method __init__

model_decoding.py:250–266  ·  view source on GitHub ↗
(self, pretrained_text_encoder, in_feature = 840, eeg_encoder_nhead=8, eeg_encoder_dim_feedforward = 2048, embed_dim = 768)

Source from the content-addressed store, hash-verified

248
249class 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

Callers

nothing calls this directly

Calls 3

PositionalEncodingClass · 0.70
PoolerClass · 0.70
__init__Method · 0.45

Tested by

no test coverage detected