MCPcopy Create free account
hub / github.com/DLVulDet/PrimeVul / DefectModel

Class DefectModel

os_expr/model.py:140–202  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

138
139
140class DefectModel(nn.Module):
141 def __init__(self, encoder, config, tokenizer, args):
142 super(DefectModel, self).__init__()
143 self.encoder = encoder
144 self.config = config
145 self.tokenizer = tokenizer
146 self.classifier = nn.Linear(config.hidden_size, 2)
147 self.args = args
148
149 def get_t5_vec(self, source_ids):
150 attention_mask = source_ids.ne(self.tokenizer.pad_token_id)
151 outputs = self.encoder(input_ids=source_ids, attention_mask=attention_mask,
152 labels=source_ids, decoder_attention_mask=attention_mask, output_hidden_states=True)
153 hidden_states = outputs['decoder_hidden_states'][-1]
154 eos_mask = source_ids.eq(self.config.eos_token_id)
155
156 if len(torch.unique(eos_mask.sum(1))) > 1:
157 print(eos_mask.sum(1))
158 print(torch.unique(eos_mask.sum(1)))
159 raise ValueError("All examples must have the same number of <eos> tokens.")
160 vec = hidden_states[eos_mask, :].view(hidden_states.size(0), -1,
161 hidden_states.size(-1))[:, -1, :]
162 return vec
163
164 def get_bart_vec(self, source_ids):
165 attention_mask = source_ids.ne(self.tokenizer.pad_token_id)
166 outputs = self.encoder(input_ids=source_ids, attention_mask=attention_mask,
167 labels=source_ids, decoder_attention_mask=attention_mask, output_hidden_states=True)
168 hidden_states = outputs['decoder_hidden_states'][-1]
169 eos_mask = source_ids.eq(self.config.eos_token_id)
170
171 if len(torch.unique(eos_mask.sum(1))) > 1:
172 raise ValueError("All examples must have the same number of <eos> tokens.")
173 vec = hidden_states[eos_mask, :].view(hidden_states.size(0), -1,
174 hidden_states.size(-1))[:, -1, :]
175 return vec
176
177 def get_roberta_vec(self, source_ids):
178 attention_mask = source_ids.ne(self.tokenizer.pad_token_id)
179 vec = self.encoder(input_ids=source_ids, attention_mask=attention_mask)[0][:, 0, :]
180 return vec
181
182 def forward(self, source_ids=None, labels=None, weight=None):
183 # source_ids = source_ids.view(-1, self.args.max_source_length)
184
185 if self.args.model_type == 'codet5':
186 vec = self.get_t5_vec(source_ids)
187 elif self.args.model_type == 'bart':
188 vec = self.get_bart_vec(source_ids)
189 elif self.args.model_type == 'roberta':
190 vec = self.get_roberta_vec(source_ids)
191 elif self.args.model_type == 't5':
192 vec = self.get_t5_vec(source_ids)
193
194 logits = self.classifier(vec)
195 prob = nn.functional.softmax(logits)
196
197 if labels is not None:

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected