diff --git a/test/test_datasets.py b/test/test_datasets.py index 22c14cbc08d..313bc4d103c 100644 --- a/test/test_datasets.py +++ b/test/test_datasets.py @@ -3596,7 +3596,9 @@ class MyVisionDataset(datasets.VisionDataset): datasets.wrap_dataset_for_transforms_v2(dataset) def test_missing_wrapper(self): - dataset = datasets.FakeData() + # Use an official dataset that is not registered for transforms v2 wrapping. + # Construct via ``__new__`` so the test does not require dataset files. + dataset = datasets.SBU.__new__(datasets.SBU) with pytest.raises(TypeError, match="please open an issue"): datasets.wrap_dataset_for_transforms_v2(dataset) diff --git a/torchvision/tv_tensors/_dataset_wrapper.py b/torchvision/tv_tensors/_dataset_wrapper.py index 23683221f60..84a1d5253a2 100644 --- a/torchvision/tv_tensors/_dataset_wrapper.py +++ b/torchvision/tv_tensors/_dataset_wrapper.py @@ -275,16 +275,40 @@ def classification_wrapper_factory(dataset, target_keys): for dataset_cls in [ - datasets.Caltech256, datasets.CIFAR10, datasets.CIFAR100, - datasets.ImageNet, - datasets.MNIST, + datasets.CLEVRClassification, + datasets.Caltech256, + datasets.Country211, + datasets.DTD, + datasets.DatasetFolder, + datasets.EMNIST, + datasets.EuroSAT, + datasets.FER2013, + datasets.FGVCAircraft, + datasets.FakeData, datasets.FashionMNIST, + datasets.Flowers102, + datasets.Food101, datasets.GTSRB, - datasets.DatasetFolder, + datasets.INaturalist, datasets.ImageFolder, + datasets.ImageNet, datasets.Imagenette, + datasets.KMNIST, + datasets.LFWPeople, + datasets.MNIST, + datasets.Omniglot, + datasets.PCAM, + datasets.Places365, + datasets.QMNIST, + datasets.RenderedSST2, + datasets.SEMEION, + datasets.STL10, + datasets.SUN397, + datasets.SVHN, + datasets.StanfordCars, + datasets.USPS, ]: WRAPPER_FACTORIES.register(dataset_cls)(classification_wrapper_factory)