MCPcopy Create free account
hub / github.com/BeastyZ/ConvSearch-R1 / load_model

Function load_model

index/dense/models.py:103–131  ·  view source on GitHub ↗
(model_type, query_or_doc, model_path)

Source from the content-addressed store, hash-verified

101'''
102
103def 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

Callers 2

dense_indexingFunction · 0.90
dense_indexingFunction · 0.90

Calls 2

TCTColBERTClass · 0.85
from_pretrainedMethod · 0.80

Tested by

no test coverage detected