MCPcopy Create free account
hub / github.com/BeastyZ/ConvSearch-R1 / TCTColBERT

Class TCTColBERT

index/dense/models.py:64–96  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

62
63# TCTColBERT model
64class TCTColBERT(nn.Module):
65 def __init__(self, model_path) -> None:
66 super(TCTColBERT, self).__init__()
67 self.model = BertModel.from_pretrained(model_path)
68
69 def forward(self, input_ids, attention_mask, **kwargs):
70 outputs = self.model(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
71
72 if "cur_utt_end_position" in kwargs:
73 device = outputs.device
74 cur_utt_end_positions = kwargs["cur_utt_end_positions"]
75 output_mask = torch.zeros(attention_mask.size()).to(device)
76 mask_row = []
77 mask_col = []
78 for i in range(len(cur_utt_end_positions)):
79 mask_row += [i] * (cur_utt_end_positions[i] - 3)
80 mask_col += list(range(4, cur_utt_end_positions[i] + 1))
81
82 mask_index = (
83 torch.tensor(mask_row).long().to(device),
84 torch.tensor(mask_col).long().to(device)
85 )
86 values = torch.ones(len(mask_row)).to(device)
87 output_mask = output_mask.index_put(mask_index, values)
88 else:
89 output_mask = attention_mask
90 output_mask[:, :4] = 0 # filter the first 4 tokens: [CLS] "[" "Q/D" "]"
91
92 # sum / length
93 sum_outputs = torch.sum(outputs * output_mask.unsqueeze(-1), dim = -2)
94 real_seq_length = torch.sum(output_mask, dim = 1).view(-1, 1)
95
96 return sum_outputs / real_seq_length
97
98
99'''

Callers 1

load_modelFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected