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']) ##################################################################### #Parameters for Dataset 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() # print("loading LLM") # ckpt_dir = args.ckpt_dir # tokenizer_path=args.tokenizer_path # max_seq_len = args.max_seq_len # max_batch_size = args.max_batch_size # llm = Llama.build( # ckpt_dir=ckpt_dir, # tokenizer_path=tokenizer_path, # max_seq_len=max_seq_len, # max_batch_size=max_batch_size, # ) # tokenizer = llm.tokenizer # print("LLM loaded") # prompt = ["hello, tell me your name"] # prompt_tokens = [tokenizer.encode(x, bos=True, eos=False) for x in prompt] # generation_tokens, generation_logprobs = llm.generate( # prompt_tokens, # max_gen_len= 20, # ) # print(generation_tokens) # print( # [tokenizer.decode(t) for t in generation_tokens] # ) # args = 1 model = T5FineTuner(args) print(model.forward()) if __name__ == "__main__": main()