"""Basic image opening/processing functionality Docs: https://docs.fast.ai/vision.core.html.md""" # AUTOGENERATED! DO NOT EDIT! File to edit: ../../nbs/07_vision.core.ipynb. # %% auto #0 __all__ = ['imagenet_stats', 'cifar_stats', 'mnist_stats', 'OpenMask', 'TensorPointCreate', 'to_image', 'load_image', 'image2tensor', 'PILBase', 'PILImage', 'PILImageBW', 'PILMask', 'AddMaskCodes', 'TensorPoint', 'get_annotations', 'TensorBBox', 'LabeledBBox', 'encodes', 'PointScaler', 'BBoxLabeler', 'decodes', 'BILINEAR', 'NEAREST', 'Image', 'ToTensor'] # %% ../../nbs/07_vision.core.ipynb #fc486f58 from ..torch_basics import * from ..data.all import * from PIL import Image try: BILINEAR,NEAREST = Image.Resampling.BILINEAR,Image.Resampling.NEAREST except AttributeError: from PIL.Image import BILINEAR,NEAREST # %% ../../nbs/07_vision.core.ipynb #402826e6 _all_ = ['BILINEAR','NEAREST'] # %% ../../nbs/07_vision.core.ipynb #cfa5917e _all_ = ['Image','ToTensor'] # %% ../../nbs/07_vision.core.ipynb #2d1a48fe @patch def __repr__(x:Image.Image): return "<%s.%s image mode=%s size=%dx%d>" % (x.__class__.__module__, x.__class__.__name__, x.mode, x.size[0], x.size[1]) # %% ../../nbs/07_vision.core.ipynb #084fdf09 imagenet_stats = ([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) cifar_stats = ([0.491, 0.482, 0.447], [0.247, 0.243, 0.261]) mnist_stats = ([0.131], [0.308]) # %% ../../nbs/07_vision.core.ipynb #3215d87d if not hasattr(Image,'_patched'): _old_sz = Image.Image.size.fget @patch(as_prop=True) def size(x:Image.Image): return fastuple(_old_sz(x)) Image._patched = True # %% ../../nbs/07_vision.core.ipynb #9ddb2bba @patch(as_prop=True) def n_px(x: Image.Image): return x.size[0] * x.size[1] # %% ../../nbs/07_vision.core.ipynb #091dbece @patch(as_prop=True) def shape(x: Image.Image): return x.size[1],x.size[0] # %% ../../nbs/07_vision.core.ipynb #599920f9 @patch(as_prop=True) def aspect(x: Image.Image): return x.size[0]/x.size[1] # %% ../../nbs/07_vision.core.ipynb #bb645193 @patch def reshape(x: Image.Image, h, w, resample=0): "`resize` `x` to `(w,h)`" return x.resize((w,h), resample=resample) # %% ../../nbs/07_vision.core.ipynb #e2897993 @patch def to_bytes_format(im:Image.Image, format='png'): "Convert to bytes, default to PNG format" arr = io.BytesIO() im.save(arr, format=format) return arr.getvalue() # %% ../../nbs/07_vision.core.ipynb #a4d5d0ab @patch def to_thumb(self:Image.Image, h, w=None): "Same as `thumbnail`, but uses a copy" if w is None: w=h im = self.copy() im.thumbnail((w,h)) return im # %% ../../nbs/07_vision.core.ipynb #926dd512 @patch def resize_max(x: Image.Image, resample=0, max_px=None, max_h=None, max_w=None): "`resize` `x` to `max_px`, or `max_h`, or `max_w`" h,w = x.shape if max_px and x.n_px>max_px: h,w = fastuple(h,w).mul(math.sqrt(max_px/x.n_px)) if max_h and h>max_h: h,w = (max_h ,max_h*w/h) if max_w and w>max_w: h,w = (max_w*h/w,max_w ) return x.reshape(round(h), round(w), resample=resample) # %% ../../nbs/07_vision.core.ipynb #82644fdc def to_image(x): "Convert a tensor or array to a PIL int8 Image" if isinstance(x,Image.Image): return x if isinstance(x,Tensor): x = to_np(x.permute((1,2,0))) if x.dtype==np.float32: x = (x*255).astype(np.uint8) return Image.fromarray(x, mode=['RGB','CMYK'][x.shape[0]==4]) # %% ../../nbs/07_vision.core.ipynb #9095dd8d def load_image(fn, mode=None): "Open and load a `PIL.Image` and convert to `mode`" im = Image.open(fn) im.load() im = im._new(im.im) return im.convert(mode) if mode else im # %% ../../nbs/07_vision.core.ipynb #fc5b8f8e def image2tensor(img): "Transform image to byte tensor in `c*h*w` dim order." res = tensor(img) if res.dim()==2: res = res.unsqueeze(-1) return res.permute(2,0,1) # %% ../../nbs/07_vision.core.ipynb #a3d83eee class PILBase(Image.Image, metaclass=BypassNewMeta): "Base class for a Pillow `Image` that can show itself and convert to a Tensor" _bypass_type=Image.Image _show_args = {'cmap':'viridis'} _open_args = {'mode': 'RGB'} @classmethod def create(cls, fn:Path|str|Tensor|ndarray|bytes|Image.Image, **kwargs): "Return an Image from `fn`" if isinstance(fn,TensorImage): fn = fn.permute(1,2,0).type(torch.uint8) if isinstance(fn,TensorMask): fn = fn.type(torch.uint8) if isinstance(fn,Tensor): fn = fn.numpy() if isinstance(fn,ndarray): return cls(Image.fromarray(fn)) if isinstance(fn,bytes): fn = io.BytesIO(fn) if isinstance(fn,Image.Image): return cls(fn) return cls(load_image(fn, **merge(cls._open_args, kwargs))) def show(self, ctx=None, **kwargs): "Show image using `merge(self._show_args, kwargs)`" return show_image(self, ctx=ctx, **merge(self._show_args, kwargs)) def __repr__(self): return f'{self.__class__.__name__} mode={self.mode} size={"x".join([str(d) for d in self.size])}' # %% ../../nbs/07_vision.core.ipynb #26ecaeec class PILImage(PILBase): "A RGB Pillow `Image` that can show itself and converts to `TensorImage`" pass # %% ../../nbs/07_vision.core.ipynb #1ed2a133 class PILImageBW(PILImage): "A BW Pillow `Image` that can show itself and converts to `TensorImageBW`" _show_args,_open_args = {'cmap':'Greys'},{'mode': 'L'} # %% ../../nbs/07_vision.core.ipynb #048590dd class PILMask(PILBase): "A Pillow `Image` Mask that can show itself and converts to `TensorMask`" _open_args,_show_args = {'mode':'L'},{'alpha':0.5, 'cmap':'tab20'} # %% ../../nbs/07_vision.core.ipynb #d03dcbca OpenMask = Transform(PILMask.create) OpenMask.loss_func = CrossEntropyLossFlat(axis=1) PILMask.create = OpenMask # %% ../../nbs/07_vision.core.ipynb #7add71f4 class AddMaskCodes(Transform): "Add the code metadata to a `TensorMask`" def __init__(self, codes=None): self.codes = codes if codes is not None: self.vocab,self.c = codes,len(codes) def decodes(self, o:TensorMask): if self.codes is not None: o.codes=self.codes return o # %% ../../nbs/07_vision.core.ipynb #dbe12daf class TensorPoint(TensorBase): "Basic type for points in an image" _show_args = dict(s=10, marker='.', c='r') @classmethod def create(cls, t, img_size=None)->None: "Convert an array or a list of points `t` to a `Tensor`" return cls(tensor(t).view(-1, 2).float(), img_size=img_size) def show(self, ctx=None, **kwargs): if 'figsize' in kwargs: del kwargs['figsize'] x = self.view(-1,2) ctx.scatter(x[:, 0], x[:, 1], **{**self._show_args, **kwargs}) return ctx # %% ../../nbs/07_vision.core.ipynb #ce9e2620 TensorPointCreate = Transform(TensorPoint.create) TensorPointCreate.loss_func = MSELossFlat() TensorPoint.create = TensorPointCreate # %% ../../nbs/07_vision.core.ipynb #5fa0c8cd def get_annotations(fname, prefix=None): "Open a COCO style json in `fname` and returns the lists of filenames (with maybe `prefix`) and labelled bboxes." annot_dict = json.load(open(fname)) id2images, id2bboxes, id2cats = {}, collections.defaultdict(list), collections.defaultdict(list) classes = {o['id']:o['name'] for o in annot_dict['categories']} for o in annot_dict['annotations']: bb = o['bbox'] id2bboxes[o['image_id']].append([bb[0],bb[1], bb[0]+bb[2], bb[1]+bb[3]]) id2cats[o['image_id']].append(classes[o['category_id']]) id2images = {o['id']:ifnone(prefix, '') + o['file_name'] for o in annot_dict['images'] if o['id'] in id2bboxes} ids = list(id2images.keys()) return [id2images[k] for k in ids], [(id2bboxes[k], id2cats[k]) for k in ids] # %% ../../nbs/07_vision.core.ipynb #f247c387 from matplotlib import patches, patheffects # %% ../../nbs/07_vision.core.ipynb #590acb21 def _draw_outline(o, lw): o.set_path_effects([patheffects.Stroke(linewidth=lw, foreground='black'), patheffects.Normal()]) def _draw_rect(ax, b, color='white', text=None, text_size=14, hw=True, rev=False): lx,ly,w,h = b if rev: lx,ly,w,h = ly,lx,h,w if not hw: w,h = w-lx,h-ly patch = ax.add_patch(patches.Rectangle((lx,ly), w, h, fill=False, edgecolor=color, lw=2)) _draw_outline(patch, 4) if text is not None: patch = ax.text(lx,ly, text, verticalalignment='top', color=color, fontsize=text_size, weight='bold') _draw_outline(patch,1) # %% ../../nbs/07_vision.core.ipynb #5a84d0e0 class TensorBBox(TensorPoint): "Basic type for a tensor of bounding boxes in an image" @classmethod def create(cls, x, img_size=None)->None: return cls(tensor(x).view(-1, 4).float(), img_size=img_size) def show(self, ctx=None, **kwargs): x = self.view(-1,4) for b in x: _draw_rect(ctx, b, hw=False, **kwargs) return ctx # %% ../../nbs/07_vision.core.ipynb #9ea16a7a class LabeledBBox(L): "Basic type for a list of bounding boxes in an image" def show(self, ctx=None, **kwargs): for b,l in zip(self.bbox, self.lbl): if l != '#na#': ctx = retain_type(b, self.bbox).show(ctx=ctx, text=l) return ctx bbox,lbl = add_props(lambda i,self: self[i]) # %% ../../nbs/07_vision.core.ipynb #e34450eb PILImage ._tensor_cls = TensorImage PILImageBW._tensor_cls = TensorImageBW PILMask ._tensor_cls = TensorMask # %% ../../nbs/07_vision.core.ipynb #52bcf5ce @ToTensor def encodes(self, o:PILBase): return o._tensor_cls(image2tensor(o)) @ToTensor def encodes(self, o:PILMask): return o._tensor_cls(image2tensor(o)[0]) # %% ../../nbs/07_vision.core.ipynb #e21b461c def _scale_pnts(y, sz, do_scale=True, y_first=False): if y_first: y = y.flip(1) res = y * 2/tensor(sz).float() - 1 if do_scale else y return TensorPoint(res, img_size=sz) def _unscale_pnts(y, sz): return TensorPoint((y+1) * tensor(sz).float()/2, img_size=sz) # %% ../../nbs/07_vision.core.ipynb #41e49f29 class PointScaler(Transform): "Scale a tensor representing points" order = 1 def __init__(self, do_scale=True, y_first=False): self.do_scale,self.y_first = do_scale,y_first def _grab_sz(self, x): self.sz = [x.shape[-1], x.shape[-2]] if isinstance(x, Tensor) else x.size return x def _get_sz(self, x): return getattr(x, 'img_size') if self.sz is None else self.sz def setups(self, dl): res = first(dl.do_item(None), risinstance(TensorPoint)) if res is not None: self.c = res.numel() def encodes(self, x:PILBase|TensorImageBase): return self._grab_sz(x) def decodes(self, x:PILBase|TensorImageBase): return self._grab_sz(x) def encodes(self, x:TensorPoint): return _scale_pnts(x, self._get_sz(x), self.do_scale, self.y_first) def decodes(self, x:TensorPoint): return _unscale_pnts(x.view(-1, 2), self._get_sz(x)) # %% ../../nbs/07_vision.core.ipynb #50e5be45 class BBoxLabeler(Transform): def setups(self, dl): self.vocab = dl.vocab def decode (self, x, **kwargs): self.bbox,self.lbls = None,None return self._call('decodes', x, **kwargs) def decodes(self, x:TensorMultiCategory): self.lbls = [self.vocab[a] for a in x] return x if self.bbox is None else LabeledBBox(self.bbox, self.lbls) def decodes(self, x:TensorBBox): self.bbox = x return self.bbox if self.lbls is None else LabeledBBox(self.bbox, self.lbls) # %% ../../nbs/07_vision.core.ipynb #97deb564 #LabeledBBox can be sent in a tl with MultiCategorize (depending on the order of the tls) but it is already decoded. @MultiCategorize def decodes(self, x:LabeledBBox): return x # %% ../../nbs/07_vision.core.ipynb #dc2bca04 @PointScaler def encodes(self, x:TensorBBox): pnts = self.encodes(cast(x.view(-1,2), TensorPoint)) return cast(pnts.view(-1, 4), TensorBBox) @PointScaler def decodes(self, x:TensorBBox): pnts = self.decodes(cast(x.view(-1,2), TensorPoint)) return cast(pnts.view(-1, 4), TensorBBox)