CoolFace
Apppublic

shekkari21/codereviewer

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
utils.py823 linesDownload Raw Back to root
1import re, json2import os, random3import torch, logging4from copy import deepcopy as cp5from torch.utils.data import Dataset6from tokenizers import ByteLevelBPETokenizer7from transformers import T5Tokenizer, RobertaTokenizer8import nltk9 10logging.basicConfig(11    format="%(asctime)s - %(levelname)s - %(name)s -   %(message)s",12    datefmt="%m/%d/%Y %H:%M:%S",13    level=logging.INFO,14)15logger = logging.getLogger(__name__)16 17 18 19class MyTokenizer(object):20    """21    Wrapper for ByteLevelBPETokenizer22    """23    def __init__(self, vocab=None, merges=None, **kwargs):24        self.tokenizer = ByteLevelBPETokenizer(vocab, merges, **kwargs)25        self.update_id2token()26 27    @staticmethod28    def from_pretrained(path):29        vocabp = os.path.join(path, "vocab.json")30        mergesp = os.path.join(path, "merges.txt")31        mytoken = MyTokenizer(vocabp, mergesp)32        return mytoken33 34    def update_id2token(self):35        vocab = self.tokenizer.get_vocab()36        self.id2token = {vocab[token]: token for token in vocab}37 38    def add_special_tokens(self, dic):39        for values in dic.values():40            self.tokenizer.add_special_tokens(values)41        self.update_id2token()42 43    def convert_ids_to_tokens(self, ids):44        vocab = self.id2token45        return [vocab[i] for i in ids]46    47    def decode(self, ids, **kwargs):    ##### to be update48        tokens = self.convert_ids_to_tokens(ids)49        return " ".join(tokens)50 51    def encode(self, text, **kwargs):52        text = text.encode("ascii", errors="ignore").decode("ascii")53        return self.tokenizer.encode(text).ids54 55    def get_vocab(self):56        return self.tokenizer.get_vocab()57 58    def __len__(self):59        return len(self.tokenizer.get_vocab())60 61 62class RefineFeatures(object):63    def __init__(self, example_id, source_ids, target_ids):64        self.example_id = example_id65        self.source_ids = source_ids66        self.target_ids = target_ids67 68class RefineDataset(Dataset):69    def __init__(self, tokenizer, pool, args, file_path, samplenum=-1):70        self.tokenizer = tokenizer71        self.args = args72        logger.info("Reading examples from {}".format(file_path))73        examples = [json.loads(line) for line in open(file_path)]74        for i in range(len(examples)):75            if "id" not in examples[i]:76                examples[i]["id"] = i77        if samplenum > 0:78            examples = examples[:samplenum]79        logger.info(f"Tokenize examples: {file_path}")80        self.feats = pool.map(self.tokenize, \81            [(example, tokenizer, args) for example in examples])82        83    def tokenize(self, item):84        example, tokenizer, args = item85        oldlines = example["old"].split("\n")86        newlines = example["new"].split("\n")87        oldlines = [line[1:].strip() for line in oldlines]88        newlines = [line[1:].strip() for line in newlines]89        oldlines = "\n".join(oldlines)90        newlines = "\n".join(newlines)91        oldlines = "<add>" + oldlines.replace("\n", "<add>")92        newlines = "<add>" + newlines.replace("\n", "<add>")93        comment = example["comment"]94        srcids = self.encode_remove(tokenizer, oldlines, args)95        srcids += [tokenizer.msg_id] + self.encode_remove(tokenizer, comment, args)96        tgtids = self.encode_remove(tokenizer, newlines, args)97        srcids, tgtids = self.pad_assert(srcids, tgtids, args, tokenizer)98        return RefineFeatures(example["id"], srcids, tgtids)99 100    @staticmethod101    def process_pred_gold(pred, gold):102        gold = gold.split("\n")103        gold = [line[1:].strip() for line in gold]104        gold = " ".join(gold)105        pred = " ".join(pred.split())106        pred = pred.replace("<add> ", "")107        return pred, gold108 109    def pad_assert(self, source_ids, target_ids, args, tokenizer):110        source_ids = source_ids[:args.max_source_length - 2]111        source_ids = [tokenizer.bos_id] + source_ids + [tokenizer.eos_id]112        pad_len = args.max_source_length - len(source_ids)113        source_ids += [tokenizer.pad_id] * pad_len114        target_ids = target_ids[:args.max_target_length - 2]115        target_ids = [tokenizer.bos_id] + target_ids + [tokenizer.eos_id]116        pad_len = args.max_target_length - len(target_ids)117        target_ids += [tokenizer.pad_id] * pad_len118        assert len(source_ids) == args.max_source_length, "Not equal length."119        assert len(target_ids) == args.max_target_length, "Not equal length."120        return source_ids, target_ids121 122    def encode_remove(self, tokenizer, text, args):123        text = tokenizer.encode(text, max_length=args.max_source_length, truncation=True)124        if type(tokenizer) == T5Tokenizer:125            return text[:-1]126        elif type(tokenizer) == RobertaTokenizer:127            return text[1:-1]128        elif type(tokenizer) == MyTokenizer:129            return text130        else:131            raise NotImplementedError132 133    def __len__(self):134        return len(self.feats)135 136    def __getitem__(self, i):137        return self.feats[i]138 139class SimpleRefineDataset(RefineDataset):140    def tokenize(self, item):141        example, tokenizer, args = item142        oldlines = example["old"].split("\n")143        newlines = example["new"].split("\n")144        oldlines = [line[1:].strip() for line in oldlines]145        newlines = [line[1:].strip() for line in newlines]146        oldlines = " ".join(oldlines)147        newlines = " ".join(newlines)148        comment = example["comment"]149        srcids = self.encode_remove(tokenizer, oldlines, args)150        srcids += [tokenizer.msg_id] + self.encode_remove(tokenizer, comment, args)151        tgtids = self.encode_remove(tokenizer, newlines, args)152        srcids, tgtids = self.pad_assert(srcids, tgtids, args, tokenizer)153        return RefineFeatures(example["id"], srcids, tgtids)154 155    @staticmethod156    def process_pred_gold(pred, gold):157        gold = gold.split("\n")158        gold = [line[1:].strip() for line in gold]159        gold = " ".join(gold)160        pred = " ".join(pred.split())161        return pred, gold162 163 164class Seq2SeqDataset(RefineDataset):165    def tokenize(self, item):166        example, tokenizer, args = item167        inputs, outputs = example["old"], example["new"]168        inputs = " ".join(inputs.split())169        outputs = " ".join(outputs.split())170        srcids = self.encode_remove(tokenizer, inputs, args)171        tgtids = self.encode_remove(tokenizer, outputs, args)172        srcids, tgtids = self.pad_assert(srcids, tgtids, args, tokenizer)173        return RefineFeatures(example["id"], srcids, tgtids)174    175    @staticmethod176    def process_pred_gold(pred, gold):177        gold = " ".join(gold.split())178        pred = " ".join(pred.split())179        return pred, gold180 181 182class TextDataset(Dataset):183    def __init__(self, tokenizer, pool, args, file_path, samplenum=-1):184        self.cnt = 0185        self.tokenizer = tokenizer186        self.args = args187        if isinstance(tokenizer, MyTokenizer):188            tokenizer_type = "mytok"189        elif isinstance(tokenizer, T5Tokenizer):190            tokenizer_type = ""191        elif isinstance(tokenizer, RobertaTokenizer):192            tokenizer_type = "rb"193        else:194            tokenizer_type = "unk"195        savep = file_path.replace(".jsonl", tokenizer_type + ".exps")196        # savep = "/home/v-zhuoli1/lzzz/processed/chunk_25.exps"197        if os.path.exists(savep):198            logger.info("Loading examples from {}".format(savep))199            examples = torch.load(savep)200        else:201            logger.info("Reading examples from {}".format(file_path))202            examples = read_review_examples(file_path, samplenum, tokenizer)203            logger.info(f"Tokenize examples: {file_path}")204            examples = pool.map(self.tokenize, \205                [(example, tokenizer, args) for example in examples])206            torch.save(examples, savep)207        logger.info("Convert examples to features...")208        self.set_start_end_ids(examples)209        self.featss = pool.map(self.convert_examples_to_features, \210            [(example, tokenizer, args) for example in examples])211        self.feats = [feat for feats in self.featss for feat in feats]  # expand the lists212 213    def __len__(self):214        return len(self.feats)215 216    def __getitem__(self, i):217        return self.feats[i]218 219    def reset_len(self, data_len):220        assert len(self.feats) >= data_len221        self.feats = self.feats[:data_len]222 223    def set_start_end_ids(self, examples):224        for example in examples:225            labels = example.labels226            start_id = 0227            end_id = len(labels) - 1228            for i, label in enumerate(labels):229                if label != -100:               # find the first label230                    start_id = i231                    break232            for i in range(len(labels) - 1, -1, -1):233                label = labels[i]234                if label != -100:235                    end_id = i236                    break237            example.start_id = start_id238            example.end_id = end_id239 240    def tokenize(self, item):241        example, tokenizer, args = item242        example.input = self.encode_remove(tokenizer, example.input, args)243        e0id = tokenizer.special_dict["<e0>"]244        inputs = " ".join(str(id) for id in example.input)245        lines = inputs.split(" " + str(e0id) + " ")246        lines = [247            [int(v) for v in line.split(" ") if len(v) > 0] for line in lines248        ]249        lens = [len(line) for line in lines]250        # if 0 in lens:251        #     logger.info("Warning: empty line in an example.")252        lens = list(map(len, lines))253        curlen = len(lens) + sum(lens)254        left, right = 0, len(lines)255        while curlen > args.max_source_length - 2:256            if left % 2 == 0:257                curlen -= 1 + len(lines[left])258                left += 1259            else:260                right -= 1261                curlen -= 1 + len(lines[right])262        lines = lines[left:right]263        labels = example.labels[left:right]264        assert len(lines) + sum(map(len, lines)) <= args.max_source_length - 2, "Too long inputs in TextDataset.tokenize."265        if len(lines) != len(labels):266            logger.info("Not equal length in TextDataset.tokenize.")267            lines = lines[:len(labels)]268            labels = labels[:len(lines)]269        example.lines = lines270        example.labels = labels271        example.msg = self.encode_remove(tokenizer, example.msg, args)272        return example273 274    def convert_examples_to_features(self, item):275        example, _, _ = item276        if len(example.msg) > 0:277            exs = []278            for _ in range(3):  # up sampling279                if random.random() < 0.5:280                    exs.append(self.genmsg_example(item))281                else:282                    exs.append(self.daemsg_example(item))283            return exs284        if random.random() < 0.5:285            return [self.encoder_example(item)]286        return [self.decoder_example(item)]287 288    def encoder_example(self, item):289        example, tokenizer, args = item290        lines = example.lines291        labels = example.labels292        target_ids = [tokenizer.pad_id] * args.max_target_length293        source_ids, input_labels = [], []294        for i, (line, label) in enumerate(zip(lines, labels)):295            if i == example.start_id:296                source_ids.append(tokenizer.start_id)297                input_labels.append(-100)298            if label != -100:       # only insert special tokens at diffs, not context299                source_ids.append(tokenizer.mask_id)300                input_labels.append(label)301            source_ids.extend(line)302            input_labels.extend([-100] * len(line))303            if i == example.end_id:304                source_ids.append(tokenizer.end_id)305                input_labels.append(-100)306        assert len(input_labels) == len(source_ids), "Not equal length."307        assert len(input_labels) <= args.max_source_length, f"Too long inputs: {len(input_labels)}."308        source_ids = source_ids[:args.max_source_length - 2]309        input_labels = input_labels[:args.max_source_length - 2]310        source_ids = [tokenizer.bos_id] + source_ids + [tokenizer.eos_id]311        input_labels = [-100] + input_labels + [-100]312        pad_len = args.max_source_length - len(source_ids)313        source_ids += [tokenizer.pad_id] * pad_len314        input_labels += [-100] * pad_len315 316        new_input_labels = []317        map_dict = {0: tokenizer.del_id, 1: tokenizer.add_id, 2: tokenizer.keep_id}318        for label in input_labels:319            if label == -100:320                new_input_labels.append(-100)321            else:322                new_input_labels.append(map_dict[label])323        input_labels = new_input_labels324        assert len(source_ids) == args.max_source_length, "Not equal length."325        assert len(input_labels) == args.max_source_length, "Not equal length."326        return ReviewFeatures(example.idx, source_ids, input_labels, target_ids, type="label")327 328    def decoder_example(self, item):329        example, tokenizer, args = item330        lines = example.lines331        labels = example.labels332 333        input_labels = [-100] * args.max_source_length334        source_ids, target_ids = [], []335        SPECIAL_ID = 0336        mask_idxs = random.choices(range(len(lines)), k=int(len(lines) * args.mask_rate))337        id_dict = {0: tokenizer.del_id, 1: tokenizer.add_id, 2: tokenizer.keep_id}338        for i, (line, label) in enumerate(zip(lines, labels)):339            if i == example.start_id:340                source_ids.append(tokenizer.start_id)341            if label in id_dict:342                source_ids.append(id_dict[label])343            if i in mask_idxs:344                source_ids.append(tokenizer.special_dict[f"<e{SPECIAL_ID}>"])345                target_ids.append(tokenizer.special_dict[f"<e{SPECIAL_ID}>"])346                target_ids.extend(line)347                if SPECIAL_ID < 99:     # only 0-99 ids in vocab348                    SPECIAL_ID += 1349            else:350                source_ids.extend(line)351            if i == example.end_id:352                source_ids.append(tokenizer.end_id)353        source_ids, target_ids = self.pad_assert(source_ids, target_ids, args, tokenizer)354        return ReviewFeatures(example.idx, source_ids, input_labels, target_ids, type="line")355 356    def genmsg_example(self, item):357        example, tokenizer, args = item358        lines = example.lines359        labels = example.labels360        input_labels = [-100] * args.max_source_length361        source_ids, target_ids = [], []362        id_dict = {0: tokenizer.del_id, 1: tokenizer.add_id, 2: tokenizer.keep_id}363        for i, (line, label) in enumerate(zip(lines, labels)):364            if i == example.start_id:365                source_ids.append(tokenizer.start_id)366            if label != -100:367                source_ids.append(id_dict[label])368            source_ids.extend(line)369            if i == example.end_id:370                source_ids.append(tokenizer.end_id)371        target_ids.append(tokenizer.msg_id)372        target_ids.extend(example.msg)373        assert len(source_ids) <= args.max_source_length, f"Too long inputs: {len(source_ids)}."374        source_ids, target_ids = self.pad_assert(source_ids, target_ids, args, tokenizer)375        return ReviewFeatures(example.idx, source_ids, input_labels, target_ids, type="genmsg")376 377    def daemsg_example(self, item):378        example, tokenizer, args = item379        input_labels = [-100] * args.max_source_length380        source_ids, target_ids = [], []381        msg_ids = cp(example.msg)382        masks = [random.random() < 0.20 for _ in range(len(msg_ids))]383        if sum(masks) == 0:384            idx = random.choice(range(len(msg_ids)))385            masks[idx] = True386        source_ids, target_ids = [], []387        i = 0388        SPECIAL_ID = 0389        while i < len(masks):390            j = i391            while j < len(masks) and not masks[j]:392                source_ids.append(msg_ids[j])393                j += 1394            if j == len(masks):395                break396            source_ids.append(tokenizer.special_dict[f"<e{SPECIAL_ID}>"])397            target_ids.append(tokenizer.special_dict[f"<e{SPECIAL_ID}>"])398            while j < len(masks) and masks[j]:399                target_ids.append(msg_ids[j])400                j += 1401            if SPECIAL_ID < 99:     # only 0-99 ids in vocab402                SPECIAL_ID += 1403            i = j404        source_ids, target_ids = self.pad_assert(source_ids, target_ids, args, tokenizer)405        return ReviewFeatures(example.idx, source_ids, input_labels, target_ids, type="daemsg")406 407    def pad_assert(self, source_ids, target_ids, args, tokenizer):408        source_ids = source_ids[:args.max_source_length - 2]409        source_ids = [tokenizer.bos_id] + source_ids + [tokenizer.eos_id]410        pad_len = args.max_source_length - len(source_ids)411        source_ids += [tokenizer.pad_id] * pad_len412        target_ids = target_ids[:args.max_target_length - 1]413        target_ids = target_ids + [tokenizer.eos_id]414        pad_len = args.max_target_length - len(target_ids)415        target_ids += [tokenizer.pad_id] * pad_len416        assert len(source_ids) == args.max_source_length, "Not equal length."417        assert len(target_ids) == args.max_target_length, "Not equal length."418        return source_ids, target_ids419 420    def encode_remove(self, tokenizer, text, args):421        text = tokenizer.encode(text, max_length=args.max_source_length, truncation=True)422        if type(tokenizer) == T5Tokenizer:423            return text[:-1]424        elif type(tokenizer) == RobertaTokenizer:425            return text[1:-1]426        elif type(tokenizer) == MyTokenizer:427            return text428        else:429            raise NotImplementedError430 431 432class CommentGenDataset(TextDataset):433    def __init__(self, tokenizer, pool, args, file_path, samplenum=-1):434        self.tokenizer = tokenizer435        if isinstance(tokenizer, MyTokenizer):436            tokenizer_type = "mytok"437        elif isinstance(tokenizer, T5Tokenizer):438            tokenizer_type = ""439        elif isinstance(tokenizer, RobertaTokenizer):440            tokenizer_type = "rb"441        else:442            tokenizer_type = "unk"443        savep = file_path.replace(".jsonl", tokenizer_type + ".exps")444        if os.path.exists(savep):445            logger.info("Loading examples from {}".format(savep))446            examples = torch.load(savep)447        else:448            logger.info("Reading examples from {}".format(file_path))449            examples = read_review_examples(file_path, samplenum, tokenizer)450            # for i in range(len(examples)):451            #     examples[i].msg = " ".join(nltk.word_tokenize(examples[i].msg))452            logger.info(f"Tokenize examples: {file_path}")453            examples = pool.map(self.tokenize, \454                [(example, tokenizer, args) for example in examples])455            torch.save(examples, savep)456        logger.info("Convert examples to features...")457        self.set_start_end_ids(examples)458        self.feats = pool.map(self.convert_examples_to_features, \459            [(example, tokenizer, args) for example in examples])460        self.feats = [feat for feat in self.feats if feat is not None]461 462    def convert_examples_to_features(self, item):463        example, tokenizer, args = item464        if len(example.msg) == 0:465            return None466        return self.genmsg_example(item)467 468 469class CommentClsDataset(TextDataset):470    def __init__(self, tokenizer, pool, args, file_path, samplenum=-1):471        self.tokenizer = tokenizer472        if isinstance(tokenizer, MyTokenizer):473            tokenizer_type = "mytok"474        elif isinstance(tokenizer, T5Tokenizer):475            tokenizer_type = ""476        elif isinstance(tokenizer, RobertaTokenizer):477            tokenizer_type = "rb"478        else:479            tokenizer_type = "unk"480        savep = file_path.replace(".jsonl", tokenizer_type + ".exps")481        if os.path.exists(savep):482            logger.info("Loading examples from {}".format(savep))483            examples = torch.load(savep)484        else:485            logger.info("Reading examples from {}".format(file_path))486            examples = read_review_examples(file_path, samplenum, tokenizer)487            logger.info(f"Tokenize examples: {file_path}")488            examples = pool.map(self.tokenize, \489                [(example, tokenizer, args) for example in examples])490            torch.save(examples, savep)491        logger.info("Convert examples to features...")492        self.set_start_end_ids(examples)493        self.feats = pool.map(self.convert_examples_to_features, \494            [(example, tokenizer, args) for example in examples])495 496    def convert_examples_to_features(self, item):497        example, tokenizer, args = item498        tmpfeature = self.genmsg_example(item)499        return ClsFeatures(tmpfeature.example_id, tmpfeature.source_ids, example.y)500 501 502class SimpleClsDataset(TextDataset):503    def __init__(self, tokenizer, pool, args, file_path, samplenum=-1):504        self.tokenizer = tokenizer505        if isinstance(tokenizer, MyTokenizer):506            tokenizer_type = "mytok"507        elif isinstance(tokenizer, T5Tokenizer):508            tokenizer_type = ""509        elif isinstance(tokenizer, RobertaTokenizer):510            tokenizer_type = "rb"511        else:512            tokenizer_type = "unk"513        savep = file_path.replace(".jsonl", tokenizer_type + ".simpexps")514        if os.path.exists(savep):515            logger.info("Loading examples from {}".format(savep))516            self.feats = torch.load(savep)517        else:518            logger.info("Reading examples from {}".format(file_path))519            examples = read_review_examples(file_path, samplenum, tokenizer)520            logger.info(f"Tokenize examples: {file_path}")521            self.feats = pool.map(self.convert_examples_to_features, \522                [(example, tokenizer, args) for example in examples])523            torch.save(self.feats, savep)524 525    def convert_examples_to_features(self, item):526        example, tokenizer, args = item527        example.input_lines = example.input.split("<e0>")528        labels_l = len(example.labels)529        example.input_lines = example.input_lines[:labels_l]530        for i in range(len(example.input_lines)):531            if example.labels[i] == 1:532                example.input_lines[i] = "+ " + example.input_lines[i]533            elif example.labels[i] == 0:534                example.input_lines[i] = "- " + example.input_lines[i]535        example.input = " ".join(example.input_lines)536        input_ids = self.encode_remove(tokenizer, example.input, args)537        exceed_l = len(input_ids) - args.max_source_length + 2538        if exceed_l > 0:539            halfexl = (exceed_l + 1) // 2540            input_ids = input_ids[halfexl:-halfexl]541        source_ids = input_ids[:args.max_source_length - 2]542        source_ids = [tokenizer.bos_id] + source_ids + [tokenizer.eos_id]543        pad_len = args.max_source_length - len(source_ids)544        source_ids += [tokenizer.pad_id] * pad_len545        example_id = example.idx546        y = example.y547        return ClsFeatures(example_id, source_ids, y)548 549 550class SimpleGenDataset(TextDataset):551    def __init__(self, tokenizer, pool, args, file_path, samplenum=-1):552        self.tokenizer = tokenizer553        if isinstance(tokenizer, MyTokenizer):554            tokenizer_type = "mytok"555        elif isinstance(tokenizer, T5Tokenizer):556            tokenizer_type = ""557        elif isinstance(tokenizer, RobertaTokenizer):558            tokenizer_type = "rb"559        else:560            tokenizer_type = "unk"561        savep = file_path.replace(".jsonl", tokenizer_type + ".simpgenexps")562        if os.path.exists(savep):563            logger.info("Loading examples from {}".format(savep))564            self.feats = torch.load(savep)565        else:566            logger.info("Reading examples from {}".format(file_path))567            data = read_jsonl(file_path)568            # data = [dic for dic in data if len(dic["patch"].split("\n")) <= 20]569            for i in range(len(data)):570                data[i]["idx"] = i571            logger.info(f"Tokenize examples: {file_path}")572            # self.feats = pool.map(self.convert_examples_to_features, \573            #     [(dic, tokenizer, args) for dic in data])574            self.feats = [self.convert_examples_to_features((dic, tokenizer, args)) for dic in data]575            torch.save(self.feats, savep)576 577    def convert_examples_to_features(self, item):578        dic, tokenizer, args = item579        diff, msg = dic["patch"], dic["msg"]580        difflines = diff.split("\n")[1:]        # remove start @@581        difflines = [line for line in difflines if len(line.strip()) > 0]582        map_dic = {"-": 0, "+": 1, " ": 2}583        def f(s):584            if s in map_dic:585                return map_dic[s]586            else:587                return 2588        labels = [f(line[0]) for line in difflines]589        difflines = [line[1:].strip() for line in difflines]590        inputstr = ""591        for label, line in zip(labels, difflines):592            if label == 1:593                inputstr += "<add>" + line594            elif label == 0:595                inputstr += "<del>" + line596            else:597                inputstr += "<keep>" + line598        source_ids = self.encode_remove(tokenizer, inputstr, args)599        target_ids = []600        target_ids.append(tokenizer.msg_id)601        msg = self.encode_remove(tokenizer, dic["msg"], args)602        target_ids.extend(msg)603        source_ids, target_ids = self.pad_assert(source_ids, target_ids, args, tokenizer)604        input_labels = [-100] * len(source_ids)605        return ReviewFeatures(dic["idx"], source_ids, input_labels, target_ids, type="genmsg")606 607 608class InputFeatures(object):609    """A single training/test features for a example."""610 611    def __init__(self, example_id, source_ids, target_ids, url=None):612        self.example_id = example_id613        self.source_ids = source_ids614        self.target_ids = target_ids615        self.url = url616 617 618class ReviewFeatures(object):619    def __init__(self, example_id, source_ids, source_labels, target_ids, type):620        self.example_id = example_id621        self.source_ids = source_ids622        self.source_labels = source_labels623        self.target_ids = target_ids624        assert type in ("label", "line", "genmsg", "daemsg")625        self.type = type626 627class ClsFeatures(object):628    def __init__(self, example_id, source_ids, y):629        self.example_id = example_id630        self.source_ids = source_ids631        self.y = y632 633class ReviewExample(object):634    """A single training/test example."""635 636    def __init__(637        self, idx, oldf, diff, msg, cmtid, max_len, y638    ):639        self.idx = idx      # idx is useless yet640        self.oldf = oldf641        self.diff = diff642        self.msg = msg643        self.cmtid = cmtid644        self.max_len = max_len645        self.y = y646        self.prevlines = []647        self.afterlines = []648        self.lines = []649        self.labels = []650        self.avail = False651        self.input = ""652        self.align_and_clean()653        self.postprocess()654 655    def postprocess(self):656        if not self.avail:657            return658        # Warning: lines is not self.lines659        # lines for rough length estimation660        lines = [source_str.split() for source_str in self.lines]661        inputl = len(lines) # line tag662        inputl += sum(map(len, lines))663        left, right = 0, len(lines)664        while inputl > self.max_len:665            if left % 2 == 0:666                inputl -= len(lines[left]) + 1667                left += 1668            else:669                right -= 1670                inputl -= len(lines[right]) + 1671        lines = lines[left:right]672        self.lines = self.lines[left:right]673        self.labels = self.labels[left:right]674        prevlines = self.prevlines675        afterlines = self.afterlines676        prev_after_len = max(len(prevlines), len(afterlines))677        i = 0678        while inputl < self.max_len and i < prev_after_len:679            if i < len(prevlines):680                newl = inputl + len(prevlines[-1-i].split()) + 1681                if newl > self.max_len:682                    break683                self.lines.insert(0, prevlines[-1-i])684                self.labels.insert(0, -100)685                inputl = newl  # tag686            if i < len(afterlines):687                newl = inputl + len(afterlines[i].split()) + 1688                if newl > self.max_len:689                    break690                self.lines.append(afterlines[i])691                self.labels.append(-100)692                inputl = newl    # tag693            i += 1694        assert inputl <= self.max_len, "Too long inputs."695        assert len(self.lines) == len(self.labels), "Not equal length."696        self.input = "<e0>".join(self.lines)697        self.prevlines, self.lines, self.afterlines = [], [], []698 699    def remove_space_clean(self, line):700        """701            Remove start and end empty chars.702        """703        rep = " \t\r"704        totallen = len(line)705        i = 0706        while i < totallen and line[i] in rep:707            i += 1708        j = totallen - 1709        while j >= 0 and line[j] in rep:710            j -= 1711        line = line[i : j + 1]712        return line713 714    def align_and_clean(self):715        oldflines = self.oldf.split("\n")716        difflines = self.diff.split("\n")717        first_line = difflines[0]718        difflines = difflines[1:]719        difflines = [line for line in difflines if line != r"\ No newline at end of file"]720        regex = r"@@ -(\d+),(\d+) \+(\d+),(\d+) @@"721        matchres = re.match(regex, first_line)722        if matchres:723            startline, rangelen, startpos, endpos = matchres.groups()724            self.avail = True725        else:726            self.avail = False727            return728        startline, rangelen = int(startline) - 1, int(rangelen)729        endline = startline + rangelen730        self.prevlines = oldflines[:startline]731        self.afterlines = oldflines[endline:]732        for line in difflines:733            if line.startswith("-"):734                self.lines.append(line[1:])735                self.labels.append(0)736            elif line.startswith("+"):737                self.lines.append(line[1:])738                self.labels.append(1)739            else:740                self.lines.append(line)741                self.labels.append(2)742        self.prevlines = [self.remove_space_clean(line) for line in self.prevlines]743        self.afterlines = [self.remove_space_clean(line) for line in self.afterlines]744        self.lines = [self.remove_space_clean(line) for line in self.lines]745        self.msg = self.remove_space_clean(self.msg)746        self.prevlines = [line for line in self.prevlines if len(line) > 0]747        self.afterlines = [line for line in self.afterlines if len(line) > 0]748        # print("\n".join(self.prevlines))749        # print("\n\n\n\n")750        # print("\n".join(self.lines))751        # print("\n\n\n\n")752        # print("\n".join(self.afterlines))753        # print("\n\n\n\n")754        assert len(self.lines) == len(self.labels), "Not equal length in align."755        topack = list(756            zip(757                *[758                    (line, label)759                    for line, label in zip(self.lines, self.labels)760                    if len(line) > 0761                ]762            )763        )764        if topack == []:765            self.avail = False766            return767        else:768            self.lines, self.labels = topack769        # tuple->list, convenient for later operation770        self.lines = list(self.lines)771        self.labels = list(self.labels)772 773 774def read_review_examples(filename, data_num=-1, tokenizer=None):775    """Read examples from filename."""776    examples = []777    idx = 0778    with open(filename) as f:779        for line in f:780            try:781                js = json.loads(line.strip())782            except:783                print("Error during reading json data.")784                continue785            maxl = 200786            if "y" not in js:787                js["y"] = 0788            if "msg" in js and len(js["msg"]) > 0:789                js["y"] = 1790            example = ReviewExample(791                        idx=idx,792                        oldf=js["oldf"],793                        diff=js["patch"],794                        msg=js["msg"] if "msg" in js else "",795                        cmtid=js["cmtid"] if "cmtid" in js else "",796                        max_len=maxl,797                        y=js["y"]798                    )799            if example.avail:800                examples.append(example)801                idx += 1802                if idx == data_num:803                    break804            else:805                # print(f"Passing {idx} because of invalid diff.")806                idx += 1 807                if idx == data_num:808                    break809                810    return examples811 812 813def read_jsonl(path):814    data = []815    with open(path) as f:816        for line in f:817            try:818                js = json.loads(line.strip())819            except:820                print("Error during reading json data.")821                continue822            data.append(js)823    return data