| 368 | |
| 369 | |
| 370 | def 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 | |
| 395 | def test_dataproto_pad_unpad(): |