taskydata/c4tasky
064
1# !pip install -q transformers datasets sentencepiece2import argparse3import gc4import json5import os6 7import datasets8import pandas as pd9import torch10from tqdm import tqdm11from transformers import AutoModelForSequenceClassification, AutoTokenizer12 13TOTAL_NUM_FILES_C4_TRAIN = 102414 15 16def parse_args():17 parser = argparse.ArgumentParser()18 19 parser.add_argument(20 "--start",21 type=int,22 required=True,23 help="Starting file number to download. Valid values: 0 - 1023",24 )25 parser.add_argument(26 "--end",27 type=int,28 required=True,29 help="Ending file number to download. Valid values: 0 - 1023",30 )31 parser.add_argument("--batch_size", type=int, default=16, help="Batch size")32 parser.add_argument(33 "--model_name",34 type=str,35 default="taskydata/deberta-v3-base_10xp3nirstbbflanseuni_10xc4",36 help="Model name",37 )38 parser.add_argument(39 "--local_cache_location",40 type=str,41 default="c4_download",42 help="local cache location from where the dataset will be loaded",43 )44 parser.add_argument(45 "--use_local_cache_location",46 type=bool,47 default=True,48 help="Set True if you want to load the dataset from local cache.",49 )50 parser.add_argument(51 "--clear_dataset_cache",52 type=bool,53 default=False,54 help="Set True if you want to delete the dataset files from the cache after inference.",55 )56 parser.add_argument(57 "--release_memory",58 type=bool,59 default=True,60 help="Set True if you want to release the memory of used variables.",61 )62 63 args = parser.parse_args()64 return args65 66 67def chunks(l, n):68 for i in range(0, len(l), n):69 yield l[i : i + n]70 71 72def batch_tokenize(data, batch_size):73 batches = list(chunks(data, batch_size))74 tokenized_batches = []75 for batch in batches:76 # max_length will automatically be set to the max length of the model (512 for deberta)77 tensor = tokenizer(78 batch,79 return_tensors="pt",80 padding="max_length",81 truncation=True,82 max_length=512,83 )84 tokenized_batches.append(tensor)85 return tokenized_batches, batches86 87 88def batch_inference(data, batch_size=16):89 preds = []90 tokenized_batches, batches = batch_tokenize(data, batch_size)91 for i in tqdm(range(len(batches))):92 with torch.no_grad():93 logits = model(**tokenized_batches[i].to(device)).logits.cpu()94 preds.extend(logits)95 return preds96 97 98if __name__ == "__main__":99 args = parse_args()100 101 tokenizer = AutoTokenizer.from_pretrained(args.model_name)102 model = AutoModelForSequenceClassification.from_pretrained(args.model_name)103 device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")104 model.to(device)105 model.eval()106 107 for file_number in range(args.start, args.end + 1):108 global_id = (109 str(file_number).zfill(5) + "-of-" + str(TOTAL_NUM_FILES_C4_TRAIN).zfill(5)110 )111 c4taskyprobas_path = f"c4taskyprobas_{global_id}.jsonl"112 c4tasky_path = f"c4tasky_{global_id}.jsonl"113 114 if os.path.exists(c4taskyprobas_path) and os.path.exists(c4tasky_path):115 print(f"Already done: {c4taskyprobas_path} & {c4tasky_path}")116 exit()117 118 if args.use_local_cache_location:119 file_name = f"c4-train.{global_id}.json.gz"120 data_files = {"train": f"{args.local_cache_location}/{file_name}"}121 c4 = datasets.load_dataset("json", data_files=data_files, split="train")122 else:123 file_name = f"en/c4-train.{global_id}.json.gz"124 data_files = {"train": file_name}125 c4 = datasets.load_dataset(126 "allenai/c4", data_files=data_files, split="train"127 )128 df = pd.DataFrame(c4, index=None)129 texts = df["text"].to_list()130 preds = batch_inference(texts, batch_size=args.batch_size)131 132 assert len(preds) == len(texts)133 134 # Write two jsonl files:135 # 1) Probas for all of C4136 # 2) Probas + texts for samples predicted as tasky137 df['timestamp'] = df['timestamp'].astype(str)138 with open(c4taskyprobas_path, "w") as f, open(c4tasky_path, "w") as g:139 for i in range(len(preds)):140 predicted_class_id = preds[i].argmax().item()141 pred = model.config.id2label[predicted_class_id]142 tasky_proba = torch.softmax(preds[i], dim=-1)[-1].item()143 f.write(json.dumps({"proba": tasky_proba}) + "\n")144 # If it's tasky, save!145 if int(predicted_class_id) == 1:146 g.write(147 json.dumps(148 {149 "proba": tasky_proba,150 "text": texts[i],151 "timestamp": df["timestamp"][i],152 "url": df["url"][i],153 }154 )155 + "\n"156 )157 # release memory158 if args.release_memory:159 del preds160 del texts161 del df162 gc.collect()163 164 # Delete the processed dataset file from the cache165 if args.clear_dataset_cache:166 os.remove(f"{args.local_cache_location}/{file_name}")167 