File size: 3,513 Bytes
6dc7c27
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
from typing import Optional

import hydra
import pytorch_lightning as pl
import torch
from transformers import get_linear_schedule_with_warmup

from src.sense_extractors import SenseExtractor
from src.utils.optimizers import RAdam


class ConsecPLModule(pl.LightningModule):
    def __init__(self, conf, *args, **kwargs) -> None:
        super().__init__(*args, **kwargs)
        self.save_hyperparameters(conf)
        self.sense_extractor: SenseExtractor = hydra.utils.instantiate(self.hparams.model.sense_extractor)

        new_embedding_size = self.sense_extractor.model.config.vocab_size + 203
        self.sense_extractor.resize_token_embeddings(new_embedding_size)

    def forward(

        self,

        input_ids: torch.Tensor,

        attention_mask: Optional[torch.Tensor] = None,

        token_type_ids: Optional[torch.Tensor] = None,

        relative_positions: Optional[torch.Tensor] = None,

        definitions_mask: Optional[torch.Tensor] = None,

        gold_markers: Optional[torch.Tensor] = None,

        *args,

        **kwargs,

    ) -> dict:

        sense_extractor_output = self.sense_extractor.extract(
            input_ids, attention_mask, token_type_ids, relative_positions, definitions_mask, gold_markers
        )

        output_dict = {
            "pred_logits": sense_extractor_output.prediction_logits,
            "pred_probs": sense_extractor_output.prediction_probs,
            "pred_markers": sense_extractor_output.prediction_markers,
            "loss": sense_extractor_output.loss,
        }

        return output_dict

    def training_step(self, batch: dict, batch_idx: int) -> torch.Tensor:
        forward_output = self.forward(**batch)
        self.log("loss", forward_output["loss"], on_step=False, on_epoch=True)
        return forward_output["loss"]

    def validation_step(self, batch: dict, batch_idx: int) -> None:
        forward_output = self.forward(**batch)
        self.log(f"val_loss", forward_output["loss"], prog_bar=True)

    def get_optimizer_and_scheduler(self):

        no_decay = self.hparams.train.no_decay_params

        optimizer_grouped_parameters = [
            {
                "params": [p for n, p in self.named_parameters() if not any(nd in n for nd in no_decay)],
                "weight_decay": self.hparams.train.weight_decay,
            },
            {
                "params": [p for n, p in self.named_parameters() if any(nd in n for nd in no_decay)],
                "weight_decay": 0.0,
            },
        ]

        if self.hparams.train.optimizer == "adamw":
            optimizer = torch.optim.AdamW(optimizer_grouped_parameters, self.hparams.train.learning_rate)
        elif self.hparams.train.optimizer == "radam":
            optimizer = RAdam(optimizer_grouped_parameters, self.hparams.train.learning_rate)
            return optimizer, None
        else:
            raise NotImplementedError

        lr_scheduler = get_linear_schedule_with_warmup(
            optimizer=optimizer,
            num_warmup_steps=self.hparams.train.num_warmup_steps,
            num_training_steps=self.hparams.train.num_training_steps,
        )

        return optimizer, lr_scheduler

    def configure_optimizers(self):
        optimizer, lr_scheduler = self.get_optimizer_and_scheduler()
        if lr_scheduler is None:
            return optimizer
        return [optimizer], [{"interval": "step", "scheduler": lr_scheduler}]