| 97 | |
| 98 | |
| 99 | def load_model(args, device): |
| 100 | tokenizer = AutoTokenizer.from_pretrained(args["lm_path"]) |
| 101 | lm_model = AutoModelForSeq2SeqLM.from_pretrained(args["lm_path"]) |
| 102 | lm_model.eval() |
| 103 | lm_model.to(device) |
| 104 | if args["sbert"]: |
| 105 | sbert_model = SentenceTransformer('paraphrase-MiniLM-L6-v2') |
| 106 | else: |
| 107 | sbert_model = None |
| 108 | |
| 109 | if args["local_llm"] == "xgen": |
| 110 | local_llm.load() |
| 111 | assert local_llm.llm_model is not None |
| 112 | assert local_llm.llm_tokenizer is not None |
| 113 | print("Testing local LLM:" + args["local_llm"]) |
| 114 | print(local_llm.generate("Hello, who are you?")) # for testing |
| 115 | llm_model = local_llm.llm_model |
| 116 | else: |
| 117 | llm_model = None |
| 118 | return lm_model, tokenizer, sbert_model, llm_model |
| 119 | |
| 120 | |
| 121 | |