CoolFace
Modelpublic

pinecone/ConstBERT

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
22likes938downloads
tokenization_utils.py191 linesDownload Raw Back to root
1import torch2from .colbert_configuration import ColBERTConfig3from transformers import AutoTokenizer4 5DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")6 7def _split_into_batches(ids, mask, bsize):8    batches = []9    for offset in range(0, ids.size(0), bsize):10        batches.append((ids[offset:offset+bsize], mask[offset:offset+bsize]))11 12    return batches13 14def _sort_by_length(ids, mask, bsize):15    if ids.size(0) <= bsize:16        return ids, mask, torch.arange(ids.size(0))17 18    indices = mask.sum(-1).sort().indices19    reverse_indices = indices.sort().indices20 21    return ids[indices], mask[indices], reverse_indices22 23class QueryTokenizer():24    def __init__(self, config: ColBERTConfig, verbose: int = 3):25        self.tok = AutoTokenizer.from_pretrained(config.checkpoint)26        self.tok.base = config.checkpoint27        self.verbose = verbose28 29        self.config = config30        self.query_maxlen = config.query_maxlen31        self.background_maxlen = 512 - self.query_maxlen + 1  # FIXME: Make this configurable32 33        self.Q_marker_token, self.Q_marker_token_id = config.query_token, self.tok.convert_tokens_to_ids(config.query_token_id)34        self.cls_token, self.cls_token_id = self.tok.cls_token, self.tok.cls_token_id35        self.sep_token, self.sep_token_id = self.tok.sep_token, self.tok.sep_token_id36        self.mask_token, self.mask_token_id = self.tok.mask_token, self.tok.mask_token_id37        self.pad_token,self.pad_token_id = self.tok.pad_token,self.tok.pad_token_id38        self.used = False39 40    def tokenize(self, batch_text, add_special_tokens=False):41        assert type(batch_text) in [list, tuple], (type(batch_text))42 43        tokens = [self.tok.tokenize(x, add_special_tokens=False) for x in batch_text]44 45        if not add_special_tokens:46            return tokens47 48        prefix, suffix = [self.cls_token, self.Q_marker_token], [self.sep_token]49        tokens = [prefix + lst + suffix + [self.mask_token] * (self.query_maxlen - (len(lst)+3)) for lst in tokens]50 51        return tokens52 53    def encode(self, batch_text, add_special_tokens=False):54        assert type(batch_text) in [list, tuple], (type(batch_text))55 56        ids = self.tok(batch_text, add_special_tokens=False).to(DEVICE)['input_ids']57 58        if not add_special_tokens:59            return ids60 61        prefix, suffix = [self.cls_token_id, self.Q_marker_token_id], [self.sep_token_id]62        ids = [prefix + lst + suffix + [self.mask_token_id] * (self.query_maxlen - (len(lst)+3)) for lst in ids]63 64        return ids65 66    def tensorize(self, batch_text, bsize=None, context=None, full_length_search=False):67        assert type(batch_text) in [list, tuple], (type(batch_text))68 69        # add placehold for the [Q] marker70        batch_text = ['. ' + x for x in batch_text]71 72        # Full length search is only available for single inference (for now)73        # Batched full length search requires far deeper changes to the code base74        assert(full_length_search == False or (type(batch_text) == list and len(batch_text) == 1))75 76        if full_length_search:77            # Tokenize each string in the batch78            un_truncated_ids = self.tok(batch_text, add_special_tokens=False).to(DEVICE)['input_ids']79            # Get the longest length in the batch80            max_length_in_batch = max(len(x) for x in un_truncated_ids)81            # Set the max length82            max_length = self.max_len(max_length_in_batch)83        else:84            # Max length is the default max length from the config85            max_length = self.query_maxlen86 87        obj = self.tok(batch_text, padding='max_length', truncation=True,88                       return_tensors='pt', max_length=max_length).to(DEVICE)89 90        ids, mask = obj['input_ids'], obj['attention_mask']91 92        # postprocess for the [Q] marker and the [MASK] augmentation93        ids[:, 1] = self.Q_marker_token_id94        ids[ids == self.pad_token_id] = self.mask_token_id95 96        if context is not None:97            assert len(context) == len(batch_text), (len(context), len(batch_text))98 99            obj_2 = self.tok(context, padding='longest', truncation=True,100                            return_tensors='pt', max_length=self.background_maxlen).to(DEVICE)101 102            ids_2, mask_2 = obj_2['input_ids'][:, 1:], obj_2['attention_mask'][:, 1:]  # Skip the first [SEP]103 104            ids = torch.cat((ids, ids_2), dim=-1)105            mask = torch.cat((mask, mask_2), dim=-1)106 107        if self.config.attend_to_mask_tokens:108            mask[ids == self.mask_token_id] = 1109            assert mask.sum().item() == mask.size(0) * mask.size(1), mask110 111        if bsize:112            batches = _split_into_batches(ids, mask, bsize)113            return batches114        115        if self.used is False:116            self.used = True117 118            firstbg = (context is None) or context[0]119            if self.verbose > 1:120                print()121                print("#> QueryTokenizer.tensorize(batch_text[0], batch_background[0], bsize) ==")122                print(f"#> Input: {batch_text[0]}, \t\t {firstbg}, \t\t {bsize}")123                print(f"#> Output IDs: {ids[0].size()}, {ids[0]}")124                print(f"#> Output Mask: {mask[0].size()}, {mask[0]}")125                print()126 127        return ids, mask128 129    # Ensure that query_maxlen <= length <= 500 tokens130    def max_len(self, length):131        return min(500, max(self.query_maxlen, length))132 133 134class DocTokenizer():135    def __init__(self, config: ColBERTConfig):136        self.tok = AutoTokenizer.from_pretrained(config.checkpoint)137        self.tok.base = config.checkpoint138 139        self.config = config140        self.doc_maxlen = config.doc_maxlen141 142        self.D_marker_token, self.D_marker_token_id = self.config.doc_token, self.tok.convert_tokens_to_ids(self.config.doc_token_id)143        self.cls_token, self.cls_token_id = self.tok.cls_token, self.tok.cls_token_id144        self.sep_token, self.sep_token_id = self.tok.sep_token, self.tok.sep_token_id145 146    def tokenize(self, batch_text, add_special_tokens=False):147        assert type(batch_text) in [list, tuple], (type(batch_text))148 149        tokens = [self.tok.tokenize(x, add_special_tokens=False).to(DEVICE) for x in batch_text]150 151        if not add_special_tokens:152            return tokens153 154        prefix, suffix = [self.cls_token, self.D_marker_token], [self.sep_token]155        tokens = [prefix + lst + suffix for lst in tokens]156 157        return tokens158 159    def encode(self, batch_text, add_special_tokens=False):160        assert type(batch_text) in [list, tuple], (type(batch_text))161 162        ids = self.tok(batch_text, add_special_tokens=False).to(DEVICE)['input_ids']163 164        if not add_special_tokens:165            return ids166 167        prefix, suffix = [self.cls_token_id, self.D_marker_token_id], [self.sep_token_id]168        ids = [prefix + lst + suffix for lst in ids]169 170        return ids171 172    def tensorize(self, batch_text, bsize=None):173        assert type(batch_text) in [list, tuple], (type(batch_text))174 175        # add placehold for the [D] marker176        batch_text = ['. ' + x for x in batch_text]177 178        obj = self.tok(batch_text, padding='max_length', truncation='longest_first',179                       return_tensors='pt', max_length=self.doc_maxlen).to(DEVICE)180 181        ids, mask = obj['input_ids'], obj['attention_mask']182 183        # postprocess for the [D] marker184        ids[:, 1] = self.D_marker_token_id185 186        if bsize:187            ids, mask, reverse_indices = _sort_by_length(ids, mask, bsize)188            batches = _split_into_batches(ids, mask, bsize)189            return batches, reverse_indices190 191        return ids, mask