CoolFace
Apppublic

q-future/Co-Instruct

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
29likes
models.py65 linesDownload Raw Back to insir_text
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