q-future/Co-Instruct
29
1import torch2from torch import nn3import torch.nn.functional as F4from transformers import DistilBertModel, DistilBertTokenizer, AutoModel, AutoTokenizer5import os6 7# Models that use mean pooling8POOL_MODELS = {"sentence-transformers/all-MiniLM-L6-v2", "TaylorAI/bge-micro-v2"}9 10#Mean Pooling - Take attention mask into account for correct averaging11def mean_pooling(model_output, attention_mask):12 token_embeddings = model_output[0] #First element of model_output contains all token embeddings13 input_mask_expanded = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()14 return torch.sum(token_embeddings * input_mask_expanded, 1) / torch.clamp(input_mask_expanded.sum(1), min=1e-9)15 16 17class LanguageModel(nn.Module):18 def __init__(self, model='distilbert-base-uncased'):19 super(LanguageModel, self).__init__()20 21 self.tokenizer = AutoTokenizer.from_pretrained(model)22 self.model = AutoModel.from_pretrained(model)23 self.model_name = model24 # Remove the CLIP vision tower25 if "clip" in self.model_name:26 self.model.vision_model = None27 # Freeze the pre-trained parameters (very important)28 for param in self.model.parameters():29 param.requires_grad = False30 31 # Make sure to set evaluation mode (also important)32 self.model.eval()33 34 def forward(self, text_batch):35 inputs = self.tokenizer(text_batch, padding=True, truncation=True, return_tensors="pt")36 with torch.no_grad(): # Ensure no gradients are computed for this forward pass37 38 if "clip" in self.model_name:39 sentence_embedding = self.model.get_text_features(**inputs)40 return sentence_embedding41 42 outputs = self.model(**inputs)43 44 if any(model in self.model_name for model in POOL_MODELS):45 sentence_embeddings = mean_pooling(outputs, inputs['attention_mask'])46 # Normalize embeddings47 sentence_embedding = F.normalize(sentence_embeddings, p=2, dim=1)48 else:49 sentence_embedding = outputs.last_hidden_state[:, 0, :]50 return sentence_embedding51 52 53class LMHead(nn.Module):54 def __init__(self, embedding_dim=384, hidden_dim=256, num_classes=4):55 super(LMHead, self).__init__()56 57 self.fc1 = nn.Linear(embedding_dim, hidden_dim)58 #self.gelu = nn.GELU()59 self.fc2 = nn.Linear(hidden_dim, num_classes)60 61 def forward(self, x):62 embd = self.fc1(x)63 embd = F.normalize(embd, p=2, dim=1)64 deg_pred = self.fc2(embd)65 return embd, deg_pred