Loading config.pydeleted 100644 → 0 +0 −30 Original line number Diff line number Diff line from dataclasses import dataclass, field from typing import List, Optional, Union @dataclass class TrainConfig: batch_size: int = 2 save_dir: str = "saved_models/" bert_model: str = "sentence-transformers/all-distilroberta-v1" multilabel: bool = True regression: bool = False five_classes: bool = False patience: int = 20 seed: int = 2 encoder_hidden_dim: int = 300 encoder_num_layers: int = 1 dropout: float = 0.5 num_classes: int = 8 attention_type: str = "hierarchical" pooling: str = "mean" binary_only: bool = True bidirectional: bool = True regularization_loss: bool = False loss_l: float = 0.1 lr: float = 3e-5 num_iters: int = 100 encoder_layers_to_freeze: Optional[Union[str, List[Union[str, int]]]] = field( default_factory=lambda: ["embeddings", 0, 1, 2, 3, 4] ) save_every_epoch: bool = False config.toml 0 → 100644 +22 −0 Original line number Diff line number Diff line save_dir = "saved_models" batch_size = 2 seed = 2 five_classes = false regularization_loss = false loss_l = 0.1 lr = 3e-5 num_iters = 100 encoder_layers_to_freeze = [ ["embeddings"], [0, 1, 2, 3, 4] ] [model] bert_model = "sentence-transformers/all-distilroberta-v1" encoder_hidden_dim = 300 encoder_num_layers = 1 dropout = 0.5 num_classes = 8 attention_type = "hierarchical" pooling = "mean" binary_only = true bidirectional = true multilabel = true regression = false No newline at end of file hdsc/model.py +1 −1 Original line number Diff line number Diff line Loading @@ -7,7 +7,7 @@ from accelerate.utils import send_to_device from torch.nn.utils.rnn import pad_sequence from transformers import AutoModel, AutoConfig from hdsc.modules import ClassificationHead, PositionalEncoding, SentenceLevelEncoder, WordLevelEncoder from hdsc.modules import ClassificationHead, SentenceLevelEncoder PAD_ID = 0 Loading train.py +59 −65 Original line number Diff line number Diff line import json import itertools from pathlib import Path from typing import Dict, List, Optional import datasets import tomli import torch import torch.nn as nn from accelerate import Accelerator Loading @@ -14,12 +16,10 @@ from torchinfo import summary from torchmetrics.functional import accuracy, f1_score, mean_absolute_error, mean_squared_error from transformers import AutoTokenizer, get_linear_schedule_with_warmup from config import TrainConfig from hdsc.losses import MultiLabelLoss from hdsc.model import PHQTotalMulticlassAttentionModelBERT from hdsc.utils import save_model train_config = TrainConfig() logger = get_logger(__name__) logger.setLevel("INFO") Loading @@ -38,6 +38,7 @@ def collate_fn(batch): def training_step( config: Dict, model: nn.Module, inputs: dict, loss_fn: nn.Module, Loading Loading @@ -70,11 +71,11 @@ def training_step( phq_score = torch.sum(inputs["labels"], dim=1) if train_config.multilabel: if config["model"]["multilabel"]: labels = inputs["labels"] elif train_config.regression: elif config["model"]["regression"]: labels = phq_score elif train_config.five_classes: elif config["five_classes"]: labels = torch.div(phq_score, 5, rounding_mode="floor") labels = labels.to(torch.long) else: Loading @@ -83,7 +84,7 @@ def training_step( with accelerator.autocast(): pred_binary = model(inputs) if train_config.multilabel: if config["model"]["multilabel"]: loss, loss_1, loss_2 = loss_fn( pred_binary, labels, Loading @@ -97,13 +98,13 @@ def training_step( scheduler.step() # Gather the outputs from all GPUs cat_fn = torch.vstack if train_config.multilabel else torch.hstack cat_fn = torch.vstack if config["model"]["multilabel"] else torch.hstack gathered_preds = accelerator.gather(pred_binary).detach() gathered_labels = accelerator.gather(labels).detach() gathered_reg_labels = accelerator.gather(phq_score).detach() # Recalculate the loss on all inputs if train_config.multilabel: if config["model"]["multilabel"]: loss, loss_1, loss_2 = loss_fn(gathered_preds, gathered_labels, gathered_reg_labels) else: loss = loss_fn(gathered_preds, gathered_labels) Loading @@ -123,6 +124,7 @@ def training_step( def evaluate( config: Dict, model: nn.Module, val_dataloader: torch.utils.data.DataLoader, loss_fn: nn.Module, Loading Loading @@ -154,17 +156,17 @@ def evaluate( all_reg_labels: List[torch.Tensor] = [] n_steps = len(val_dataloader) cat_fn = torch.vstack if train_config.multilabel else torch.hstack cat_fn = torch.vstack if config["model"]["multilabel"] else torch.hstack model.eval() with torch.no_grad(): for batch in val_dataloader: phq_score = torch.sum(batch["labels"], dim=1) if train_config.multilabel: if config["model"]["multilabel"]: labels = batch["labels"] elif train_config.regression: elif config["model"]["regression"]: labels = phq_score elif train_config.five_classes: elif config["five_classes"]: labels = torch.div(phq_score, 5, rounding_mode="floor") labels = labels.to(torch.long) else: Loading @@ -177,7 +179,7 @@ def evaluate( gathered_reg_labels = accelerator.gather(phq_score).detach() # Recalculate the loss on all inputs if train_config.multilabel: if config["model"]["multilabel"]: loss, loss_1, loss_2 = loss_fn(gathered_preds, gathered_labels, gathered_reg_labels) else: loss = loss_fn(gathered_preds, gathered_labels) Loading @@ -196,7 +198,7 @@ def evaluate( avg_loss_2 = running_loss_2.item() / n_steps all_preds = torch.vstack(all_preds) if not train_config.multilabel and not train_config.regression: if not config["model"]["multilabel"] and not config["model"]["regression"]: all_preds = all_preds.topk(k=1, dim=1)[1].squeeze(-1) all_labels = cat_fn(all_labels) all_reg_labels = torch.hstack(all_reg_labels) Loading @@ -217,85 +219,77 @@ def main(): accelerator = Accelerator() device = accelerator.device with open("config.toml", "rb") as f: config = tomli.load(f) model_name = "_".join( [ "lstm" if train_config.bert_model == "lstm" else "robert", "multilabel" if train_config.multilabel else "binary", "regression" if train_config.regression else "no-regression", "five_classes" if train_config.five_classes else "", "lstm" if config["model"]["bert_model"] == "lstm" else "robert", "multilabel" if config["model"]["multilabel"] else "binary", "regression" if config["model"]["regression"] else "no-regression", "five_classes" if config["five_classes"] else "", ] ) save_dir = Path(train_config.save_dir) / model_name save_dir = Path(config["save_dir"]) / model_name if not save_dir.exists(): save_dir.mkdir(parents=True) tokenizer = AutoTokenizer.from_pretrained(train_config.bert_model) tokenizer = AutoTokenizer.from_pretrained(config["model"]["bert_model"]) dataset = load_dataset("daic_woz.py", "lines") encoded_dataset = dataset.map( lambda examples: tokenizer(examples["turns"], padding="max_length", truncation=True), load_from_cache_file=False, ) encoded_dataset.set_format(type="torch", columns=["input_ids", "attention_mask", "labels"]) cat_fn = torch.vstack if train_config.multilabel else torch.hstack cat_fn = torch.vstack if config["model"]["multilabel"] else torch.hstack train_dataloader = torch.utils.data.DataLoader( encoded_dataset["train"], batch_size=train_config.batch_size, batch_size=config["batch_size"], collate_fn=collate_fn, shuffle=True, ) validation_dataloader = torch.utils.data.DataLoader( encoded_dataset["validation"], batch_size=train_config.batch_size, batch_size=config["batch_size"], collate_fn=collate_fn, shuffle=False, ) model = PHQTotalMulticlassAttentionModelBERT( train_config.bert_model, train_config.encoder_hidden_dim, train_config.encoder_num_layers, train_config.dropout, train_config.num_classes, train_config.attention_type, train_config.pooling, train_config.binary_only, train_config.bidirectional, train_config.multilabel, train_config.regression, device, ) model = PHQTotalMulticlassAttentionModelBERT(device=device, **config["model"]) model.initialize_encoder_weights() if train_config.multilabel: if config["model"]["multilabel"]: loss_fn = MultiLabelLoss( regularization=train_config.regularization_loss, l=train_config.loss_l, regularization=config["regularization_loss"], l=config["loss_l"], ) elif args.regression: elif config["model"]["regression"]: loss_fn = nn.SmoothL1Loss() else: loss_fn = nn.NLLLoss() optimizer = AdamW(model.parameters(), lr=train_config.lr) total_steps = len(train_dataloader) * train_config.num_iters optimizer = AdamW(model.parameters(), lr=config["lr"]) total_steps = len(train_dataloader) * config["num_iters"] scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=5, num_training_steps=total_steps) model = accelerator.prepare(model) optimizer, train_dataloader, scheduler = accelerator.prepare(optimizer, train_dataloader, scheduler) validation_dataloader = accelerator.prepare(validation_dataloader) if train_config.encoder_layers_to_freeze is not None: if train_config.encoder_layers_to_freeze == "all": if config["encoder_layers_to_freeze"] is not None: if config["encoder_layers_to_freeze"] == "all": model.encoder.requires_grad_(False) else: layers_to_freeze = [str(x) for x in train_config.encoder_layers_to_freeze] layers_to_freeze = [ str(x) for x in list(itertools.chain.from_iterable(config["encoder_layers_to_freeze"])) ] for name, param in model.encoder.named_parameters(): if any([True for layer in layers_to_freeze if layer in name]): param.requires_grad_(False) accelerator.print(train_config) accelerator.print(config) model_stats = summary(model, depth=4, verbose=0) accelerator.print(model_stats) Loading @@ -313,16 +307,16 @@ def main(): global_step = 0 eval_every = len(train_dataloader) with open(save_dir / f"log_{model_name}_{train_config.seed}.tsv", "w", encoding="utf-8") as f: with open(save_dir / f"log_{model_name}_{config['seed']}.tsv", "w", encoding="utf-8") as f: f.write( "epoch\ttrain_loss\ttrain_loss_bin\ttrain_loss_reg\tdev_loss\tdev_loss_bin\tdev_loss_reg\t" + "train_acc\tdev_acc\tf1_micro\tf1_macro\tdev_mae\tdev_mse\n" ) for epoch in range(train_config.num_iters): for epoch in range(config["num_iters"]): running_loss = 0.0 preds = [] for batch in train_dataloader: outputs = training_step(model, batch, loss_fn, optimizer, accelerator, scheduler) outputs = training_step(config, model, batch, loss_fn, optimizer, accelerator, scheduler) running_loss += outputs["loss"] running_loss_bin += outputs["loss_1"] Loading @@ -334,7 +328,7 @@ def main(): global_step += 1 if global_step % eval_every == 0: eval_outputs = evaluate(model, validation_dataloader, loss_fn, accelerator) eval_outputs = evaluate(config, model, validation_dataloader, loss_fn, accelerator) average_train_loss = running_loss.item() / eval_every average_train_loss_bin = running_loss_bin.item() / eval_every Loading @@ -344,7 +338,7 @@ def main(): average_dev_loss_reg = eval_outputs["loss_2"] preds = torch.vstack(preds) if not train_config.multilabel and not train_config.regression: if not config["model"]["multilabel"] and not config["model"]["regression"]: preds = preds.topk(k=1, dim=1)[1].squeeze(-1) dev_preds = eval_outputs["preds"] preds = preds.float() Loading @@ -353,7 +347,7 @@ def main(): dev_labels = eval_outputs["labels"] dev_labels_reg = eval_outputs["reg_labels"] if train_config.multilabel: if config["model"]["multilabel"]: average_train_acc = mean_absolute_error(preds, labels).item() average_dev_acc = mean_absolute_error(dev_preds, dev_labels).item() Loading @@ -371,7 +365,7 @@ def main(): dev_preds >= 1.5, task="multilabel", average="micro", num_labels=train_config.num_classes, num_labels=config["model"]["num_classes"], ).item() mae = mean_absolute_error( Loading @@ -383,7 +377,7 @@ def main(): dev_labels_reg, ).item() f1_mae = f1_micro / mae elif train_config.regression: elif config["model"]["regression"]: preds = preds.squeeze(-1) dev_preds = dev_preds.squeeze(-1) average_train_acc = mean_absolute_error(preds, labels).item() Loading @@ -408,7 +402,7 @@ def main(): dev_labels_reg, ).item() f1_mae = f1_micro / mae elif train_config.five_classes: elif config["five_classes"]: average_train_acc = accuracy(preds.to(torch.long), labels).item() average_dev_acc = accuracy(dev_preds.to(torch.long), dev_labels).item() Loading Loading @@ -443,7 +437,7 @@ def main(): f1_mae = 0.0 f1_samples = 0.0 preds_save_path = save_dir / f"preds_{train_config.seed}_{epoch}.json" preds_save_path = save_dir / f"preds_{config['seed']}_{epoch}.json" preds_to_json = { "pred": dev_preds.tolist(), "true": dev_labels.tolist(), Loading @@ -464,7 +458,7 @@ def main(): dev_labels_reg = [] stats_message = ( f"Epoch [{epoch+1}/{train_config.num_iters}], Step [{global_step}/{train_config.num_iters*len(train_dataloader)}], " f"Epoch [{epoch+1}/{config['num_iters']}], Step [{global_step}/{config['num_iters']*len(train_dataloader)}], " + f"Train Loss: {average_train_loss:.4f}, Dev Loss: {average_dev_loss:.4f}, " + f"Train Loss Binary: {average_train_loss_bin:.4f}, Dev Loss Binary: {average_dev_loss_bin:.4f}, " + f"Train Loss Regression: {average_train_loss_reg:.4f}, Dev Loss Regression: {average_dev_loss_reg:.4f}, " Loading Loading @@ -498,20 +492,20 @@ def main(): f.write("\n") accelerator.wait_for_everyone() if train_config.save_every_epoch: save_path = save_dir / f"model_{train_config.seed}_{epoch}.pt" if config["save_every_epoch"]: save_path = save_dir / f"model_{config['seed']}_{epoch}.pt" save_model(model, save_path, accelerator) else: if train_config.multilabel: save_path = save_dir / f"model_best_loss_{train_config.seed}.pt" if config["model"]["multilabel"]: save_path = save_dir / f"model_best_loss_{config['seed']}.pt" if average_dev_loss < best_loss: accelerator.print( f"Dev loss decreased ({best_loss} -> {average_dev_loss}). Saving the model to {save_path}..." ) save_model(model, save_path, accelerator) best_loss = average_dev_loss elif train_config.regression: save_path = save_dir / f"model_best_mae_{train_config.seed}.pt" elif config["model"]["regression"]: save_path = save_dir / f"model_best_mae_{config['seed']}.pt" if mae < best_loss: accelerator.print( f"Dev MAE decreased ({best_loss} -> {mae}). Saving the model to {save_path}..." Loading @@ -519,7 +513,7 @@ def main(): save_model(model, save_path, accelerator) best_loss = mae else: save_path = save_dir / f"model_best_f1_{train_config.seed}.pt" save_path = save_dir / f"model_best_f1_{config['seed']}.pt" if f1_macro > best_f1: accelerator.print( f"Dev F1 increased ({best_f1} -> {f1_macro}). Saving the model to {save_path}..." Loading Loading
config.pydeleted 100644 → 0 +0 −30 Original line number Diff line number Diff line from dataclasses import dataclass, field from typing import List, Optional, Union @dataclass class TrainConfig: batch_size: int = 2 save_dir: str = "saved_models/" bert_model: str = "sentence-transformers/all-distilroberta-v1" multilabel: bool = True regression: bool = False five_classes: bool = False patience: int = 20 seed: int = 2 encoder_hidden_dim: int = 300 encoder_num_layers: int = 1 dropout: float = 0.5 num_classes: int = 8 attention_type: str = "hierarchical" pooling: str = "mean" binary_only: bool = True bidirectional: bool = True regularization_loss: bool = False loss_l: float = 0.1 lr: float = 3e-5 num_iters: int = 100 encoder_layers_to_freeze: Optional[Union[str, List[Union[str, int]]]] = field( default_factory=lambda: ["embeddings", 0, 1, 2, 3, 4] ) save_every_epoch: bool = False
config.toml 0 → 100644 +22 −0 Original line number Diff line number Diff line save_dir = "saved_models" batch_size = 2 seed = 2 five_classes = false regularization_loss = false loss_l = 0.1 lr = 3e-5 num_iters = 100 encoder_layers_to_freeze = [ ["embeddings"], [0, 1, 2, 3, 4] ] [model] bert_model = "sentence-transformers/all-distilroberta-v1" encoder_hidden_dim = 300 encoder_num_layers = 1 dropout = 0.5 num_classes = 8 attention_type = "hierarchical" pooling = "mean" binary_only = true bidirectional = true multilabel = true regression = false No newline at end of file
hdsc/model.py +1 −1 Original line number Diff line number Diff line Loading @@ -7,7 +7,7 @@ from accelerate.utils import send_to_device from torch.nn.utils.rnn import pad_sequence from transformers import AutoModel, AutoConfig from hdsc.modules import ClassificationHead, PositionalEncoding, SentenceLevelEncoder, WordLevelEncoder from hdsc.modules import ClassificationHead, SentenceLevelEncoder PAD_ID = 0 Loading
train.py +59 −65 Original line number Diff line number Diff line import json import itertools from pathlib import Path from typing import Dict, List, Optional import datasets import tomli import torch import torch.nn as nn from accelerate import Accelerator Loading @@ -14,12 +16,10 @@ from torchinfo import summary from torchmetrics.functional import accuracy, f1_score, mean_absolute_error, mean_squared_error from transformers import AutoTokenizer, get_linear_schedule_with_warmup from config import TrainConfig from hdsc.losses import MultiLabelLoss from hdsc.model import PHQTotalMulticlassAttentionModelBERT from hdsc.utils import save_model train_config = TrainConfig() logger = get_logger(__name__) logger.setLevel("INFO") Loading @@ -38,6 +38,7 @@ def collate_fn(batch): def training_step( config: Dict, model: nn.Module, inputs: dict, loss_fn: nn.Module, Loading Loading @@ -70,11 +71,11 @@ def training_step( phq_score = torch.sum(inputs["labels"], dim=1) if train_config.multilabel: if config["model"]["multilabel"]: labels = inputs["labels"] elif train_config.regression: elif config["model"]["regression"]: labels = phq_score elif train_config.five_classes: elif config["five_classes"]: labels = torch.div(phq_score, 5, rounding_mode="floor") labels = labels.to(torch.long) else: Loading @@ -83,7 +84,7 @@ def training_step( with accelerator.autocast(): pred_binary = model(inputs) if train_config.multilabel: if config["model"]["multilabel"]: loss, loss_1, loss_2 = loss_fn( pred_binary, labels, Loading @@ -97,13 +98,13 @@ def training_step( scheduler.step() # Gather the outputs from all GPUs cat_fn = torch.vstack if train_config.multilabel else torch.hstack cat_fn = torch.vstack if config["model"]["multilabel"] else torch.hstack gathered_preds = accelerator.gather(pred_binary).detach() gathered_labels = accelerator.gather(labels).detach() gathered_reg_labels = accelerator.gather(phq_score).detach() # Recalculate the loss on all inputs if train_config.multilabel: if config["model"]["multilabel"]: loss, loss_1, loss_2 = loss_fn(gathered_preds, gathered_labels, gathered_reg_labels) else: loss = loss_fn(gathered_preds, gathered_labels) Loading @@ -123,6 +124,7 @@ def training_step( def evaluate( config: Dict, model: nn.Module, val_dataloader: torch.utils.data.DataLoader, loss_fn: nn.Module, Loading Loading @@ -154,17 +156,17 @@ def evaluate( all_reg_labels: List[torch.Tensor] = [] n_steps = len(val_dataloader) cat_fn = torch.vstack if train_config.multilabel else torch.hstack cat_fn = torch.vstack if config["model"]["multilabel"] else torch.hstack model.eval() with torch.no_grad(): for batch in val_dataloader: phq_score = torch.sum(batch["labels"], dim=1) if train_config.multilabel: if config["model"]["multilabel"]: labels = batch["labels"] elif train_config.regression: elif config["model"]["regression"]: labels = phq_score elif train_config.five_classes: elif config["five_classes"]: labels = torch.div(phq_score, 5, rounding_mode="floor") labels = labels.to(torch.long) else: Loading @@ -177,7 +179,7 @@ def evaluate( gathered_reg_labels = accelerator.gather(phq_score).detach() # Recalculate the loss on all inputs if train_config.multilabel: if config["model"]["multilabel"]: loss, loss_1, loss_2 = loss_fn(gathered_preds, gathered_labels, gathered_reg_labels) else: loss = loss_fn(gathered_preds, gathered_labels) Loading @@ -196,7 +198,7 @@ def evaluate( avg_loss_2 = running_loss_2.item() / n_steps all_preds = torch.vstack(all_preds) if not train_config.multilabel and not train_config.regression: if not config["model"]["multilabel"] and not config["model"]["regression"]: all_preds = all_preds.topk(k=1, dim=1)[1].squeeze(-1) all_labels = cat_fn(all_labels) all_reg_labels = torch.hstack(all_reg_labels) Loading @@ -217,85 +219,77 @@ def main(): accelerator = Accelerator() device = accelerator.device with open("config.toml", "rb") as f: config = tomli.load(f) model_name = "_".join( [ "lstm" if train_config.bert_model == "lstm" else "robert", "multilabel" if train_config.multilabel else "binary", "regression" if train_config.regression else "no-regression", "five_classes" if train_config.five_classes else "", "lstm" if config["model"]["bert_model"] == "lstm" else "robert", "multilabel" if config["model"]["multilabel"] else "binary", "regression" if config["model"]["regression"] else "no-regression", "five_classes" if config["five_classes"] else "", ] ) save_dir = Path(train_config.save_dir) / model_name save_dir = Path(config["save_dir"]) / model_name if not save_dir.exists(): save_dir.mkdir(parents=True) tokenizer = AutoTokenizer.from_pretrained(train_config.bert_model) tokenizer = AutoTokenizer.from_pretrained(config["model"]["bert_model"]) dataset = load_dataset("daic_woz.py", "lines") encoded_dataset = dataset.map( lambda examples: tokenizer(examples["turns"], padding="max_length", truncation=True), load_from_cache_file=False, ) encoded_dataset.set_format(type="torch", columns=["input_ids", "attention_mask", "labels"]) cat_fn = torch.vstack if train_config.multilabel else torch.hstack cat_fn = torch.vstack if config["model"]["multilabel"] else torch.hstack train_dataloader = torch.utils.data.DataLoader( encoded_dataset["train"], batch_size=train_config.batch_size, batch_size=config["batch_size"], collate_fn=collate_fn, shuffle=True, ) validation_dataloader = torch.utils.data.DataLoader( encoded_dataset["validation"], batch_size=train_config.batch_size, batch_size=config["batch_size"], collate_fn=collate_fn, shuffle=False, ) model = PHQTotalMulticlassAttentionModelBERT( train_config.bert_model, train_config.encoder_hidden_dim, train_config.encoder_num_layers, train_config.dropout, train_config.num_classes, train_config.attention_type, train_config.pooling, train_config.binary_only, train_config.bidirectional, train_config.multilabel, train_config.regression, device, ) model = PHQTotalMulticlassAttentionModelBERT(device=device, **config["model"]) model.initialize_encoder_weights() if train_config.multilabel: if config["model"]["multilabel"]: loss_fn = MultiLabelLoss( regularization=train_config.regularization_loss, l=train_config.loss_l, regularization=config["regularization_loss"], l=config["loss_l"], ) elif args.regression: elif config["model"]["regression"]: loss_fn = nn.SmoothL1Loss() else: loss_fn = nn.NLLLoss() optimizer = AdamW(model.parameters(), lr=train_config.lr) total_steps = len(train_dataloader) * train_config.num_iters optimizer = AdamW(model.parameters(), lr=config["lr"]) total_steps = len(train_dataloader) * config["num_iters"] scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=5, num_training_steps=total_steps) model = accelerator.prepare(model) optimizer, train_dataloader, scheduler = accelerator.prepare(optimizer, train_dataloader, scheduler) validation_dataloader = accelerator.prepare(validation_dataloader) if train_config.encoder_layers_to_freeze is not None: if train_config.encoder_layers_to_freeze == "all": if config["encoder_layers_to_freeze"] is not None: if config["encoder_layers_to_freeze"] == "all": model.encoder.requires_grad_(False) else: layers_to_freeze = [str(x) for x in train_config.encoder_layers_to_freeze] layers_to_freeze = [ str(x) for x in list(itertools.chain.from_iterable(config["encoder_layers_to_freeze"])) ] for name, param in model.encoder.named_parameters(): if any([True for layer in layers_to_freeze if layer in name]): param.requires_grad_(False) accelerator.print(train_config) accelerator.print(config) model_stats = summary(model, depth=4, verbose=0) accelerator.print(model_stats) Loading @@ -313,16 +307,16 @@ def main(): global_step = 0 eval_every = len(train_dataloader) with open(save_dir / f"log_{model_name}_{train_config.seed}.tsv", "w", encoding="utf-8") as f: with open(save_dir / f"log_{model_name}_{config['seed']}.tsv", "w", encoding="utf-8") as f: f.write( "epoch\ttrain_loss\ttrain_loss_bin\ttrain_loss_reg\tdev_loss\tdev_loss_bin\tdev_loss_reg\t" + "train_acc\tdev_acc\tf1_micro\tf1_macro\tdev_mae\tdev_mse\n" ) for epoch in range(train_config.num_iters): for epoch in range(config["num_iters"]): running_loss = 0.0 preds = [] for batch in train_dataloader: outputs = training_step(model, batch, loss_fn, optimizer, accelerator, scheduler) outputs = training_step(config, model, batch, loss_fn, optimizer, accelerator, scheduler) running_loss += outputs["loss"] running_loss_bin += outputs["loss_1"] Loading @@ -334,7 +328,7 @@ def main(): global_step += 1 if global_step % eval_every == 0: eval_outputs = evaluate(model, validation_dataloader, loss_fn, accelerator) eval_outputs = evaluate(config, model, validation_dataloader, loss_fn, accelerator) average_train_loss = running_loss.item() / eval_every average_train_loss_bin = running_loss_bin.item() / eval_every Loading @@ -344,7 +338,7 @@ def main(): average_dev_loss_reg = eval_outputs["loss_2"] preds = torch.vstack(preds) if not train_config.multilabel and not train_config.regression: if not config["model"]["multilabel"] and not config["model"]["regression"]: preds = preds.topk(k=1, dim=1)[1].squeeze(-1) dev_preds = eval_outputs["preds"] preds = preds.float() Loading @@ -353,7 +347,7 @@ def main(): dev_labels = eval_outputs["labels"] dev_labels_reg = eval_outputs["reg_labels"] if train_config.multilabel: if config["model"]["multilabel"]: average_train_acc = mean_absolute_error(preds, labels).item() average_dev_acc = mean_absolute_error(dev_preds, dev_labels).item() Loading @@ -371,7 +365,7 @@ def main(): dev_preds >= 1.5, task="multilabel", average="micro", num_labels=train_config.num_classes, num_labels=config["model"]["num_classes"], ).item() mae = mean_absolute_error( Loading @@ -383,7 +377,7 @@ def main(): dev_labels_reg, ).item() f1_mae = f1_micro / mae elif train_config.regression: elif config["model"]["regression"]: preds = preds.squeeze(-1) dev_preds = dev_preds.squeeze(-1) average_train_acc = mean_absolute_error(preds, labels).item() Loading @@ -408,7 +402,7 @@ def main(): dev_labels_reg, ).item() f1_mae = f1_micro / mae elif train_config.five_classes: elif config["five_classes"]: average_train_acc = accuracy(preds.to(torch.long), labels).item() average_dev_acc = accuracy(dev_preds.to(torch.long), dev_labels).item() Loading Loading @@ -443,7 +437,7 @@ def main(): f1_mae = 0.0 f1_samples = 0.0 preds_save_path = save_dir / f"preds_{train_config.seed}_{epoch}.json" preds_save_path = save_dir / f"preds_{config['seed']}_{epoch}.json" preds_to_json = { "pred": dev_preds.tolist(), "true": dev_labels.tolist(), Loading @@ -464,7 +458,7 @@ def main(): dev_labels_reg = [] stats_message = ( f"Epoch [{epoch+1}/{train_config.num_iters}], Step [{global_step}/{train_config.num_iters*len(train_dataloader)}], " f"Epoch [{epoch+1}/{config['num_iters']}], Step [{global_step}/{config['num_iters']*len(train_dataloader)}], " + f"Train Loss: {average_train_loss:.4f}, Dev Loss: {average_dev_loss:.4f}, " + f"Train Loss Binary: {average_train_loss_bin:.4f}, Dev Loss Binary: {average_dev_loss_bin:.4f}, " + f"Train Loss Regression: {average_train_loss_reg:.4f}, Dev Loss Regression: {average_dev_loss_reg:.4f}, " Loading Loading @@ -498,20 +492,20 @@ def main(): f.write("\n") accelerator.wait_for_everyone() if train_config.save_every_epoch: save_path = save_dir / f"model_{train_config.seed}_{epoch}.pt" if config["save_every_epoch"]: save_path = save_dir / f"model_{config['seed']}_{epoch}.pt" save_model(model, save_path, accelerator) else: if train_config.multilabel: save_path = save_dir / f"model_best_loss_{train_config.seed}.pt" if config["model"]["multilabel"]: save_path = save_dir / f"model_best_loss_{config['seed']}.pt" if average_dev_loss < best_loss: accelerator.print( f"Dev loss decreased ({best_loss} -> {average_dev_loss}). Saving the model to {save_path}..." ) save_model(model, save_path, accelerator) best_loss = average_dev_loss elif train_config.regression: save_path = save_dir / f"model_best_mae_{train_config.seed}.pt" elif config["model"]["regression"]: save_path = save_dir / f"model_best_mae_{config['seed']}.pt" if mae < best_loss: accelerator.print( f"Dev MAE decreased ({best_loss} -> {mae}). Saving the model to {save_path}..." Loading @@ -519,7 +513,7 @@ def main(): save_model(model, save_path, accelerator) best_loss = mae else: save_path = save_dir / f"model_best_f1_{train_config.seed}.pt" save_path = save_dir / f"model_best_f1_{config['seed']}.pt" if f1_macro > best_f1: accelerator.print( f"Dev F1 increased ({best_f1} -> {f1_macro}). Saving the model to {save_path}..." Loading