(data_location='data',
model_dir=None,
batch_size=128,
maxlen=100,
seed=2,
bf16=False)
| 266 | |
| 267 | |
| 268 | def test(data_location='data', |
| 269 | model_dir=None, |
| 270 | batch_size=128, |
| 271 | maxlen=100, |
| 272 | seed=2, |
| 273 | bf16=False): |
| 274 | train_file = os.path.join(data_location, "local_train_splitByUser") |
| 275 | test_file = os.path.join(data_location, "local_test_splitByUser") |
| 276 | uid_voc = os.path.join(data_location, "uid_voc.pkl") |
| 277 | mid_voc = os.path.join(data_location, "mid_voc.pkl") |
| 278 | cat_voc = os.path.join(data_location, "cat_voc.pkl") |
| 279 | model_type = 'DIEN' |
| 280 | model_path = os.path, join(model_dir, |
| 281 | "ckpt_noshuff" + model_type + str(seed)) |
| 282 | |
| 283 | with tf.Session() as sess: |
| 284 | train_data = DataIterator(train_file, |
| 285 | uid_voc, |
| 286 | mid_voc, |
| 287 | cat_voc, |
| 288 | batch_size, |
| 289 | maxlen, |
| 290 | data_location=data_location) |
| 291 | test_data = DataIterator(test_file, |
| 292 | uid_voc, |
| 293 | mid_voc, |
| 294 | cat_voc, |
| 295 | batch_size, |
| 296 | maxlen, |
| 297 | data_location=data_location) |
| 298 | n_uid, n_mid, n_cat = train_data.get_n() |
| 299 | |
| 300 | if bf16: |
| 301 | model = Model_DIN_V2_Gru_Vec_attGru_Neg_bf16( |
| 302 | n_uid, n_mid, n_cat, EMBEDDING_DIM, HIDDEN_SIZE, |
| 303 | ATTENTION_SIZE) |
| 304 | else: |
| 305 | model = Model_DIN_V2_Gru_Vec_attGru_Neg(n_uid, n_mid, n_cat, |
| 306 | EMBEDDING_DIM, HIDDEN_SIZE, |
| 307 | ATTENTION_SIZE) |
| 308 | |
| 309 | model.restore(sess, model_path) |
| 310 | print( |
| 311 | 'test_auc: %.4f ----test_loss: %.4f ---- test_accuracy: %.4f ---- test_aux_loss: %.4f' |
| 312 | % eval(sess, test_data, model, model_path)) |
| 313 | |
| 314 | |
| 315 | def get_arg_parser(): |
no test coverage detected