| 410 | |
| 411 | |
| 412 | def 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 | |
| 437 | def test_dataproto_pad_unpad(): |