(self, name, use_priority_queue)
| 32 | ] |
| 33 | ) |
| 34 | async def test_queue_buffer(self, name, use_priority_queue): |
| 35 | config = ExperienceBufferConfig( |
| 36 | name=name, |
| 37 | schema_type="experience", |
| 38 | storage_type=StorageType.QUEUE.value, |
| 39 | max_read_timeout=3, |
| 40 | path=BUFFER_FILE_PATH, |
| 41 | batch_size=self.train_batch_size, |
| 42 | ) |
| 43 | config.replay_buffer.enable = use_priority_queue |
| 44 | config = config.to_storage_config() |
| 45 | writer = QueueWriter(config) |
| 46 | reader = QueueReader(config) |
| 47 | self.assertEqual(await writer.acquire(), 1) |
| 48 | exps = [ |
| 49 | Experience( |
| 50 | tokens=torch.tensor([float(j) for j in range(i + 1)]), |
| 51 | prompt_length=i, |
| 52 | reward=float(i), |
| 53 | logprobs=torch.tensor([0.1]), |
| 54 | ) |
| 55 | for i in range(1, self.put_batch_size + 1) |
| 56 | ] |
| 57 | for exp in exps: |
| 58 | exp.info = {"model_version": 0, "use_count": 0} |
| 59 | for _ in range(self.total_num // self.put_batch_size): |
| 60 | await writer.write(exps) |
| 61 | for _ in range(self.total_num // self.train_batch_size): |
| 62 | exps = await reader.read() |
| 63 | self.assertEqual(len(exps), self.train_batch_size) |
| 64 | print(f"finish read {self.train_batch_size} experience") |
| 65 | exps = [ |
| 66 | Experience( |
| 67 | tokens=torch.tensor([float(j) for j in range(i + 1)]), |
| 68 | reward=float(i), |
| 69 | logprobs=torch.tensor([0.1]), |
| 70 | action_mask=torch.tensor([j % 2 for j in range(i + 1)]), |
| 71 | ) |
| 72 | for i in range(1, self.put_batch_size * 2 + 1) |
| 73 | ] |
| 74 | for exp in exps: |
| 75 | exp.info = {"model_version": 1, "use_count": 0} |
| 76 | await writer.write(exps) |
| 77 | exps = await reader.read(batch_size=self.put_batch_size * 2) |
| 78 | self.assertEqual(len(exps), self.put_batch_size * 2) |
| 79 | |
| 80 | def thread_read(reader, result_queue): |
| 81 | try: |
| 82 | batch = asyncio.run(reader.read()) |
| 83 | result_queue.put(batch) |
| 84 | except StopAsyncIteration as e: |
| 85 | result_queue.put(e) |
| 86 | |
| 87 | result_queue = queue.Queue() |
| 88 | t = threading.Thread(target=thread_read, args=(reader, result_queue)) |
| 89 | t.start() |
| 90 | time.sleep(2) # make sure the thread is waiting for data |
| 91 | self.assertEqual(await writer.release(), 0) |
nothing calls this directly
no test coverage detected