Views
No views yet
1from datasets import load_dataset
2from transformers import GPT2Tokenizer, GPT2LMHeadModel, Trainer, TrainingArguments, DataCollatorForLanguageModeling
3from transformers import AutoTokenizer, DataCollatorWithPadding
4from transformers import AutoTokenizer, AutoModelForCausalLM
5import math
6from transformers import LogitsProcessorList, LogitsProcessor
7import torch
8
9
10# 加载 GPT-2 分词器
11tokenizer = AutoTokenizer.from_pretrained("dnagpt/gene_eng_gpt2_summary")
12tokenizer.pad_token = tokenizer.eos_token # 设置填充标记为 EOS 标记
13
14# 6. 加载 GPT-2 模型
15model = GPT2LMHeadModel.from_pretrained("dnagpt/gene_eng_gpt2_summary")
16model.config.pad_token_id = model.config.eos_token_id
17
18def classify_sequence(sequence):
19 # 定义字符集(所有字符都假设为大写)
20 dna_chars = set('ACGT')
21 protein_chars = set('ACDEFGHIKLMNPQRSTVWY')
22 english_chars = set('ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789 ,.!?:;-"\'()')
23
24 # 去除空格并检查长度
25 sequence = sequence.strip() #
26
27 # 检查是否为DNA序列
28 if all(c in dna_chars for c in sequence):
29 return "DNA"
30
31 # 检查是否为蛋白质序列
32 if all(c in protein_chars for c in sequence):
33 return "Protein"
34
35 # 检查是否为英文文本(允许大小写字母、数字及常见标点符号)
36 if all(c in english_chars for c in sequence):
37 return "English"
38
39 # 如果不符合上述任何条件,则无法明确分类
40 return "Unknown"
41
42#获得DNA和英文词表 只要长度2个及以上的词
43word_dict = tokenizer.get_vocab()
44
45DNA_token_list = []
46
47for word in word_dict:
48 word_type = classify_sequence(word)
49 if "DNA"==word_type:
50 DNA_token_list.append(word)
51
52
53class DNAOnlyLogitsProcessor(LogitsProcessor):
54 def __init__(self, allowed_tokens, tokenizer):
55 self.allowed_token_ids = tokenizer.convert_tokens_to_ids(allowed_tokens)
56
57 def __call__(self, input_ids, scores):
58 # 创建掩码,将不允许的 token 的分数设为 -inf
59 mask = torch.full_like(scores, float("-inf"))
60 mask[:, self.allowed_token_ids] = 0
61 scores += mask
62 return scores
63
64def get_summary_with_constraints(text, DNA_token_list):
65 # 确保输入文本的预处理
66 text = text.strip() + " TL;DR:"
67
68 # 对输入文本进行编码
69 encoded_input = tokenizer(
70 text,
71 return_tensors="pt",
72 truncation=True,
73 max_length=256, # 输入文本的最大长度
74 )
75
76 # 创建 DNA 限制的 LogitsProcessor
77 logits_processor = LogitsProcessorList([
78 DNAOnlyLogitsProcessor(DNA_token_list, tokenizer)
79 ])
80
81 # 使用 max_new_tokens 控制生成长度
82 output = model.generate(
83 input_ids=encoded_input["input_ids"],
84 attention_mask=encoded_input["attention_mask"],
85 max_new_tokens=16, # 控制生成的新增文本长度
86 num_beams=5, # 控制生成文本的多样性
87 logits_processor=logits_processor,
88 no_repeat_ngram_size=3, # 避免生成重复内容
89 early_stopping=True, # 提前终止生成
90 )
91
92 # 对生成的输出进行解码
93 generated_text = tokenizer.decode(output[0], skip_special_tokens=True)
94
95 # 提取生成的摘要部分
96 summary = generated_text[len(text)+len(encoded_input["input_ids"][0])-1:].strip() #字符长度+多出来的空格-1
97
98 return summary
99
100# 示例用法
101#test_text = "The DNA sequence analysis showed remarkable results."
102test_text = "GTTATAACCTGTGAGAGTATGTTGGCGGTTTGTTGCACCTACCTTTCAAACCTCTTGTTCTTCCTGTGATTTATTTGAGGCACTCAAGTGGACAGAGACCATGAGAAATTTGAGTGGAGGCCATGTCGAAGAGTTTGTCTTGGTGGGTTTCCCTACCACTCCTCCCTTCCAGCTGCTCCTCTTTGTCCTTTTCTTTGCAATTTACCTTCTGACATTGTTGGAGAATGCACTCATTGTCTTCACAATATGGCTCACTCCAAGCCTTCATCGCCCCATGTACTTTTTCCTTGGCCATCTTTCTTTCCTGGAGCTTTGGTACATCAACGTCACCATTCCTCAGCTCTTGGCAGCCTTTCTTACCCAGGATAGTAGAGTCTCCTATGTAGGTTGCATGACCCAACTCTACTTCTTTATTGCCTTAGCCTGTACTGAATGTGTGCTGTTGGCAGTTATGGCCTATGACCGC"
103
104print(get_summary_with_constraints(test_text, DNA_token_list))
105