chendl/compositional_test
1
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 