| 380 | |
| 381 | |
| 382 | def test_get(): |
| 383 | obs = torch.randn(3, 10) |
| 384 | act = torch.randn(3, 3) |
| 385 | labels = ["a", ["b"], []] |
| 386 | dataset = tu.get_tensordict({"obs": obs, "act": act, "labels": labels}, non_tensor_dict={"2": 2, "1": 1}) |
| 387 | |
| 388 | # test pop keys |
| 389 | popped_dataset = tu.get_keys(dataset, keys=["obs", "2"]) |
| 390 | |
| 391 | assert popped_dataset.batch_size[0] == 3 |
| 392 | |
| 393 | assert torch.all(torch.eq(popped_dataset["obs"], dataset["obs"])).item() |
| 394 | |
| 395 | assert popped_dataset["2"] == dataset["2"] |
| 396 | |
| 397 | # test pop non-exist key |
| 398 | with pytest.raises(KeyError): |
| 399 | tu.get_keys(dataset, keys=["obs", "3"]) |
| 400 | |
| 401 | # test single pop |
| 402 | # NonTensorData |
| 403 | assert tu.get(dataset, key="2") == 2 |
| 404 | # NonTensorStack |
| 405 | assert tu.get(dataset, key="labels") == ["a", ["b"], []] |
| 406 | # Tensor |
| 407 | assert torch.all(torch.eq(tu.get(dataset, key="obs"), obs)).item() |
| 408 | # Non-exist key |
| 409 | assert tu.get(dataset, key="3", default=3) == 3 |
| 410 | |
| 411 | |
| 412 | def test_repeat(): |