spam_detector.py (1248B)
1 import pickle 2 import os 3 import tensorflow as tf 4 from keras.preprocessing.sequence import pad_sequences 5 6 7 class SpamDetector: 8 def __init__(self): 9 current_dir = os.getcwd() 10 model_save_path = os.path.join(current_dir, 'MessageTagging/saved_model') 11 tokenizer_save_path = os.path.join(current_dir, 'MessageTagging/tokenizer.pickle') 12 13 self.model = self.load_model(model_save_path) 14 self.tokenizer = self.load_tokenizer(tokenizer_save_path) 15 16 @staticmethod 17 def load_model(model_save_path): 18 return tf.keras.models.load_model(model_save_path) 19 20 @staticmethod 21 def load_tokenizer(tokenizer_save_path): 22 with open(tokenizer_save_path, 'rb') as handle: 23 return pickle.load(handle) 24 25 def detect_spam(self, text): 26 sequences = self.tokenizer.texts_to_sequences([text]) 27 padded_sequences = pad_sequences(sequences, padding='post', maxlen=5530) 28 29 return self.model.predict(padded_sequences)[0][0] 30 31 32 if __name__ == "__main__": 33 spam_detector = SpamDetector() 34 35 prediction = spam_detector.detect_spam('This is a test message') 36 37 if prediction < .05: 38 print("not spam") 39 else: 40 print("spam") 41 42 print(f'Prediction: {prediction:.2f}')