File size: 4,657 Bytes
8c9ba62 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 | import os
from typing import List
import ray
import torch
from tests.tools import RayUnittestBaseAsync, get_template_config
from trinity.buffer import get_buffer_reader
from trinity.buffer.pipelines.experience_pipeline import ExperiencePipeline
from trinity.common.config import (
ExperienceBufferConfig,
ExperiencePipelineConfig,
OperatorConfig,
)
from trinity.common.constants import SELECTOR_METRIC
from trinity.common.experience import EID, Experience
BUFFER_FILE_PATH = os.path.join(os.path.dirname(__file__), "test_pipeline_buffer.jsonl")
def get_experiences(task_num: int, repeat_times: int = 1, step_num: int = 1) -> List[Experience]:
"""Generate a list of experiences for testing."""
return [
Experience(
eid=EID(task=i, run=j, step=k),
tokens=torch.zeros((5,)),
prompt_length=4,
reward=j,
logprobs=torch.tensor([0.1]),
)
for i in range(task_num)
for j in range(repeat_times)
for k in range(step_num)
]
class TestExperiencePipeline(RayUnittestBaseAsync):
def setUp(self):
if os.path.exists(BUFFER_FILE_PATH):
os.remove(BUFFER_FILE_PATH)
async def test_experience_pipeline(self):
# test input cache
config = get_template_config()
config.data_processor.experience_pipeline = ExperiencePipelineConfig(
save_input=True,
input_save_path=BUFFER_FILE_PATH,
operators=[
OperatorConfig(
name="reward_filter",
args={"threshold": 0.5},
)
],
)
config.algorithm.algorithm_type = "grpo"
config.algorithm.advantage_fn = (
"grpo" # grpo will add an operator at the end of the pipeline
)
config.buffer.trainer_input.experience_buffer = ExperienceBufferConfig(
name="pipeline_test_experience_buffer",
max_read_timeout=3,
)
config.check_and_update()
pipeline = (
ray.remote(ExperiencePipeline)
.options(name=f"{config.explorer.name}_pipeline")
.remote(config)
)
await pipeline.prepare.remote()
task_num = 8
repeat_times = 4
experiences = get_experiences(task_num=task_num, repeat_times=repeat_times)
metrics = await pipeline.process.remote(experiences)
self.assertEqual(
metrics["experience_pipeline/experience_count"], task_num * (repeat_times - 1)
) # first experience of each task will be filtered out by the reward filter
# tests
reader = get_buffer_reader(config.buffer.trainer_input.experience_buffer)
exps = await reader.read_async(batch_size=task_num * (repeat_times - 1))
self.assertEqual(len(exps), task_num * (repeat_times - 1))
with self.assertRaises(TimeoutError):
await reader.read_async(batch_size=task_num)
with open(config.data_processor.experience_pipeline.input_save_path, "r") as f:
input_data = f.readlines()
self.assertEqual(len(input_data), len(experiences))
async def test_pass_rate_calculation(self) -> None:
config = get_template_config()
config.data_processor.experience_pipeline = ExperiencePipelineConfig(
save_input=True,
input_save_path=BUFFER_FILE_PATH,
operators=[
OperatorConfig(
name="pass_rate_calculator",
)
],
)
config.check_and_update()
config.buffer.trainer_input.experience_buffer.name = "pipeline_test_experience_buffer"
config.buffer.trainer_input.experience_buffer.max_read_timeout = 3
pipeline = (
ray.remote(ExperiencePipeline)
.options(name=f"{config.explorer.name}_pipeline")
.remote(config)
)
await pipeline.prepare.remote()
task_num = 8
repeat_times = 4
experiences = get_experiences(task_num=task_num, repeat_times=repeat_times)
for exp in experiences:
exp.info["task_index"] = {
"taskset_id": 0,
"index": exp.eid.task,
}
metrics = await pipeline.process.remote(experiences)
self.assertIn(SELECTOR_METRIC, metrics)
selector_metrics = metrics[SELECTOR_METRIC]
self.assertEqual(len(selector_metrics), 1)
self.assertEqual(set(selector_metrics[0]["indices"]), set(range(task_num)))
self.assertEqual(selector_metrics[0]["values"], [(repeat_times - 1.0) / 2] * task_num)
|