A util function for plotting a grid of images. Args: images: (N, H, W, 4/1) array of RGBA images rows: number of rows in the grid cols: number of columns in the grid fill: boolean indicating if the space between images should be filled show_axes: bool
(
images,
rows=None,
cols=None,
fill: bool = True,
show_axes: bool = False,
rgb: bool = True,
)
| 290 | |
| 291 | |
| 292 | def image_grid( |
| 293 | images, |
| 294 | rows=None, |
| 295 | cols=None, |
| 296 | fill: bool = True, |
| 297 | show_axes: bool = False, |
| 298 | rgb: bool = True, |
| 299 | ): |
| 300 | """ |
| 301 | A util function for plotting a grid of images. |
| 302 | Args: |
| 303 | images: (N, H, W, 4/1) array of RGBA images |
| 304 | rows: number of rows in the grid |
| 305 | cols: number of columns in the grid |
| 306 | fill: boolean indicating if the space between images should be filled |
| 307 | show_axes: boolean indicating if the axes of the plots should be visible |
| 308 | rgb: boolean, If True, only RGB channels are plotted. |
| 309 | If False, only the alpha channel is plotted. |
| 310 | Returns: |
| 311 | None |
| 312 | """ |
| 313 | if (rows is None) != (cols is None): |
| 314 | raise ValueError("Specify either both rows and cols or neither.") |
| 315 | |
| 316 | if rows is None: |
| 317 | rows = len(images) |
| 318 | cols = 1 |
| 319 | |
| 320 | gridspec_kw = {"wspace": 0.0, "hspace": 0.0} if fill else {} |
| 321 | fig, axarr = plt.subplots(rows, cols, gridspec_kw=gridspec_kw, figsize=(15, 9)) |
| 322 | bleed = 0 |
| 323 | fig.subplots_adjust(left=bleed, bottom=bleed, right=(1 - bleed), top=(1 - bleed)) |
| 324 | |
| 325 | for ax, im in zip(axarr.ravel(), images): |
| 326 | if rgb: |
| 327 | # only render RGB channels |
| 328 | ax.imshow(im[..., :3]) |
| 329 | else: |
| 330 | if im.shape[-1] == 4: |
| 331 | # only render Alpha channel |
| 332 | ax.imshow(im[..., 3]) |
| 333 | else: # depth only |
| 334 | ax.imshow(im) |
| 335 | if not show_axes: |
| 336 | ax.set_axis_off() |
| 337 | |
| 338 | plt.show() |
| 339 | |
| 340 | class OrderedSet(collections.Set): |
| 341 | def __init__(self, iterable=()): |
nothing calls this directly
no outgoing calls
no test coverage detected