From ce1e50970fc629a18bc0ad91956acbceaecb3234 Mon Sep 17 00:00:00 2001 From: Sasha Mitchell Date: Tue, 29 Sep 2026 03:24:08 +0700 Subject: [PATCH] Round pixels before storing them as bytes. tensor2img_fast of 0.5 stored 127. The other converter rounds that value to 128. --- basicsr/utils/img_util.py | 2 +- tests/test_img_util.py | 38 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 39 insertions(+), 1 deletion(-) create mode 100644 tests/test_img_util.py diff --git a/basicsr/utils/img_util.py b/basicsr/utils/img_util.py index 3a5f1da09..97d794d64 100644 --- a/basicsr/utils/img_util.py +++ b/basicsr/utils/img_util.py @@ -105,7 +105,7 @@ def tensor2img_fast(tensor, rgb2bgr=True, min_max=(0, 1)): """ output = tensor.squeeze(0).detach().clamp_(*min_max).permute(1, 2, 0) output = (output - min_max[0]) / (min_max[1] - min_max[0]) * 255 - output = output.type(torch.uint8).cpu().numpy() + output = output.round().type(torch.uint8).cpu().numpy() if rgb2bgr: output = cv2.cvtColor(output, cv2.COLOR_RGB2BGR) return output diff --git a/tests/test_img_util.py b/tests/test_img_util.py new file mode 100644 index 000000000..15d75af95 --- /dev/null +++ b/tests/test_img_util.py @@ -0,0 +1,38 @@ +import importlib.util +import sys +import types +import unittest +from pathlib import Path + +import torch + + +def _load_img_util(): + if "torchvision.utils" not in sys.modules: + torchvision = types.ModuleType("torchvision") + utils = types.ModuleType("torchvision.utils") + utils.make_grid = lambda *args, **kwargs: None + torchvision.utils = utils + sys.modules["torchvision"] = torchvision + sys.modules["torchvision.utils"] = utils + path = Path(__file__).resolve().parents[1] / "basicsr" / "utils" / "img_util.py" + spec = importlib.util.spec_from_file_location("basicsr_img_util", path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +img_util = _load_img_util() + + +class TestTensor2ImgFast(unittest.TestCase): + def test_a_half_step_is_rounded(self): + pixel = torch.tensor([0.5, 0.0, 1.0]).view(1, 3, 1, 1) + image = img_util.tensor2img_fast(pixel, rgb2bgr=False) + self.assertEqual(int(image[0, 0, 0]), 128) + self.assertEqual(int(image[0, 0, 1]), 0) + self.assertEqual(int(image[0, 0, 2]), 255) + + +if __name__ == "__main__": + unittest.main()