KKvision/PSTU_AI_sem-1_bert-language-classifier
0
1import os2import json3import numpy as np4import torch5from torch.utils.data import DataLoader, Dataset6from torch.optim import AdamW7from transformers import BertTokenizer, BertForSequenceClassification, get_linear_schedule_with_warmup8from datasets import load_dataset9from sklearn.metrics import classification_report, accuracy_score10from tqdm import tqdm11 12MODEL_NAME = "bert-base-multilingual-cased"13DATASET_NAME = "papluca/language-identification"14MAX_LEN = 12815BATCH_SIZE = 1616EPOCHS = 317LR = 2e-518SAVE_DIR = "saved_model"19DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")20 21torch.manual_seed(42)22np.random.seed(42)23 24 25class LanguageDataset(Dataset):26 def __init__(self, data, tokenizer):27 self.texts = data["text"]28 self.labels = data["labels"]29 self.tokenizer = tokenizer30 31 def __len__(self):32 return len(self.texts)33 34 def __getitem__(self, idx):35 enc = self.tokenizer(36 self.texts[idx],37 max_length=MAX_LEN,38 padding="max_length",39 truncation=True,40 return_tensors="pt",41 )42 return {43 "input_ids": enc["input_ids"].squeeze(0),44 "attention_mask": enc["attention_mask"].squeeze(0),45 "labels": torch.tensor(self.labels[idx], dtype=torch.long),46 }47 48 49def train_epoch(model, loader, optimizer, scheduler):50 model.train()51 total_loss, correct, total = 0, 0, 052 for batch in tqdm(loader, desc="train", leave=False):53 ids = batch["input_ids"].to(DEVICE)54 mask = batch["attention_mask"].to(DEVICE)55 y = batch["labels"].to(DEVICE)56 optimizer.zero_grad()57 out = model(input_ids=ids, attention_mask=mask, labels=y)58 out.loss.backward()59 torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)60 optimizer.step()61 scheduler.step()62 total_loss += out.loss.item()63 correct += (out.logits.argmax(-1) == y).sum().item()64 total += y.size(0)65 return total_loss / len(loader), correct / total66 67 68def evaluate(model, loader, desc="eval"):69 model.eval()70 total_loss, preds, targets = 0, [], []71 with torch.no_grad():72 for batch in tqdm(loader, desc=desc, leave=False):73 ids = batch["input_ids"].to(DEVICE)74 mask = batch["attention_mask"].to(DEVICE)75 y = batch["labels"].to(DEVICE)76 out = model(input_ids=ids, attention_mask=mask, labels=y)77 total_loss += out.loss.item()78 preds.extend(out.logits.argmax(-1).cpu().numpy())79 targets.extend(y.cpu().numpy())80 return total_loss / len(loader), accuracy_score(targets, preds), preds, targets81 82 83def main():84 print(f"device: {DEVICE}\n")85 86 # датасет87 dataset = load_dataset(DATASET_NAME)88 labels = sorted(set(dataset["train"]["labels"]))89 label2id = {l: i for i, l in enumerate(labels)}90 id2label = {i: l for i, l in enumerate(labels)}91 92 dataset = dataset.map(lambda x: {"labels": label2id[x["labels"]]})93 print(f"languages ({len(labels)}): {', '.join(labels)}")94 95 # модель96 tokenizer = BertTokenizer.from_pretrained(MODEL_NAME)97 model = BertForSequenceClassification.from_pretrained(98 MODEL_NAME, num_labels=len(labels), id2label=id2label, label2id=label2id99 ).to(DEVICE)100 101 train_loader = DataLoader(LanguageDataset(dataset["train"], tokenizer), BATCH_SIZE, shuffle=True, num_workers=2)102 val_loader = DataLoader(LanguageDataset(dataset["validation"], tokenizer), BATCH_SIZE, shuffle=False, num_workers=2)103 test_loader = DataLoader(LanguageDataset(dataset["test"], tokenizer), BATCH_SIZE, shuffle=False, num_workers=2)104 105 total_steps = len(train_loader) * EPOCHS106 optimizer = AdamW(model.parameters(), lr=LR, weight_decay=0.01)107 scheduler = get_linear_schedule_with_warmup(optimizer, int(total_steps * 0.1), total_steps)108 109 # обучение110 history = {"train_loss": [], "train_acc": [], "val_loss": [], "val_acc": []}111 best_val_acc = 0.0112 113 for epoch in range(1, EPOCHS + 1):114 tr_loss, tr_acc = train_epoch(model, train_loader, optimizer, scheduler)115 vl_loss, vl_acc, _, _ = evaluate(model, val_loader, "val")116 117 history["train_loss"].append(tr_loss)118 history["train_acc"].append(tr_acc)119 history["val_loss"].append(vl_loss)120 history["val_acc"].append(vl_acc)121 122 print(f"epoch {epoch}/{EPOCHS} train loss {tr_loss:.4f} acc {tr_acc:.4f} val loss {vl_loss:.4f} acc {vl_acc:.4f}")123 124 if vl_acc > best_val_acc:125 best_val_acc = vl_acc126 os.makedirs(SAVE_DIR, exist_ok=True)127 model.save_pretrained(SAVE_DIR)128 tokenizer.save_pretrained(SAVE_DIR)129 with open(f"{SAVE_DIR}/labels.json", "w") as f:130 json.dump({"id2label": id2label, "label2id": label2id}, f)131 print(f" saved (val_acc={vl_acc:.4f})")132 133 # тест134 model = BertForSequenceClassification.from_pretrained(SAVE_DIR).to(DEVICE)135 test_loss, test_acc, preds, targets = evaluate(model, test_loader, "test")136 print(f"\ntest acc: {test_acc:.4f}\n")137 print(classification_report(targets, preds, target_names=labels, digits=3))138 139 with open(f"{SAVE_DIR}/training_history.json", "w") as f:140 json.dump(history, f, indent=2)141 142 143if __name__ == "__main__":144 main()145 