MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / DPRReaderFinalMixin

Class DPRReaderFinalMixin

SwissArmyTransformer/sat/model/official/dpr_model.py:13–41  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

11 return logits
12
13class DPRReaderFinalMixin(BaseMixin):
14 def __init__(self, hidden_size, projection_dim):
15 super().__init__()
16 if projection_dim > 0:
17 embeddings_size = projection_dim
18 else:
19 embeddings_size = hidden_size
20 self.qa_outputs = nn.Linear(embeddings_size, 2)
21 self.qa_classifier = nn.Linear(embeddings_size, 1)
22
23 def final_forward(self, logits, **kwargs):
24 # notations: N - number of questions in a batch, M - number of passages per questions, L - sequence length
25 print("Before final_forward: logits = ", logits)
26 n_passages, sequence_length = logits.size()[:2]
27 sequence_output = logits
28
29 # compute logits
30 logits = self.qa_outputs(sequence_output)
31 start_logits, end_logits = logits.split(1, dim=-1)
32 start_logits = start_logits.squeeze(-1).contiguous()
33 end_logits = end_logits.squeeze(-1).contiguous()
34 relevance_logits = self.qa_classifier(sequence_output[:, 0, :])
35
36 # resize
37 start_logits = start_logits.view(n_passages, sequence_length)
38 end_logits = end_logits.view(n_passages, sequence_length)
39 relevance_logits = relevance_logits.view(n_passages)
40
41 return (start_logits, end_logits, relevance_logits)
42
43class DPRTypeMixin(BaseMixin):
44 def __init__(self, num_types, hidden_size):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected