Sample an image patch. args: im: Image pos: center position of crop sample_sz: size to crop output_sz: size to resize to mode: how to treat image borders: 'replicate' (default), 'inside' or 'inside_major' max_scale_change: maximum allowed scale ch
(im: torch.Tensor, pos: torch.Tensor, sample_sz: torch.Tensor, output_sz: torch.Tensor = None,
mode: str = 'replicate', max_scale_change=None, is_mask=False)
| 53 | |
| 54 | |
| 55 | def sample_patch(im: torch.Tensor, pos: torch.Tensor, sample_sz: torch.Tensor, output_sz: torch.Tensor = None, |
| 56 | mode: str = 'replicate', max_scale_change=None, is_mask=False): |
| 57 | """Sample an image patch. |
| 58 | |
| 59 | args: |
| 60 | im: Image |
| 61 | pos: center position of crop |
| 62 | sample_sz: size to crop |
| 63 | output_sz: size to resize to |
| 64 | mode: how to treat image borders: 'replicate' (default), 'inside' or 'inside_major' |
| 65 | max_scale_change: maximum allowed scale change when using 'inside' and 'inside_major' mode |
| 66 | """ |
| 67 | |
| 68 | # if mode not in ['replicate', 'inside']: |
| 69 | # raise ValueError('Unknown border mode \'{}\'.'.format(mode)) |
| 70 | |
| 71 | # copy and convert |
| 72 | posl = pos.long().clone() |
| 73 | |
| 74 | pad_mode = mode |
| 75 | |
| 76 | # Get new sample size if forced inside the image |
| 77 | if mode == 'inside' or mode == 'inside_major': |
| 78 | pad_mode = 'replicate' |
| 79 | im_sz = torch.Tensor([im.shape[2], im.shape[3]]) |
| 80 | shrink_factor = (sample_sz.float() / im_sz) |
| 81 | if mode == 'inside': |
| 82 | shrink_factor = shrink_factor.max() |
| 83 | elif mode == 'inside_major': |
| 84 | shrink_factor = shrink_factor.min() |
| 85 | shrink_factor.clamp_(min=1, max=max_scale_change) |
| 86 | sample_sz = (sample_sz.float() / shrink_factor).long() |
| 87 | |
| 88 | # Compute pre-downsampling factor |
| 89 | if output_sz is not None: |
| 90 | resize_factor = torch.min(sample_sz.float() / output_sz.float()).item() |
| 91 | df = int(max(int(resize_factor - 0.1), 1)) |
| 92 | else: |
| 93 | df = int(1) |
| 94 | |
| 95 | sz = sample_sz.float() / df # new size |
| 96 | |
| 97 | # Do downsampling |
| 98 | if df > 1: |
| 99 | os = posl % df # offset |
| 100 | posl = (posl - os) / df # new position |
| 101 | im2 = im[..., os[0].item()::df, os[1].item()::df] # downsample |
| 102 | else: |
| 103 | im2 = im |
| 104 | |
| 105 | # compute size to crop |
| 106 | szl = torch.max(sz.round(), torch.Tensor([2])).long() |
| 107 | |
| 108 | # Extract top and bottom coordinates |
| 109 | tl = posl - (szl - 1)/2 |
| 110 | br = posl + szl/2 + 1 |
| 111 | |
| 112 | # Shift the crop to inside |
no outgoing calls
no test coverage detected