Views
No views yet
1from transformers import AutoTokenizer, AutoConfig,AutoModel
2from transformers import DataCollatorForLanguageModeling
3from transformers import Trainer, TrainingArguments
4from transformers import AutoConfig, AutoModelForCausalLM,LlamaForCausalLM,LlamaTokenizer
5from tokenizers import Tokenizer
6from datasets import load_dataset
7
8dna_ft_dataset = load_dataset('dnagpt/dna_multi_task_finetune')
9
10data = dna_ft_dataset["train"].train_test_split(train_size=0.1, seed=42)
11
12tokenizer = LlamaTokenizer.from_pretrained("dnagpt/llama-dna-sft")
13tokenizer.pad_token = tokenizer.eos_token
14
15model = LlamaForCausalLM.from_pretrained("dnagpt/llama-dna-sft") #sft
16
17#构建提示词
18def format_input(entry):
19 instruction_text = (
20 f"Below is an instruction that describes a task. "
21 f"Write a response that appropriately completes the request."
22 f"\n\n### Instruction:\n{entry['instruction']}"
23 )
24
25 input_text = f"\n\n### Input:\n{entry['input']}" if entry["input"] else ""
26
27 return instruction_text + input_text + "\n\n### Response:\n"
28
29#构建提示词
30def build_prompt(entry):
31
32 input_data = format_input(entry)
33
34 desired_response = entry['output']
35
36 return input_data + desired_response
37
38example = data["test"][0]
39
40prompt = build_prompt(example)
41
42def inference(text, model, tokenizer, max_input_tokens=1000, max_output_tokens=1000):
43 # Tokenize
44 input_ids = tokenizer.encode(
45 text,
46 return_tensors="pt",
47 truncation=True,
48 max_length=max_input_tokens
49 # return_attention_mask=True,
50 )
51
52 # Generate
53 device = model.device
54 generated_tokens_with_prompt = model.generate(
55 input_ids=input_ids.to(device),
56 #max_length=max_output_tokens,
57 max_new_tokens=8,
58 temperature=0.01 # 控制生成的多样性
59 )
60
61 # Decode
62 generated_text_with_prompt = tokenizer.decode(generated_tokens_with_prompt[0], skip_special_tokens=True)
63 generated_text_answer = generated_text_with_prompt[len(text):]
64
65
66 return generated_text_answer
67
68
69input_text = format_input(data["test"][0])
70
71print("input (test):", input_text)
72
73print("real answer:", data["test"][0]["output"])
74
75print("--------------------------\n")
76
77print("model's answer: \n")
78print(inference(input_text, model, tokenizer))
79
80
81
82
83
84