Commit 46f90861 authored by Kirill Milintsevich's avatar Kirill Milintsevich
Browse files

Config now uses TOML format

Old code clean-up
parent 2258db61
Loading
Loading
Loading
Loading

config.py

deleted100644 → 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
+1 −1
Original line number Diff line number Diff line
@@ -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

+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
@@ -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")

@@ -38,6 +38,7 @@ def collate_fn(batch):


def training_step(
    config: Dict,
    model: nn.Module,
    inputs: dict,
    loss_fn: nn.Module,
@@ -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:
@@ -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,
@@ -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)
@@ -123,6 +124,7 @@ def training_step(


def evaluate(
    config: Dict,
    model: nn.Module,
    val_dataloader: torch.utils.data.DataLoader,
    loss_fn: nn.Module,
@@ -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:
@@ -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)
@@ -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)
@@ -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)

@@ -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"]
@@ -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
@@ -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()
@@ -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()

@@ -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(
@@ -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()
@@ -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()

@@ -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(),
@@ -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}, "
@@ -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}..."
@@ -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}..."