CNN恶意URL分类模型训练正常,Flask部署后预测失效求助
恶意URL分类CNN模型加载后预测失效问题
我近期开展了一个机器学习项目,训练了一款用于恶意URL分类的CNN模型。该模型在训练集和测试集上表现均良好,但将模型保存后,在EC2实例的Flask环境中重新加载并进行预测时,表现极差,仿佛从未经过训练。项目思路参考了一个GitHub仓库,恳请告知可能的问题原因。
模型构建代码
class CNN(object): def __init__(self) -> None: super(CNN, self).__init__() self.max_len = 75 self.emb_dim = 32 self.max_vocab_len = 100 self.W_reg = regularizers.l2(1e-4) def sum_1d(self, X): return K.sum(X, axis=1) def get_conv_layer(self, emb, kernel_size=5, filters=256): conv = Convolution1D(kernel_size=kernel_size, filters=filters, padding='same')(emb) conv = ELU()(conv) conv = Lambda(self.sum_1d, output_shape=(filters,))(conv) # conv = BatchNormalization()(conv) conv = Dropout(0.5)(conv) return conv def build_model(self): main_input = Input(shape=(self.max_len,), dtype='int32', name='main_input') emb = Embedding(input_dim=self.max_vocab_len, output_dim=self.emb_dim, input_length=self.max_len, embeddings_regularizer=self.W_reg)(main_input) emb = Dropout(0.25)(emb) conv1 = self.get_conv_layer(emb, kernel_size=2, filters=256) conv2 = self.get_conv_layer(emb, kernel_size=3, filters=256) conv3 = self.get_conv_layer(emb, kernel_size=4, filters=256) conv4 = self.get_conv_layer(emb, kernel_size=5, filters=256) merged = concatenate([conv1, conv2, conv3, conv4], axis=1) hidden1 = Dense(1024)(merged) hidden1 = ELU()(hidden1) hidden1 = BatchNormalization()(hidden1) hidden1 = Dropout(0.5)(hidden1) hidden2 = Dense(1024)(hidden1) hidden2 = ELU()(hidden2) hidden2 = BatchNormalization()(hidden2) hidden2 = Dropout(0.5)(hidden2) output = Dense(1, activation='sigmoid', name='output')(hidden2) model = Model(inputs=[main_input], outputs=[output]) adam = Adam(lr=1e-4, beta_1=0.9, beta_2=0.999, epsilon=1e-08, decay=0.0) model.compile(optimizer=adam, loss='binary_crossentropy', metrics=['accuracy']) return
模型保存代码
# Function to Save Model def save_model(model, json_dir, weights_dir): # have h5py installed if Path(json_dir).is_file(): os.remove(json_dir) json_string = model.to_json() with open(json_dir, 'w') as f: json.dump(json_string, f) if Path(weights_dir).is_file(): os.remove(weights_dir) model.save_weights(weights_dir)
模型加载代码
def load_model(): json_dir = os.path.join('models\cnn_lstm\conv_lstm.json') weights_dir = os.path.join('models\cnn_lstm\conv_lstm.h5') with open(json_dir, 'r') as f: model_json = json.load(f) model = model_from_json(model_json) print("here") model.load_weights(weights_dir) return model
Flask服务代码
import argparse import json import os import pandas as pd from string import printable import pickle as p import flask from flask import Flask, jsonify, request import tensorflow as tf from tensorflow import keras from keras_preprocessing.sequence import pad_sequences from keras.models import model_from_json def load_model(): json_dir = os.path.join('models\cnn_lstm\conv_lstm.json') weights_dir = os.path.join('models\cnn_lstm\conv_lstm.h5') with open(json_dir, 'r') as f: model_json = json.load(f) model = model_from_json(model_json) print("here") model.load_weights(weights_dir) return model def prepare_urls(urls): url_int_tokens = [ [printable.index(x) + 1 for x in url if x in printable] for url in urls] # Cut URL string at max_len or pad with zeros if shorter, return result return pad_sequences(url_int_tokens, maxlen=75) def predict_type(urls): model = load_model() predictions = model.predict(prepare_urls(urls)) print(predictions) url_type = [] for prediction in predictions: if prediction > 0.5: url_type.append('Malicious') else: url_type.append('Safe') return url_type app = Flask(__name__) @app.route('/') def welcome(): return 'Malicious URL Identification' @app.route('/predict', methods=['GET', 'POST']) def predict(): if flask.request.method == 'POST': try: print(request.json) urls = pd.DataFrame(request.json, index=[0]) return jsonify(predict_type(urls)) except Exception as e: return jsonify({ "Exception": e }) if __name__ == "__main__": parser = argparse.ArgumentParser() parser.addargument('--model_name', type=str, default='cnn_lstm_model', metavar='N', help='Provide the name of the model (default value is cnn_lstm_model)') args = parser.parse_args() app.run(port=8080, host='0.0.0.0')
异常发现
打印预测值时,所有输入的预测结果均相同,如下图所示:
内容的提问来源于stack exchange,提问作者YKM
相关产品推荐
相关产品推荐

