MCPcopy Create free account
hub / github.com/MichSchli/AVeriTeC / JustificationGenerationModule

Class JustificationGenerationModule

models/JustificationGenerationModule.py:18–193  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

16 layer.requires_grade = False
17
18class 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

Calls

no outgoing calls

Tested by

no test coverage detected