ContactBridge

Log | Files | Refs | README

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}')