| 16 | layer.requires_grade = False |
| 17 | |
| 18 | class JustificationGenerationModule(pl.LightningModule): |
| 19 | |
| 20 | def __init__(self, tokenizer, model, learning_rate=1e-3, gen_num_beams=2, gen_max_length=100, should_pad_gen=True): |
| 21 | super().__init__() |
| 22 | self.tokenizer = tokenizer |
| 23 | self.model = model |
| 24 | self.learning_rate = learning_rate |
| 25 | |
| 26 | self.gen_num_beams = gen_num_beams |
| 27 | self.gen_max_length = gen_max_length |
| 28 | self.should_pad_gen = should_pad_gen |
| 29 | |
| 30 | #self.metrics = datasets.load_metric('meteor') |
| 31 | |
| 32 | freeze_params(self.model.get_encoder()) |
| 33 | self.freeze_embeds() |
| 34 | |
| 35 | def freeze_embeds(self): |
| 36 | ''' freeze the positional embedding parameters of the model; adapted from finetune.py ''' |
| 37 | freeze_params(self.model.model.shared) |
| 38 | for d in [self.model.model.encoder, self.model.model.decoder]: |
| 39 | freeze_params(d.embed_positions) |
| 40 | freeze_params(d.embed_tokens) |
| 41 | |
| 42 | # Do a forward pass through the model |
| 43 | def forward(self, input_ids, **kwargs): |
| 44 | return self.model(input_ids, **kwargs) |
| 45 | |
| 46 | def configure_optimizers(self): |
| 47 | optimizer = AdamW(self.parameters(), lr = self.learning_rate) |
| 48 | return optimizer |
| 49 | |
| 50 | def shift_tokens_right(self, input_ids: torch.Tensor, pad_token_id: int, decoder_start_token_id: int): |
| 51 | """ |
| 52 | Shift input ids one token to the right. |
| 53 | https://github.com/huggingface/transformers/blob/main/src/transformers/models/bart/modeling_bart.py. |
| 54 | """ |
| 55 | shifted_input_ids = input_ids.new_zeros(input_ids.shape) |
| 56 | shifted_input_ids[:, 1:] = input_ids[:, :-1].clone() |
| 57 | shifted_input_ids[:, 0] = decoder_start_token_id |
| 58 | |
| 59 | if pad_token_id is None: |
| 60 | raise ValueError("self.model.config.pad_token_id has to be defined.") |
| 61 | # replace possible -100 values in labels by `pad_token_id` |
| 62 | shifted_input_ids.masked_fill_(shifted_input_ids == -100, pad_token_id) |
| 63 | |
| 64 | return shifted_input_ids |
| 65 | |
| 66 | def run_model(self, batch): |
| 67 | src_ids, src_mask, tgt_ids = batch[0], batch[1], batch[2] |
| 68 | |
| 69 | decoder_input_ids = self.shift_tokens_right( |
| 70 | tgt_ids, self.tokenizer.pad_token_id, self.tokenizer.pad_token_id # BART uses the EOS token to start generation as well. Might have to change for other models. |
| 71 | ) |
| 72 | |
| 73 | outputs = self(src_ids, attention_mask=src_mask, decoder_input_ids=decoder_input_ids, use_cache=False) |
| 74 | return outputs |
| 75 |
no outgoing calls
no test coverage detected