| 35 | args = parser.parse_args() |
| 36 | |
| 37 | def load_model(name, size, path): |
| 38 | if name == "llama": |
| 39 | tokenizer = LlamaTokenizer.from_pretrained(path) |
| 40 | tokenizer.bos_token_id = 1 |
| 41 | tokenizer.eos_token_id = 2 |
| 42 | tokenizer.pad_token_id = 0 |
| 43 | model = LlamaForCausalLM.from_pretrained(path) |
| 44 | wrapper = LMTokenizerWrapper |
| 45 | if size == "7b": |
| 46 | hidden_size = 4096 |
| 47 | elif size == "13b": |
| 48 | hidden_size = 5120 |
| 49 | elif name == "alpaca": |
| 50 | tokenizer = LlamaTokenizer.from_pretrained(path) |
| 51 | tokenizer.bos_token_id = 1 |
| 52 | tokenizer.eos_token_id = 2 |
| 53 | tokenizer.pad_token_id = 0 |
| 54 | model = AutoModelForCausalLM.from_pretrained(path) |
| 55 | wrapper = LMTokenizerWrapper |
| 56 | hidden_size = 4096 |
| 57 | elif name == "vicuna": |
| 58 | tokenizer = LlamaTokenizer.from_pretrained(path) |
| 59 | tokenizer.bos_token_id = 1 |
| 60 | tokenizer.eos_token_id = 2 |
| 61 | tokenizer.pad_token_id = 0 |
| 62 | model = AutoModelForCausalLM.from_pretrained(path) |
| 63 | wrapper = LMTokenizerWrapper |
| 64 | hidden_size = 5120 |
| 65 | else: |
| 66 | model, tokenizer, model_config, wrapper = load_plm(args.model, args.model_name_or_path) |
| 67 | hidden_size = model_config.hidden_size |
| 68 | |
| 69 | return model, tokenizer, hidden_size, wrapper |
| 70 | |
| 71 | |
| 72 | def build_dataloader(dataset, template, verbalizer, tokenizer, tokenizer_wrapper_class, batch_size): |