e2e2 / src /NCI /test.py
Qiwei2000's picture
initialize
6548a2f
Raw
History Blame Contribute Delete
8.41 kB
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()