snorkelai/instruction-response-quality
4
1 2from typing import Dict, List, Union, Optional3import os4from pathlib import Path5import json6import joblib7import pandas as pd8import nltk9from transformers import AutoModel, AutoTokenizer10import torch11import numpy as np12from sklearn.base import TransformerMixin13 14LOCAL_PATH = Path(__file__).parent15nltk.data.path.append(str(LOCAL_PATH/"nltk_data"))16 17class SimcseGenerator(TransformerMixin):18 def __init__(19 self, batch_size: int =16, model_name: str = "princeton-nlp/unsup-simcse-bert-base-uncased"20 ) -> None:21 22 self.model_name = model_name23 24 self.device = torch.device('cpu')25 26 tokenizer = AutoTokenizer.from_pretrained(model_name)27 model = AutoModel.from_pretrained(model_name).to(self.device)28 29 self.tokenizer = tokenizer30 self.model = model31 self.batch_size = batch_size32 33 def transform(self, X: np.ndarray) -> np.ndarray:34 batch_size = (35 16 # any larger, and we risk running out of memory on EC2 dev instances36 )37 38 embeddings = []39 40 for start in range(0, len(X), batch_size):41 end = min(len(X), start + batch_size)42 inputs = self.tokenizer(43 X[start:end],44 padding=True,45 truncation=True,46 return_tensors="pt",47 )48 with torch.no_grad():49 inputs = inputs.to(self.device)50 batch_embeddings = self.model(51 **inputs, output_hidden_states=True, return_dict=True52 ).pooler_output53 embeddings.append(batch_embeddings.cpu().detach().numpy())54 55 embeddings = np.concatenate(embeddings)56 embeddings /= np.sqrt(np.square(embeddings).sum(axis=1))[:,np.newaxis]57 58 return embeddings59 60class EndpointHandler():61 def __init__(self, path: str = ""):62 63 if len(path)==0:64 path = LOCAL_PATH65 else:66 path = Path(path)67 68 with open(path/'stop_words.json','r') as fp:69 self.stop_words = set(json.load(fp))70 71 with open(path/'instruction_label_map.json','r') as fp:72 self.instruction_label_map = json.load(fp)73 self.instruction_label_map = {int(k):v for k,v in self.instruction_label_map.items()}74 75 self.instruction_pipeline = joblib.load(path/'instruction_classification_pipeline.joblib')76 self.response_pipeline = joblib.load(path/'response_quality_pipeline.joblib')77 78 self.simcse_generator = SimcseGenerator()79 80 def _get_stop_word_proportion(self, s):81 s = s.lower()82 try:83 words = nltk.tokenize.word_tokenize(s)84 except:85 words = nltk.tokenize.word_tokenize(s[1:])86 87 if len(words)==0:88 return 089 else:90 return sum(x in self.stop_words for x in words) / len(words)91 92 93 def predict_instruction_classes(self, df: pd.DataFrame) -> np.ndarray:94 instruction_classes = self.instruction_pipeline.predict(df)95 instruction_class_confidence = self.instruction_pipeline.predict_proba(df).max(axis=1)96 return np.array(list(map(lambda x: self.instruction_label_map[x], instruction_classes))), instruction_class_confidence97 98 def compute_response_quality_feature_space(self, df: pd.DataFrame, instruction_classes: Optional[np.ndarray] = None):99 100 if instruction_classes is None:101 instruction_classes, _ = self.predict_instruction_classes(df)102 103 instruction_class_set = [self.instruction_label_map[i] for i in range(len(self.instruction_label_map))]104 105 instruction_classes_onehot = pd.DataFrame(instruction_classes[:,np.newaxis]==np.array(instruction_class_set)[np.newaxis,:], columns=instruction_class_set).astype(float)106 107 df1 = pd.concat([df,instruction_classes_onehot], axis=1)108 109 df1['instruction_response_similarity'] = (self.simcse_generator.transform(df['instruction'].tolist()) * self.simcse_generator.transform(df['response'].tolist())).sum(axis=1)110 111 df1['token_number'] = df1['response'].str.split().apply(len)112 df1['stop_word_proportion'] = df1['response'].apply(self._get_stop_word_proportion)113 114 return df1115 116 def predict_response_quality(self, df, instruction_classes):117 df1 = self.compute_response_quality_feature_space(df, instruction_classes)118 return self.response_pipeline.predict_proba(df1)[:,1]119 120 121 def __call__(self, data: Dict[str, Union[Dict, List]]):122 123 inputs = data['inputs']124 125 is_dict = isinstance(inputs, dict)126 127 if is_dict:128 df = pd.DataFrame([inputs])129 else:130 df = pd.DataFrame(inputs)131 132 df = df.fillna('')133 134 if 'dataset' not in df.columns:135 df['dataset'] = ''136 137 instruction_classes, instruction_class_confidences = self.predict_instruction_classes(df)138 139 predictions = [{'instruction class': instruction_class, 'instruction class confidence': instruction_class_confidence} for instruction_class, instruction_class_confidence in zip(instruction_classes, instruction_class_confidences)]140 141 if 'response' in df.columns:142 response_qualities = self.predict_response_quality(df, instruction_classes)143 for i,response_quality in enumerate(response_qualities):144 predictions[i].update({'response quality': response_quality})145 146 if is_dict:147 return predictions[0]148 else:149 return predictions