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

Function test_repeat

verl/tests/test_protocol_v2_on_cpu.py:412–434  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

410
411
412def test_repeat():
413 # Create a DataProto object with some batch and non-tensor data
414 obs = torch.tensor([[1, 2], [3, 4], [5, 6]])
415 labels = ["a", "b", "c"]
416 data = tu.get_tensordict({"obs": obs, "labels": labels}, non_tensor_dict={"info": "test_info"})
417
418 # Test interleave=True
419 repeated_data_interleave = data.repeat_interleave(repeats=2)
420 expected_obs_interleave = torch.tensor([[1, 2], [1, 2], [3, 4], [3, 4], [5, 6], [5, 6]])
421 expected_labels_interleave = ["a", "a", "b", "b", "c", "c"]
422
423 assert torch.all(torch.eq(repeated_data_interleave["obs"], expected_obs_interleave))
424 assert repeated_data_interleave["labels"] == expected_labels_interleave
425 assert repeated_data_interleave["info"] == "test_info"
426
427 # Test interleave=False
428 repeated_data_no_interleave = data.repeat(2)
429 expected_obs_no_interleave = torch.tensor([[1, 2], [3, 4], [5, 6], [1, 2], [3, 4], [5, 6]])
430 expected_labels_no_interleave = ["a", "b", "c", "a", "b", "c"]
431
432 assert torch.all(torch.eq(repeated_data_no_interleave["obs"], expected_obs_no_interleave))
433 assert repeated_data_no_interleave["labels"] == expected_labels_no_interleave
434 assert repeated_data_no_interleave["info"] == "test_info"
435
436
437def test_dataproto_pad_unpad():

Callers

nothing calls this directly

Calls 1

repeatMethod · 0.80

Tested by

no test coverage detected