diff --git a/test/test_datasets.py b/test/test_datasets.py index 22c14cbc08d..88fb3af343a 100644 --- a/test/test_datasets.py +++ b/test/test_datasets.py @@ -3297,6 +3297,16 @@ def test_splits(self): for left, right, disparity in dataset: datasets_utils.shape_test_for_stereo(left, right, disparity) + def test_disparity_scale(self): + with self.create_dataset(split="train") as (dataset, _): + disparity_path = dataset._disparities[0][0] + disparity_image = PIL.Image.fromarray( + np.full((100, 200), 100, dtype=np.uint16) + ) + disparity_image.save(disparity_path) + disparity = dataset._read_disparity(disparity_path)[0] + np.testing.assert_array_equal(disparity, np.ones((1, 100, 200))) + def test_bad_input(self): with pytest.raises( ValueError, match="Unknown value 'bad' for argument split. Valid values are {'train', 'test'}." diff --git a/torchvision/datasets/_stereo_matching.py b/torchvision/datasets/_stereo_matching.py index bc2236e97b8..5bca961da20 100644 --- a/torchvision/datasets/_stereo_matching.py +++ b/torchvision/datasets/_stereo_matching.py @@ -1104,7 +1104,7 @@ def __init__(self, root: Union[str, Path], split: str = "train", transforms: Opt def _read_disparity(self, file_path: str) -> tuple[np.ndarray, None]: disparity_map = np.asarray(Image.open(file_path), dtype=np.float32) # unsqueeze disparity to (C, H, W) - disparity_map = disparity_map[None, :, :] / 1024.0 + disparity_map = disparity_map[None, :, :] / 100.0 valid_mask = None return disparity_map, valid_mask