(args, interaction_list, user2index)
| 352 | |
| 353 | |
| 354 | def write_atomic_files(args, interaction_list, user2index): |
| 355 | print("[INFO] Writing atomic train/valid/test files") |
| 356 | |
| 357 | out_dir = os.path.join(args.output_path, args.dataset) |
| 358 | check_path(out_dir) |
| 359 | |
| 360 | total = len(interaction_list) |
| 361 | train_end = int(total * 0.8) |
| 362 | valid_end = int(total * 0.9) |
| 363 | |
| 364 | train_data = interaction_list[:train_end] |
| 365 | valid_data = interaction_list[train_end:valid_end] |
| 366 | test_data = interaction_list[valid_end:] |
| 367 | |
| 368 | print(f"[Split] train={len(train_data)}, valid={len(valid_data)}, test={len(test_data)}") |
| 369 | |
| 370 | def write_file(path, data): |
| 371 | with open(path, "w") as f: |
| 372 | f.write("user_id:token\titem_id_list:token_seq\titem_id:token\n") |
| 373 | for it in data: |
| 374 | user_original = it[0] |
| 375 | uid = user2index[user_original] |
| 376 | |
| 377 | hist = [str(x) for x in it[3]] |
| 378 | target = str(it[4]) |
| 379 | |
| 380 | hist = hist[-50:] # cap history length = 50 |
| 381 | f.write(f"{uid}\t{' '.join(hist)}\t{target}\n") |
| 382 | |
| 383 | write_file(os.path.join(out_dir, f"{args.dataset}.train.inter"), train_data) |
| 384 | write_file(os.path.join(out_dir, f"{args.dataset}.valid.inter"), valid_data) |
| 385 | write_file(os.path.join(out_dir, f"{args.dataset}.test.inter"), test_data) |
| 386 | |
| 387 | return train_data, valid_data, test_data |
| 388 | |
| 389 | |
| 390 | def build_item_features_amazon23(asin2meta, item2index): |
no test coverage detected