Transform ndarray image to torch tensor. Parameters ---------- img: numpy.ndarray An ndarray with shape: `(H, W, 3)`. Returns ------- torch.Tensor A tensor with shape: `(3, H, W)`.
(img)
| 73 | |
| 74 | |
| 75 | def im_to_torch(img): |
| 76 | """Transform ndarray image to torch tensor. |
| 77 | |
| 78 | Parameters |
| 79 | ---------- |
| 80 | img: numpy.ndarray |
| 81 | An ndarray with shape: `(H, W, 3)`. |
| 82 | |
| 83 | Returns |
| 84 | ------- |
| 85 | torch.Tensor |
| 86 | A tensor with shape: `(3, H, W)`. |
| 87 | |
| 88 | """ |
| 89 | img = np.transpose(img, (2, 0, 1)) # C*H*W |
| 90 | img = to_torch(img).float() |
| 91 | if img.max() > 1: |
| 92 | img /= 255 |
| 93 | return img |
| 94 | |
| 95 | |
| 96 | def torch_to_im(img): |