muellerzr/performance-debugging
0
1# coding=utf-82# Copyright 2021 The HuggingFace Inc. team. All rights reserved.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15 16import evaluate17import torch18from datasets import load_dataset19from torch.optim import AdamW20from torch.utils.data import DataLoader21from transformers import AutoModelForSequenceClassification, AutoTokenizer, get_linear_schedule_with_warmup22 23from accelerate import Accelerator, DistributedType24from accelerate.utils import set_seed25 26import transformers27 28transformers.logging.set_verbosity_error()29 30 31 32def get_dataloaders(batch_size: int = 16):33 """34 Creates a set of `DataLoader`s for the `glue` dataset,35 using "bert-base-cased" as the tokenizer.36 37 Args:38 accelerator (`Accelerator`):39 An `Accelerator` object40 batch_size (`int`, *optional*):41 The batch size for the train and validation DataLoaders.42 """43 tokenizer = AutoTokenizer.from_pretrained("bert-base-cased")44 datasets = load_dataset("glue", "mrpc")45 46 def tokenize_function(examples):47 outputs = tokenizer(examples["sentence1"], examples["sentence2"], truncation=True, max_length=None)48 return outputs49 50 tokenized_datasets = datasets.map(51 tokenize_function,52 batched=True,53 remove_columns=["idx", "sentence1", "sentence2"],54 )55 tokenized_datasets = tokenized_datasets.rename_column("label", "labels")56 57 def collate_fn(examples):58 return tokenizer.pad(59 examples,60 padding="longest",61 max_length=None,62 pad_to_multiple_of=8,63 return_tensors="pt",64 )65 66 train_dataloader = DataLoader(67 tokenized_datasets["train"], shuffle=True, collate_fn=collate_fn, batch_size=batch_size, drop_last=True68 )69 eval_dataloader = DataLoader(70 tokenized_datasets["validation"],71 shuffle=False,72 collate_fn=collate_fn,73 batch_size=32,74 drop_last=False,75 )76 77 return train_dataloader, eval_dataloader78 79 80def training_function():81 config = {"lr": 2e-5, "num_epochs": 3, "seed": 42}82 seed = int(config["seed"])83 batch_size = 3284 config["batch_size"] = batch_size85 metric = evaluate.load("glue", "mrpc")86 87 set_seed(seed, device_specific=False)88 train_dataloader, eval_dataloader = get_dataloaders(batch_size)89 model = AutoModelForSequenceClassification.from_pretrained("bert-base-cased", return_dict=True)90 model.cuda()91 92 optimizer = AdamW(params=model.parameters(), lr=config["lr"])93 lr_scheduler = get_linear_schedule_with_warmup(94 optimizer=optimizer,95 num_warmup_steps=0,96 num_training_steps=(len(train_dataloader) * config["num_epochs"]),97 )98 99 current_step = 0100 for epoch in range(config["num_epochs"]):101 model.train()102 total_loss = 0103 for _, batch in enumerate(train_dataloader):104 batch = batch.to("cuda")105 outputs = model(**batch)106 loss = outputs.loss107 total_loss += loss.detach().cpu().float()108 current_step += 1109 loss.backward()110 optimizer.step()111 lr_scheduler.step()112 optimizer.zero_grad()113 114 model.eval()115 for step, batch in enumerate(eval_dataloader):116 # We could avoid this line since we set the accelerator with `device_placement=True`.117 batch = batch.to("cuda")118 with torch.no_grad():119 outputs = model(**batch)120 predictions = outputs.logits.argmax(dim=-1)121 metric.add_batch(122 predictions=predictions,123 references=batch["labels"],124 )125 126 eval_metric = metric.compute()127 128 # Use accelerator.print to print only on the main process.129 print(f"epoch {epoch}:", eval_metric)130 print("train_loss: ", total_loss.item() / len(train_dataloader))131 132 133def main():134 training_function()135 136 137if __name__ == "__main__":138 main()139 