fastai--fastai
311 行
12 KiB
Python
311 行
12 KiB
Python
"""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)
|