MCPcopy Create free account
hub / github.com/OpenBMB/DecT / load_model

Function load_model

src/run_dect.py:37–69  ·  view source on GitHub ↗
(name, size, path)

Source from the content-addressed store, hash-verified

35args = parser.parse_args()
36
37def 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
72def build_dataloader(dataset, template, verbalizer, tokenizer, tokenizer_wrapper_class, batch_size):

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected