Downloads · 30 days
95
4% of all-time downloads
prhegde/t5-query-reformulation-RL
t5-query-reformulation-RL is a text generation model from prhegde. Use it when you need the model to write or continue text. It is set up for transformers. The card lists the license as apache-2.0.
This is a generative model designed specifically for search query rewriting, employing a sequence-to-sequence architecture for generating reformulated queries. It leverages a Reinforcement Learning framework to furthe…
Downloads · 30 days
95
4% of all-time downloads
All-time downloads
2.7K
Public
Parameters
223M
1.8 GB on disk
Likes
7
Public
Click a slice to open those files.
.safetensors892 MB · 100%
From the Hugging Face model README
This is a generative model designed specifically for search query rewriting, employing a sequence-to-sequence architecture for generating reformulated queries. It leverages a Reinforcement Learning framework to further boost performance, integrating a policy gradient algorithm. The model is trained with reward functions aimed at diversifying the generated queries by paraphrasing keywords. It can be integrated with sparse retrieval methods, such as bm25-based retrieval, to enhance document recall in search.
Query rewriting for search (web, e-commerce), Virtual assistants and chatbots, Information retrieval
Training Procedure
Refer here for more details.
For optimal utilization of this model, use sampling with repetition penalty to generate diverse samples. Below is the provided sample code.
import torch
from transformers import T5ForConditionalGeneration, T5Tokenizer
MODEL_ID = "prhegde/t5-query-reformulation-RL"
tokenizer = T5Tokenizer.from_pretrained(MODEL_ID)
model = T5ForConditionalGeneration.from_pretrained(MODEL_ID)
model.eval()
input_sequence = "how to bake great cookie"
input_ids = tokenizer(input_sequence, return_tensors="pt").input_ids
print(f'Input: {input_sequence}')
nsent = 4
with torch.no_grad():
for i in range(nsent):
output = model.generate(input_ids, max_length=35, num_beams=1, do_sample=True, repetition_penalty=1.8)
target_sequence = tokenizer.decode(output[0], skip_special_tokens=True)
print(f'Target: {target_sequence}')