from https://www.kaggle.com/knowledgegrappler/a-simple-nn-solution-with-keras-0-48611-pl
from keras.layers import Input, Dropout, Dense, BatchNormalization, Activation, concatenate, GRU, Embedding, Flatten, BatchNormalization
from keras.models import Model
from keras.callbacks import ModelCheckpoint, Callback, EarlyStopping
from keras import backend as K
def get_callbacks(filepath, patience=2):
es = EarlyStopping('val_loss', patience=patience, mode="min")
msave = ModelCheckpoint(filepath, save_best_only=True)
return [es, msave]
def rmsle_cust(y_true, y_pred):
first_log = K.log(K.clip(y_pred, K.epsilon(), None) + 1.)
second_log = K.log(K.clip(y_true, K.epsilon(), None) + 1.)
return K.sqrt(K.mean(K.square(first_log - second_log), axis=-1))
def get_model():
#params
dr_r = 0.1
#Inputs
name = Input(shape=[X_train["name"].shape[1]], name="name")
item_desc = Input(shape=[X_train["item_desc"].shape[1]], name="item_desc")
brand_name = Input(shape=[1], name="brand_name")
category_name = Input(shape=[1], name="category_name")
item_condition = Input(shape=[1], name="item_condition")
num_vars = Input(shape=[X_train["num_vars"].shape[1]], name="num_vars")
#Embeddings layers
emb_name = Embedding(MAX_TEXT, 50)(name)
emb_item_desc = Embedding(MAX_TEXT, 50)(item_desc)
emb_brand_name = Embedding(MAX_BRAND, 10)(brand_name)
emb_category_name = Embedding(MAX_CATEGORY, 10)(category_name)
emb_item_condition = Embedding(MAX_CONDITION, 5)(item_condition)
#rnn layer
rnn_layer1 = GRU(16) (emb_item_desc)
rnn_layer2 = GRU(8) (emb_name)
#main layer
main_l = concatenate([
Flatten() (emb_brand_name)
, Flatten() (emb_category_name)
, Flatten() (emb_item_condition)
, rnn_layer1
, rnn_layer2
, num_vars
])
main_l = Dropout(dr_r) (Dense(128) (main_l))
main_l = Dropout(dr_r) (Dense(64) (main_l))
#output
output = Dense(1, activation="linear") (main_l)
#model
model = Model([name, item_desc, brand_name
, category_name, item_condition, num_vars], output)
model.compile(loss="mse", optimizer="adam", metrics=["mae", rmsle_cust])
return model
model = get_model()
model.summary()