Views
No views yet
1from transformers import AutoModelForCausalLM, AutoTokenizer
2from peft import PeftModel
3import torch
4
5base_model_name = "baichuan-inc/Baichuan2-7B-Base"
6base_model = AutoModelForCausalLM.from_pretrained(
7 base_model_name,
8 torch_dtype="auto",
9 device_map="cuda:0"
10)
11lora_model_name = "yanghh7/Alirector-baichuan2-7b-lora"
12model = PeftModel.from_pretrained(
13 base_model,
14 lora_model_name
15)
16tokenizer = AutoTokenizer.from_pretrained(base_model_name, trust_remote_code=True)
17template = "对输入句子进行语法纠错,并输出正确的句子。\nInput:{source}\nOutput:"
18
19while True:
20 source = input("输入句子:")
21
22 prompt = template.format(source=source)
23 model_inputs = tokenizer(
24 [prompt],
25 return_tensors="pt",
26 ).to(model.device)
27
28 with torch.no_grad():
29 output = model.generate(**model_inputs)
30
31 response = tokenizer.batch_decode(output, skip_special_tokens=True)[0]
32 print(response.split('\nOutput:')[-1])@inproceedings{yang-quan-2024-alirector,
title = "Alirector: Alignment-Enhanced {C}hinese Grammatical Error Corrector",
author = "Yang, Haihui and Quan, Xiaojun",
booktitle = "Findings of the Association for Computational Linguistics: ACL 2024",
year = "2024",
}