(model_type, query_or_doc, model_path)
| 101 | ''' |
| 102 | |
| 103 | def load_model(model_type, query_or_doc, model_path): |
| 104 | assert query_or_doc in ("query", "doc") |
| 105 | if model_type.lower() == "ance": |
| 106 | config = RobertaConfig.from_pretrained( |
| 107 | model_path, |
| 108 | finetuning_task="MSMarco", |
| 109 | ) |
| 110 | tokenizer = RobertaTokenizer.from_pretrained( |
| 111 | model_path, |
| 112 | do_lower_case=True |
| 113 | ) |
| 114 | model = ANCE.from_pretrained(model_path, config=config) |
| 115 | elif model_type.lower() == "dpr-nq": |
| 116 | if query_or_doc == "query": |
| 117 | tokenizer = DPRQuestionEncoderTokenizer.from_pretrained(model_path) |
| 118 | model = DPRQuestionEncoder.from_pretrained(model_path) |
| 119 | else: |
| 120 | tokenizer = DPRContextEncoderTokenizer.from_pretrained(model_path) |
| 121 | model = DPRContextEncoder.from_pretrained(model_path) |
| 122 | elif model_type.lower() == "tctcolbert": |
| 123 | tokenizer = AutoTokenizer.from_pretrained(model_path) |
| 124 | model = TCTColBERT(model_path) |
| 125 | else: |
| 126 | raise ValueError |
| 127 | |
| 128 | # tokenizer.add_tokens(["<CUR_Q>", "<CTX>", "<CTX_R>", "<CTX_Q>"]) |
| 129 | # model.resize_token_embeddings(len(tokenizer)) |
| 130 | |
| 131 | return tokenizer, model |
no test coverage detected