MCPcopy Create free account
hub / github.com/InternLM/InternBootcamp / test_repeat

Function test_repeat

verl/tests/test_protocol_on_cpu.py:370–392  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

368
369
370def test_repeat():
371 # Create a DataProto object with some batch and non-tensor data
372 obs = torch.tensor([[1, 2], [3, 4], [5, 6]])
373 labels = ["a", "b", "c"]
374 data = DataProto.from_dict(tensors={"obs": obs}, non_tensors={"labels": labels}, meta_info={"info": "test_info"})
375
376 # Test interleave=True
377 repeated_data_interleave = data.repeat(repeat_times=2, interleave=True)
378 expected_obs_interleave = torch.tensor([[1, 2], [1, 2], [3, 4], [3, 4], [5, 6], [5, 6]])
379 expected_labels_interleave = ["a", "a", "b", "b", "c", "c"]
380
381 assert torch.all(torch.eq(repeated_data_interleave.batch["obs"], expected_obs_interleave))
382 assert (repeated_data_interleave.non_tensor_batch["labels"] == expected_labels_interleave).all()
383 assert repeated_data_interleave.meta_info == {"info": "test_info"}
384
385 # Test interleave=False
386 repeated_data_no_interleave = data.repeat(repeat_times=2, interleave=False)
387 expected_obs_no_interleave = torch.tensor([[1, 2], [3, 4], [5, 6], [1, 2], [3, 4], [5, 6]])
388 expected_labels_no_interleave = ["a", "b", "c", "a", "b", "c"]
389
390 assert torch.all(torch.eq(repeated_data_no_interleave.batch["obs"], expected_obs_no_interleave))
391 assert (repeated_data_no_interleave.non_tensor_batch["labels"] == expected_labels_no_interleave).all()
392 assert repeated_data_no_interleave.meta_info == {"info": "test_info"}
393
394
395def test_dataproto_pad_unpad():

Callers

nothing calls this directly

Calls 2

repeatMethod · 0.80
from_dictMethod · 0.45

Tested by

no test coverage detected