Views
No views yet
RagRetriever instance. The question encoder can be any model that can be loaded with AutoModel and the generator can be any model that can be loaded with AutoModelForSeq2SeqLM.1from transformers import RagTokenizer, RagRetriever, RagSequenceForGeneration, AutoTokenizer
2
3model = RagSequenceForGeneration.from_pretrained_question_encoder_generator("facebook/dpr-question_encoder-single-nq-base", "facebook/bart-large")
4
5question_encoder_tokenizer = AutoTokenizer.from_pretrained("facebook/dpr-question_encoder-single-nq-base")
6generator_tokenizer = AutoTokenizer.from_pretrained("facebook/bart-large")
7
8tokenizer = RagTokenizer(question_encoder_tokenizer, generator_tokenizer)
9model.config.use_dummy_dataset = True
10model.config.index_name = "exact"
11retriever = RagRetriever(model.config, question_encoder_tokenizer, generator_tokenizer)
12
13model.save_pretrained("./")
14tokenizer.save_pretrained("./")
15retriever.save_pretrained("./")config.index_name="legacy" and config.use_dummy_dataset=False.
The model can be fine-tuned as follows:1from transformers import RagTokenizer, RagRetriever, RagTokenForGeneration
2
3tokenizer = RagTokenizer.from_pretrained("facebook/rag-sequence-base")
4retriever = RagRetriever.from_pretrained("facebook/rag-sequence-base")
5model = RagTokenForGeneration.from_pretrained("facebook/rag-sequence-base", retriever=retriever)
6
7input_dict = tokenizer.prepare_seq2seq_batch("who holds the record in 100m freestyle", "michael phelps", return_tensors="pt")
8
9outputs = model(input_dict["input_ids"], labels=input_dict["labels"])
10
11loss = outputs.loss
12
13# train on loss