NSChristopher's picture
Upload ProjectionHead.py
93e5552
Raw
History Blame Contribute Delete
977 Bytes
import pandas as pd
import numpy as np
import ast
import os
import torch
import torch.nn as nn
import torch.nn.functional as F
class ProjectionHead(nn.Module):
def __init__(self, input_dims, num_projection_layers, projection_dims, dropout_rate):
super(ProjectionHead, self).__init__()
# layers
self.dense = nn.Linear(input_dims, projection_dims)
self.dense_2 = nn.Linear(projection_dims, projection_dims)
self.dropout = nn.Dropout(dropout_rate)
self.layer_norm = nn.LayerNorm(projection_dims)
self.num_projection_layers = num_projection_layers
def forward(self, embeddings):
project_embeddings = self.dense(embeddings)
for _ in range(self.num_projection_layers):
x = F.gelu(project_embeddings)
x = self.dense_2(x)
x = self.dropout(x)
x = x + project_embeddings
project_embeddings = self.layer_norm(x)
return project_embeddings