CoolFace
Apppublic

chendl/compositional_test

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
_test_finetune_rag.py112 linesDownload Raw Back to rag
1import json2import logging3import os4import sys5from pathlib import Path6 7import finetune_rag8 9from transformers.file_utils import is_apex_available10from transformers.testing_utils import (11    TestCasePlus,12    execute_subprocess_async,13    require_ray,14    require_torch_gpu,15    require_torch_multi_gpu,16)17 18 19logging.basicConfig(level=logging.DEBUG)20logger = logging.getLogger()21 22stream_handler = logging.StreamHandler(sys.stdout)23logger.addHandler(stream_handler)24 25 26class RagFinetuneExampleTests(TestCasePlus):27    def _create_dummy_data(self, data_dir):28        os.makedirs(data_dir, exist_ok=True)29        contents = {"source": "What is love ?", "target": "life"}30        n_lines = {"train": 12, "val": 2, "test": 2}31        for split in ["train", "test", "val"]:32            for field in ["source", "target"]:33                content = "\n".join([contents[field]] * n_lines[split])34                with open(os.path.join(data_dir, f"{split}.{field}"), "w") as f:35                    f.write(content)36 37    def _run_finetune(self, gpus: int, distributed_retriever: str = "pytorch"):38        tmp_dir = self.get_auto_remove_tmp_dir()39        output_dir = os.path.join(tmp_dir, "output")40        data_dir = os.path.join(tmp_dir, "data")41        self._create_dummy_data(data_dir=data_dir)42 43        testargs = f"""44                --data_dir {data_dir} \45                --output_dir {output_dir} \46                --model_name_or_path facebook/rag-sequence-base \47                --model_type rag_sequence \48                --do_train \49                --do_predict \50                --n_val -1 \51                --val_check_interval 1.0 \52                --train_batch_size 2 \53                --eval_batch_size 1 \54                --max_source_length 25 \55                --max_target_length 25 \56                --val_max_target_length 25 \57                --test_max_target_length 25 \58                --label_smoothing 0.1 \59                --dropout 0.1 \60                --attention_dropout 0.1 \61                --weight_decay 0.001 \62                --adam_epsilon 1e-08 \63                --max_grad_norm 0.1 \64                --lr_scheduler polynomial \65                --learning_rate 3e-04 \66                --num_train_epochs 1 \67                --warmup_steps 4 \68                --gradient_accumulation_steps 1 \69                --distributed-port 8787 \70                --use_dummy_dataset 1 \71                --distributed_retriever {distributed_retriever} \72            """.split()73 74        if gpus > 0:75            testargs.append(f"--gpus={gpus}")76            if is_apex_available():77                testargs.append("--fp16")78        else:79            testargs.append("--gpus=0")80            testargs.append("--distributed_backend=ddp_cpu")81            testargs.append("--num_processes=2")82 83        cmd = [sys.executable, str(Path(finetune_rag.__file__).resolve())] + testargs84        execute_subprocess_async(cmd, env=self.get_env())85 86        metrics_save_path = os.path.join(output_dir, "metrics.json")87        with open(metrics_save_path) as f:88            result = json.load(f)89        return result90 91    @require_torch_gpu92    def test_finetune_gpu(self):93        result = self._run_finetune(gpus=1)94        self.assertGreaterEqual(result["test"][0]["test_avg_em"], 0.2)95 96    @require_torch_multi_gpu97    def test_finetune_multigpu(self):98        result = self._run_finetune(gpus=2)99        self.assertGreaterEqual(result["test"][0]["test_avg_em"], 0.2)100 101    @require_torch_gpu102    @require_ray103    def test_finetune_gpu_ray_retrieval(self):104        result = self._run_finetune(gpus=1, distributed_retriever="ray")105        self.assertGreaterEqual(result["test"][0]["test_avg_em"], 0.2)106 107    @require_torch_multi_gpu108    @require_ray109    def test_finetune_multigpu_ray_retrieval(self):110        result = self._run_finetune(gpus=1, distributed_retriever="ray")111        self.assertGreaterEqual(result["test"][0]["test_avg_em"], 0.2)112