shekkari21/codereviewer
0
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