CoolFace
Datasetpublic

taskydata/c4tasky

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes64downloads
inference_c4.py167 linesDownload Raw Back to root
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