MCPcopy Create free account
hub / github.com/agentscope-ai/Trinity-RFT / test_queue_buffer

Method test_queue_buffer

tests/buffer/queue_test.py:34–97  ·  view source on GitHub ↗
(self, name, use_priority_queue)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 12

to_storage_configMethod · 0.95
acquireMethod · 0.95
writeMethod · 0.95
readMethod · 0.95
releaseMethod · 0.95
QueueWriterClass · 0.90
QueueReaderClass · 0.90
ExperienceClass · 0.90
sleepMethod · 0.80
startMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected