CoolFace
Apppublic

tcyang/TransDis-CreativityAutoAssessment-V2

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
models.py84 linesDownload Raw Back to utils
1from functools import lru_cache2 3import torch4from loguru import logger5from sentence_transformers import SentenceTransformer6from transformers import AutoTokenizer, AutoModel7 8DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'9 10list_models = [11    'sentence-transformers/paraphrase-multilingual-mpnet-base-v2',12    'sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2',13    'sentence-transformers/all-mpnet-base-v2',14    'sentence-transformers/all-MiniLM-L12-v2',15    'cyclone/simcse-chinese-roberta-wwm-ext',16    'bert-base-chinese',17    'IDEA-CCNL/Erlangshen-SimCSE-110M-Chinese',18    'Qwen/Qwen3-Embedding-0.6B',19]20 21 22class SBert:23    def __init__(self, path):24        logger.info(f'Start loading {self.__class__} from {path} ...')25        self.model = SentenceTransformer(path, device=DEVICE)26        logger.info(f'Load {self.__class__} from {path} ...')27 28    @lru_cache(maxsize=10000)29    def __call__(self, x) -> torch.Tensor:30        y = self.model.encode(x, convert_to_tensor=True)31        return y32 33 34class ModelWithPooling:35    def __init__(self, path):36        logger.info(f'Start loading {self.__class__} from {path} ...')37        self.tokenizer = AutoTokenizer.from_pretrained(path)38        self.model = AutoModel.from_pretrained(path)39        logger.info(f'Load {self.__class__} from {path} ...')40 41    @lru_cache(maxsize=100)42    @torch.no_grad()43    def __call__(self, text: str, pooling='mean'):44        inputs = self.tokenizer(text, padding=True, truncation=True, return_tensors="pt")45        outputs = self.model(**inputs, output_hidden_states=True)46 47        if pooling == 'cls':48            o = outputs.last_hidden_state[:, 0]  # [b, h]49 50        elif pooling == 'pooler':51            o = outputs.pooler_output  # [b, h]52 53        elif pooling in ['mean', 'last-avg']:54            last = outputs.last_hidden_state.transpose(1, 2)  # [b, h, s]55            o = torch.avg_pool1d(last, kernel_size=last.shape[-1]).squeeze(-1)  # [b, h]56 57        elif pooling == 'first-last-avg':58            first = outputs.hidden_states[1].transpose(1, 2)  # [b, h, s]59            last = outputs.hidden_states[-1].transpose(1, 2)  # [b, h, s]60            first_avg = torch.avg_pool1d(first, kernel_size=last.shape[-1]).squeeze(-1)  # [b, h]61            last_avg = torch.avg_pool1d(last, kernel_size=last.shape[-1]).squeeze(-1)  # [b, h]62            avg = torch.cat((first_avg.unsqueeze(1), last_avg.unsqueeze(1)), dim=1)  # [b, 2, h]63            o = torch.avg_pool1d(avg.transpose(1, 2), kernel_size=2).squeeze(-1)  # [b, h]64 65        else:66            raise Exception(f'Unknown pooling {pooling}')67 68        o = o.squeeze(0)69        return o70 71 72def test_sbert():73    m = SBert('bert-base-chinese')74    o = m('hello')75    print(o.size())76    assert o.size() == (768,)77 78 79def test_hf_model():80    m = ModelWithPooling('IDEA-CCNL/Erlangshen-SimCSE-110M-Chinese')81    o = m('hello', pooling='cls')82    print(o.size())83    assert o.size() == (768,)84