| 62 | |
| 63 | # TCTColBERT model |
| 64 | class 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 | ''' |