Views
No views yet
1from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
2import torch
3model_name = 'csdc-atl/doc2query'
4tokenizer = AutoTokenizer.from_pretrained(model_name)
5model = AutoModelForSeq2SeqLM.from_pretrained(model_name)
6def create_queries(history, next_question):
7 inputs_ids = []
8 for line in history:
9 inputs_ids.extend([32127]+tokenizer.encode(line[0], add_special_tokens=False)+[32126]+tokenizer.encode(line[1], add_special_tokens=False))
10 inputs_ids.extend([32127]+tokenizer.encode(next_question, add_special_tokens=False))
11 inputs_ids = inputs_ids + [1]
12 inputs_ids = torch.Tensor([inputs_ids]).long()
13 with torch.no_grad():
14 sampling_outputs = model.generate(
15 input_ids=inputs_ids,
16 max_length=512,
17 do_sample=True,
18 top_p=0.95,
19 top_k=10
20 )
21 print("\nSampling Outputs:")
22 for i in range(len(sampling_outputs)):
23 rewrite_question = tokenizer.decode(sampling_outputs[i], skip_special_tokens=True)
24 print(f'{i + 1}: {rewrite_question}')
25history = [['loghub是什么', 'AWS 上的loghub解决方案可帮助组织在单个控制面板上收集、分析和显示 Amazon CloudWatch Logs。该解决方案可整合、管理和分析来自各种来源的日志文件,例如访问、配置更改和计费事件的审计日志。您也可以从多个账户和 AWS 区域收集 Amazon CloudWatch Logs。']]
26next_question = '它的优点是什么?'
27create_queries(history, next_question)
28# 1: loghub解决方案的优点是什么?