| import argparse |
|
|
| from typing import Tuple, Dict, Optional |
|
|
| from llm import Llama |
|
|
| import torch |
|
|
| from main_models import T5FineTuner |
|
|
|
|
| def parse_arguments(): |
| parser = argparse.ArgumentParser(description="Run LLM Generation.") |
| parser.add_argument('--output_dir', type=str, default='data/gen_res', help='Output directory') |
| parser.add_argument('--ckpt_dir', type=str, default='Llama/Meta-Llama-3-8B', help='Checkpoint directory') |
| parser.add_argument('--tokenizer_path', type=str, default='Llama/Meta-Llama-3-8B/tokenizer.model', help='Checkpoint directory') |
| parser.add_argument('--max_seq_len', type=int, help='Maximum sequence length for LLM', default=512) |
| parser.add_argument('--max_batch_size', type=int, help='Maximum batch length for LLM', default=4) |
|
|
|
|
| |
| parser.add_argument('--model_name_or_path', type=str, default="t5-") |
| parser.add_argument('--tokenizer_name_or_path', type=str, default="t5-") |
| parser.add_argument('--model_info', type=str, default='base', choices=['small', 'large', 'base', '3b', '11b']) |
| |
| |
| parser.add_argument('--max_output_length', type=int, default=10) |
| parser.add_argument('--max_input_length', type=int, default=40) |
| parser.add_argument('--inf_max_input_length', type=int, default=40) |
| parser.add_argument('--random_gen', type=int, default=0, choices=[0, 1]) |
| parser.add_argument('--aug', type=int, default=0, choices=[0, 1]) |
| parser.add_argument('--contrastive_variant', type=str, default="", help='E_CL, ED_CL, doc_Reweight') |
| parser.add_argument('--query_type', type=str, default='gtq_qg', |
| help='gtq -- use ground turth query;' |
| 'qg -- use qg; ' |
| 'doc -- just use top64 doc token; ' |
| 'doc_aug -- use random doc token. ') |
| |
| parser.add_argument('--id_class', type=str, default='bert_k30_c30_1') |
| parser.add_argument('--trivia', type=int, default=0) |
| parser.add_argument('--nq', type=int, default=1) |
| parser.add_argument('--kary', type=int, default=30) |
| |
| parser.add_argument('--hard_negative', type=int, default=0) |
| parser.add_argument('--aug_query', type=int, default=0) |
| parser.add_argument('--aug_query_type', type=str, default='aug_query', help='aug_query, corrupted_query') |
| |
| parser.add_argument('--trivia_train_path', type=str, default='/home/t-qiweidi/Neural-Corpus-Indexer-NCI/Data_process/Trivia_dataset/train.tsv') |
| parser.add_argument('--nq_train_doc_newid_path', type=str, default='/home/t-qiweidi/Neural-Corpus-Indexer-NCI/Data_process/NQ_dataset/nq_train_doc_newid.tsv') |
| parser.add_argument('--trivia_qg_path', type=str, default='/home/t-qiweidi/Neural-Corpus-Indexer-NCI/Data_process/Trivia_dataset/trivia_512_qg.tsv') |
| parser.add_argument('--nq_qg_path', type=str, default='/home/t-qiweidi/Neural-Corpus-Indexer-NCI/Data_process/NQ_dataset/NQ_512_qg.tsv') |
| parser.add_argument('--trivia_title_cont_path', type=str, default='/home/t-qiweidi/Neural-Corpus-Indexer-NCI/Data_process/Trivia_dataset/trivia_title_cont.tsv') |
| parser.add_argument('--nq_title_abs_path', type=str, default='/home/t-qiweidi/Neural-Corpus-Indexer-NCI/Data_process/NQ_dataset/nq_title_abs.tsv') |
| parser.add_argument('--trivia_doc_aug_path', type=str, default='/home/t-qiweidi/Neural-Corpus-Indexer-NCI/Data_process/Trivia_dataset/trivia_doc_aug.tsv') |
| parser.add_argument('--nq_doc_aug_path', type=str, default='/home/t-qiweidi/Neural-Corpus-Indexer-NCI/Data_process/NQ_dataset/NQ_doc_aug.tsv') |
| |
| parser.add_argument('--retriever_ckpt', type=str, default='') |
|
|
| |
| parser.add_argument('--tree', type=int, default=1) |
| parser.add_argument('--position', type=int, default=1) |
| parser.add_argument('--num_layers', type=int, default=12) |
| parser.add_argument('--softmax', type=int, default=0, choices=[0, 1]) |
| parser.add_argument('--num_decoder_layers', type=int, default=6) |
| parser.add_argument('--d_ff', type=int, default=3072) |
| parser.add_argument('--d_model', type=int, default=768) |
| parser.add_argument('--num_heads', type=int, default=12) |
|
|
| parser.add_argument('--dropout_rate', type=float, default=0.1) |
| parser.add_argument('--decode_embedding', type=int, default=2, choices=[0, 1, 2]) |
| parser.add_argument('--hierarchic_decode', type=int, default=0, choices=[0, 1]) |
| parser.add_argument('--output_vocab_size', type=int, default=10) |
| parser.add_argument('--tie_word_embedding', type=int, default=0, choices=[0, 1]) |
| parser.add_argument('--tie_decode_embedding', type=int, default=1, choices=[0, 1]) |
| parser.add_argument('--contrastive', type=int, default=0) |
| parser.add_argument('--Rdrop', type=float, default=0.15, help='default to 0-0.3') |
| parser.add_argument('--Rdrop_only_decoder', type=int, default=0, |
| help='1-RDrop only for decoder, 0-RDrop only for all model', choices=[0,1]) |
| parser.add_argument('--Rdrop_loss', type=str, default='KL', choices=['KL', 'L2']) |
| parser.add_argument('--adaptor_decode', type=int, default=1, help='default to 0,1') |
| parser.add_argument('--adaptor_efficient', type=int, default=1, help='default to 0,1') |
| parser.add_argument('--adaptor_layer_num', type=int, default=4) |
| parser.add_argument('--embedding_distillation', type=float, default=0.0) |
| parser.add_argument('--weight_distillation', type=float, default=0.0) |
| parser.add_argument('--input_dropout', type=int, default=0) |
| parser.add_argument('--denoising', type=int, default=0) |
| parser.add_argument('--multiple_decoder', type=int, default=0) |
| parser.add_argument('--decoder_num', type=int, default=1) |
|
|
| parser.add_argument('--train_batch_size', type=int, default=4) |
| parser.add_argument('--eval_batch_size', type=int, default=2) |
|
|
| parser.add_argument('--t5_model_info', type=str, default='base', choices=['small', 'large', 'base', '3b', '11b']) |
|
|
| parser_args = parser.parse_args() |
| parser_args.tokenizer_name_or_path += parser_args.model_info |
| parser_args.model_name_or_path += parser_args.model_info |
|
|
| if parser_args.t5_model_info == 'base': |
| parser_args.num_layers = 12 |
| parser_args.num_decoder_layers = 6 |
| parser_args.d_ff = 3072 |
| parser_args.d_model = 768 |
| parser_args.num_heads = 12 |
| parser_args.d_kv = 64 |
| elif parser_args.t5_model_info == 'large': |
| parser_args.num_layers = 24 |
| parser_args.num_decoder_layers = 12 |
| parser_args.d_ff = 4096 |
| parser_args.d_model = 1024 |
| parser_args.num_heads = 16 |
| parser_args.d_kv = 64 |
| elif parser_args.t5_model_info == 'small': |
| parser_args.num_layers = 6 |
| parser_args.num_decoder_layers = 3 |
| parser_args.d_ff = 2048 |
| parser_args.d_model = 512 |
| parser_args.num_heads = 8 |
| parser_args.d_kv = 64 |
|
|
| return parser_args |
|
|
| |
|
|
|
|
|
|
| def print_info(args: argparse.Namespace): |
| print("INFO:") |
| print(f"MODEL: {args.llm_id}") |
|
|
|
|
|
|
|
|
| def main(): |
| args = parse_arguments() |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| model = T5FineTuner(args) |
| print(model.forward()) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|