Views
No views yet
1from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
2import torch
3
4model_name = 'doc2query/msmarco-german-mt5-base-v1'
5tokenizer = AutoTokenizer.from_pretrained(model_name)
6model = AutoModelForSeq2SeqLM.from_pretrained(model_name)
7
8text = "Python ist eine universelle, üblicherweise interpretierte, höhere Programmiersprache. Sie hat den Anspruch, einen gut lesbaren, knappen Programmierstil zu fördern. So werden beispielsweise Blöcke nicht durch geschweifte Klammern, sondern durch Einrückungen strukturiert."
9
10
11def create_queries(para):
12 input_ids = tokenizer.encode(para, return_tensors='pt')
13 with torch.no_grad():
14 # Here we use top_k / top_k random sampling. It generates more diverse queries, but of lower quality
15 sampling_outputs = model.generate(
16 input_ids=input_ids,
17 max_length=64,
18 do_sample=True,
19 top_p=0.95,
20 top_k=10,
21 num_return_sequences=5
22 )
23
24 # Here we use Beam-search. It generates better quality queries, but with less diversity
25 beam_outputs = model.generate(
26 input_ids=input_ids,
27 max_length=64,
28 num_beams=5,
29 no_repeat_ngram_size=2,
30 num_return_sequences=5,
31 early_stopping=True
32 )
33
34
35 print("Paragraph:")
36 print(para)
37
38 print("\nBeam Outputs:")
39 for i in range(len(beam_outputs)):
40 query = tokenizer.decode(beam_outputs[i], skip_special_tokens=True)
41 print(f'{i + 1}: {query}')
42
43 print("\nSampling Outputs:")
44 for i in range(len(sampling_outputs)):
45 query = tokenizer.decode(sampling_outputs[i], skip_special_tokens=True)
46 print(f'{i + 1}: {query}')
47
48create_queries(text)
49model.generate() is non-deterministic for top_k/top_n sampling. It produces different queries each time you run it.train_script.py in this repository.