Views
No views yet
1from transformers import AutoModel, AutoTokenizer, HfArgumentParser, TrainingArguments, Trainer
2from transformers.data.data_collator import DataCollatorWithPadding
3from transformers.trainer_pt_utils import get_parameter_names
4from transformers.pytorch_utils import ALL_LAYERNORM_LAYERS
5from datasets import load_dataset
6import functools
7import numpy as np
8from sklearn.metrics import accuracy_score, matthews_corrcoef
9import sys
10import torch
11import logging
12import datasets
13import transformers
14+ import habana_frameworks.torch
15+ from optimum.habana import GaudiConfig, GaudiTrainer, GaudiTrainingArguments
16
17
18logging.basicConfig(level=logging.INFO)
19logger = logging.getLogger(__name__)
20
21def create_optimizer(opt_model, lr_ratio=0.1):
22 head_names = []
23 for n, p in opt_model.named_parameters():
24 if "classifier" in n:
25 head_names.append(n)
26 else:
27 p.requires_grad = False
28 # turn a list of tuple to 2 lists
29 for n, p in opt_model.named_parameters():
30 if n in head_names:
31 assert p.requires_grad
32 backbone_names = []
33 for n, p in opt_model.named_parameters():
34 if n not in head_names and p.requires_grad:
35 backbone_names.append(n)
36 # for weight_decay policy, see
37 # https://github.com/huggingface/transformers/blob/50573c648ae953dcc1b94d663651f07fb02268f4/src/transformers/trainer.py#L947
38 decay_parameters = get_parameter_names(opt_model, ALL_LAYERNORM_LAYERS) # forbidden layer norm
39 decay_parameters = [name for name in decay_parameters if "bias" not in name]
40 # training_args.learning_rate
41 head_decay_parameters = [name for name in head_names if name in decay_parameters]
42 head_not_decay_parameters = [name for name in head_names if name not in decay_parameters]
43 # training_args.learning_rate * model_config.lr_ratio
44 backbone_decay_parameters = [name for name in backbone_names if name in decay_parameters]
45 backbone_not_decay_parameters = [name for name in backbone_names if name not in decay_parameters]
46 optimizer_grouped_parameters = [
47 {
48 "params": [p for n, p in opt_model.named_parameters() if (n in head_decay_parameters and p.requires_grad)],
49 "weight_decay": training_args.weight_decay,
50 "lr": training_args.learning_rate
51 },
52 {
53 "params": [p for n, p in opt_model.named_parameters() if (n in backbone_decay_parameters and p.requires_grad)],
54 "weight_decay": training_args.weight_decay,
55 "lr": training_args.learning_rate * lr_ratio
56 },
57 {
58 "params": [p for n, p in opt_model.named_parameters() if (n in head_not_decay_parameters and p.requires_grad)],
59 "weight_decay": 0.0,
60 "lr": training_args.learning_rate
61 },
62 {
63 "params": [p for n, p in opt_model.named_parameters() if (n in backbone_not_decay_parameters and p.requires_grad)],
64 "weight_decay": 0.0,
65 "lr": training_args.learning_rate * lr_ratio
66 },
67 ]
68 - optimizer_cls, optimizer_kwargs = Trainer.get_optimizer_cls_and_kwargs(training_args)
69 + optimizer_cls, optimizer_kwargs = GaudiTrainer.get_optimizer_cls_and_kwargs(training_args)
70 optimizer = optimizer_cls(optimizer_grouped_parameters, **optimizer_kwargs)
71
72 return optimizer
73
74def create_scheduler(training_args, optimizer):
75 from transformers.optimization import get_scheduler
76 return get_scheduler(
77 training_args.lr_scheduler_type,
78 optimizer=optimizer if optimizer is None else optimizer,
79 num_warmup_steps=training_args.get_warmup_steps(training_args.max_steps),
80 num_training_steps=training_args.max_steps,
81 )
82
83def compute_metrics(eval_preds):
84 probs, labels = eval_preds
85 preds = np.argmax(probs, axis=-1)
86 result = {"accuracy": accuracy_score(labels, preds), "mcc": matthews_corrcoef(labels, preds)}
87 return result
88
89def preprocess_logits_for_metrics(logits, labels):
90 return torch.softmax(logits, dim=-1)
91
92
93if __name__ == "__main__":
94 - device = torch.device("cpu")
95 + device = torch.device("hpu")
96 raw_dataset = load_dataset("Jiqing/ProtST-BinaryLocalization")
97 model = AutoModel.from_pretrained("Jiqing/protst-esm1b-for-sequential-classification", trust_remote_code=True, torch_dtype=torch.bfloat16).to(device)
98 tokenizer = AutoTokenizer.from_pretrained("facebook/esm1b_t33_650M_UR50S")
99
100 output_dir = "/home/jiqingfe/protst/protst_2/ProtST-HuggingFace/output_dir/ProtSTModel/default/ESM-1b_PubMedBERT-abs/240123_015856"
101 training_args = {'output_dir': output_dir, 'overwrite_output_dir': True, 'do_train': True, 'per_device_train_batch_size': 32, 'gradient_accumulation_steps': 1, \
102 'learning_rate': 5e-05, 'weight_decay': 0, 'num_train_epochs': 100, 'max_steps': -1, 'lr_scheduler_type': 'constant', 'do_eval': True, \
103 'evaluation_strategy': 'epoch', 'per_device_eval_batch_size': 32, 'logging_strategy': 'epoch', 'save_strategy': 'epoch', 'save_steps': 820, \
104 'dataloader_num_workers': 0, 'run_name': 'downstream_esm1b_localization_fix', 'optim': 'adamw_torch', 'resume_from_checkpoint': False, \
105 - 'label_names': ['labels'], 'load_best_model_at_end': True, 'metric_for_best_model': 'accuracy', 'bf16': True, "save_total_limit": 3}
106 + 'label_names': ['labels'], 'load_best_model_at_end': True, 'metric_for_best_model': 'accuracy', 'bf16': True, "save_total_limit": 3, "use_habana":True, "use_lazy_mode": True, "use_hpu_graphs_for_inference": True}
107 - training_args = HfArgumentParser(TrainingArguments).parse_dict(training_args, allow_extra_keys=False)[0]
108 + training_args = HfArgumentParser(GaudiTrainingArguments).parse_dict(training_args, allow_extra_keys=False)[0]
109
110 def tokenize_protein(example, tokenizer=None):
111 protein_seq = example["prot_seq"]
112 protein_seq_str = tokenizer(protein_seq, add_special_tokens=True)
113 example["input_ids"] = protein_seq_str["input_ids"]
114 example["attention_mask"] = protein_seq_str["attention_mask"]
115 example["labels"] = example["localization"]
116
117 return example
118
119 func_tokenize_protein = functools.partial(tokenize_protein, tokenizer=tokenizer)
120
121 for split in ["train", "validation", "test"]:
122 raw_dataset[split] = raw_dataset[split].map(func_tokenize_protein, batched=False, remove_columns=["Unnamed: 0", "prot_seq", "localization"])
123
124 - data_collator = DataCollatorWithPadding(tokenizer=tokenizer)
125 + data_collator = DataCollatorWithPadding(tokenizer=tokenizer, padding="max_length", max_length=1024)
126
127 transformers.utils.logging.set_verbosity_info()
128 log_level = training_args.get_process_log_level()
129 logger.setLevel(log_level)
130
131 optimizer = create_optimizer(model)
132 scheduler = create_scheduler(training_args, optimizer)
133
134 + gaudi_config = GaudiConfig()
135 + gaudi_config.use_fused_adam = True
136 + gaudi_config.use_fused_clip_norm =True
137
138
139 # build trainer
140 - trainer = Trainer(
141 + trainer = GaudiTrainer(
142 model=model,
143 + gaudi_config=gaudi_config,
144 args=training_args,
145 train_dataset=raw_dataset["train"],
146 eval_dataset=raw_dataset["validation"],
147 data_collator=data_collator,
148 optimizers=(optimizer, scheduler),
149 compute_metrics=compute_metrics,
150 preprocess_logits_for_metrics=preprocess_logits_for_metrics,
151 )
152
153 train_result = trainer.train()
154
155 trainer.save_model()
156 # Saves the tokenizer too for easy upload
157 tokenizer.save_pretrained(training_args.output_dir)
158
159 metrics = train_result.metrics
160 metrics["train_samples"] = len(raw_dataset["train"])
161
162 trainer.log_metrics("train", metrics)
163 trainer.save_metrics("train", metrics)
164 trainer.save_state()
165
166 metric = trainer.evaluate(raw_dataset["test"], metric_key_prefix="test")
167 print("test metric: ", metric)
168
169 metric = trainer.evaluate(raw_dataset["validation"], metric_key_prefix="valid")
170 print("valid metric: ", metric)