Views
No views yet
pip install datasets ai2-olmo| Histone Mark | Accuracy | F1 Score | Matthews Correlation | Precision | Recall |
|---|---|---|---|---|---|
| H3 | 0.8778 | 0.8775 | 0.7558 | 0.8787 | 0.8771 |
| H4 | 0.8919 | 0.8897 | 0.7808 | 0.8934 | 0.8873 |
| H3K9ac | 0.8543 | 0.8531 | 0.7069 | 0.8522 | 0.8548 |
| H3K14ac | 0.9540 | 0.9532 | 0.9065 | 0.9526 | 0.9538 |
| H4ac | 0.9396 | 0.9392 | 0.8785 | 0.9390 | 0.9395 |
| H3K4me1 | 0.7926 | 0.7912 | 0.5825 | 0.7909 | 0.7916 |
| H3K4me2 | 0.8664 | 0.8623 | 0.7248 | 0.8611 | 0.8638 |
| H3K4me3 | 0.9046 | 0.9041 | 0.8084 | 0.9049 | 0.9034 |
| H3K36me3 | 0.8581 | 0.8571 | 0.7142 | 0.8571 | 0.8570 |
| H3K79me3 | 0.8651 | 0.8641 | 0.7291 | 0.8660 | 0.8631 |
zl6222@ic.ac.uk1import argparse
2import os
3import torch
4import re
5import numpy as np
6import sklearn
7from tqdm import tqdm
8from datasets import load_dataset, DatasetDict
9from transformers import AutoModelForCausalLM, AutoTokenizer
10
11def parse_args():
12 parser = argparse.ArgumentParser(description="Run DNA task inference with a specified model and tokenizer.")
13 parser.add_argument(
14 "--model_tokenizer_path",
15 type=str,
16 default="zehui127/Omni-DNA-Multitask", # Set default value
17 help="Path to the pretrained model and tokenizer. Default: zehui127/Omni-DNA-Multitask"
18 )
19 return parser.parse_args()
20
21def load_model_and_tokenizer(model_tokenizer_path):
22 tokenizer = AutoTokenizer.from_pretrained(model_tokenizer_path)
23 model = AutoModelForCausalLM.from_pretrained(model_tokenizer_path).to('cuda')
24 return model, tokenizer
25
26def generate(message, task_type, model, tokenizer, sample_num=1):
27 tokenized_message = tokenizer([message], return_tensors='pt', return_token_type_ids=False, add_special_tokens=True).to('cuda')
28 response = model.generate(**tokenized_message, max_new_tokens=sample_num, do_sample=False)
29 reply = tokenizer.batch_decode(response, skip_special_tokens=False)[0].replace(" ", "")
30 return extract_label(reply, task_type)
31
32def extract_label(message, task_type):
33 task_type = '[MASK]'
34 answer = message.split(task_type)[1]
35 match = re.search(r'\d+', answer)
36 return match.group() if match else None
37
38def load_and_format_dataset():
39 raw_dataset = load_dataset("zehui127/Omni-DNA-dataset-nt-downstream-multitask")
40 dataset = raw_dataset['test']
41
42 def formatting_prompts_func(example):
43 output_texts = [f"{instr}[MASK]" for instr in example['instruction']]
44 labels = [output[-1] for output in example['output']]
45 task_types = example['task']
46 return {'formatted_text': output_texts, 'label': labels, 'task_type': task_types}
47
48 formatted_dataset = dataset.map(formatting_prompts_func, batched=True, remove_columns=dataset.column_names, desc="Formatting dataset")
49 return formatted_dataset
50
51def group_by_task_type(dataset):
52 task_types = set(dataset['task_type'])
53 task_datasets = DatasetDict()
54
55 for task_type in task_types:
56 filtered_dataset = dataset.filter(lambda x: x['task_type'] == task_type, num_proc=1, desc=f"Filtering {task_type} examples")
57 if len(filtered_dataset) > 0:
58 task_datasets[task_type] = filtered_dataset
59 print(f"\nTask type '{task_type}': {len(filtered_dataset)} examples")
60
61 return task_datasets
62
63def calculate_metrics(predictions, labels):
64 valid_mask = labels != -100
65 valid_predictions = predictions[valid_mask]
66 valid_labels = labels[valid_mask]
67 return {
68 "accuracy": sklearn.metrics.accuracy_score(valid_labels, valid_predictions),
69 "f1": sklearn.metrics.f1_score(valid_labels, valid_predictions, average="macro", zero_division=0),
70 "matthews_correlation": sklearn.metrics.matthews_corrcoef(valid_labels, valid_predictions),
71 "precision": sklearn.metrics.precision_score(valid_labels, valid_predictions, average="macro", zero_division=0),
72 "recall": sklearn.metrics.recall_score(valid_labels, valid_predictions, average="macro", zero_division=0),
73 }
74
75def inference(dataset, model, tokenizer):
76 predictions, labels = [], []
77
78 for element in tqdm(dataset):
79 prediction = generate(element['formatted_text'], element['task_type'], model, tokenizer)
80 sample_num = 2
81 while prediction is None:
82 prediction = generate(element['formatted_text'], element['task_type'], model, tokenizer, sample_num)
83 sample_num += 1
84 if sample_num >= 20:
85 prediction = '0'
86 print("Warning: No valid result")
87 break
88 predictions.append(int(str(prediction)[0]))
89 labels.append(int(element['label']))
90
91 return calculate_metrics(np.array(predictions), np.array(labels))
92
93def main():
94 args = parse_args()
95 model, tokenizer = load_model_and_tokenizer(args.model_tokenizer_path)
96 formatted_dataset = load_and_format_dataset()
97 task_specific_datasets = group_by_task_type(formatted_dataset)
98
99 tasks = ['H3', 'H4', 'H3K9ac', 'H3K14ac', 'H4ac', 'H3K4me1', 'H3K4me2', 'H3K4me3', 'H3K36me3', 'H3K79me3']
100
101 for task in tasks:
102 print(f"==========={task}=========")
103 dataset_test = task_specific_datasets.get(task, None)
104 if dataset_test:
105 print(inference(dataset_test, model, tokenizer))
106 else:
107 print(f"No data for task {task}")
108
109if __name__ == "__main__":
110 main()
111