diff --git a/NAMESPACE b/NAMESPACE index 6c202a4a..2dae38fa 100644 --- a/NAMESPACE +++ b/NAMESPACE @@ -176,6 +176,9 @@ S3method(transform_to_tensor,list) S3method(transform_to_tensor,matrix) S3method(transform_vflip,default) S3method(transform_vflip,torch_tensor) +S3method(vision_make_grid,"magick-image") +S3method(vision_make_grid,default) +S3method(vision_make_grid,torch_tensor) export(base_loader) export(batched_nms) export(box_area) diff --git a/NEWS.md b/NEWS.md index 4321e546..2a1f7c25 100644 --- a/NEWS.md +++ b/NEWS.md @@ -39,6 +39,7 @@ ## Bug fixes and improvements +* `vision_make_grid()` now accepts multiple 3D tensors with mixed uint8 and float dtype (#398). * `transform_random_affine()` now accepts a bare number for `shear`. It used to widen `degrees` instead of `shear`, which left the shear range incomplete and made the sampling fail (#390). * `item_transform_rotate()` no longer truncates the rotation angle to a whole number of degrees diff --git a/R/collection-rf100-biology.R b/R/collection-rf100-biology.R index 8c86b27e..ed88c261 100644 --- a/R/collection-rf100-biology.R +++ b/R/collection-rf100-biology.R @@ -104,6 +104,7 @@ rf100_biology_collection <- torch::dataset( size = c(81, 24, 12, 6.4, 1.8, 0.9 , 65.1,17.9,9, 0.3, 0.1, 0.05, 1.4, 2.5, 0.8, 62, 16.8, 9, 19, 5.3, 2.7, 69,9.0, 5.7, 192,55.6, 28, - 45, 13, 4.5) * 1e6 + 45, 13, 4.5) * 1e6, + stringsAsFactors = FALSE ) ) diff --git a/R/collection-rf100-damage.R b/R/collection-rf100-damage.R index 56c197f6..8834a39e 100644 --- a/R/collection-rf100-damage.R +++ b/R/collection-rf100-damage.R @@ -54,6 +54,7 @@ rf100_damage_collection <- torch::dataset( # asbestos "f53af703c4ce594c3950d7e732003f2d", "8e99904ce49e7f0e830735fb22986868", "4fef507d057690d1a55fa043696248cc" ), - size = c(21.5, 5.7, 1.4, 8.5, 2.3, 1.5, 28, 7.8, 4) * 1e6 + size = c(21.5, 5.7, 1.4, 8.5, 2.3, 1.5, 28, 7.8, 4) * 1e6, + stringsAsFactors = FALSE ) ) diff --git a/R/collection-rf100-doc.R b/R/collection-rf100-doc.R index 9e4003de..b5ce289f 100644 --- a/R/collection-rf100-doc.R +++ b/R/collection-rf100-doc.R @@ -133,7 +133,8 @@ rf100_document_collection <- torch::dataset( "ca0b96eb696da512eed629895a433bad" ), - size = c(rep(50, 18), 32, 9, 5, 108, 32, 23) * 1e6 # placeholder; optional + size = c(rep(50, 18), 32, 9, 5, 108, 32, 23) * 1e6, # placeholder; optional + stringsAsFactors = FALSE ), initialize = function( diff --git a/R/dataset-pascal.R b/R/dataset-pascal.R index 57a32860..a129411a 100644 --- a/R/dataset-pascal.R +++ b/R/dataset-pascal.R @@ -112,7 +112,8 @@ pascal_segmentation_dataset <- torch::dataset( "da459979d0c395079b5c75ee67908abb", "6c3384ef61512963050cb5d687e5bf1e", "6cd6e144f989b92b3379bac3b3de84fd"), - size = c("440 MB", "440 MB", "550 MB", "890 MB", "1.3 GB", "1.7 GB", "1.9 GB") + size = c("440 MB", "440 MB", "550 MB", "890 MB", "1.3 GB", "1.7 GB", "1.9 GB"), + stringsAsFactors = FALSE ), classes = pascal_voc_classes(), voc_colormap = c( diff --git a/R/dataset-plankton.R b/R/dataset-plankton.R index b85b7bb4..f0b088b0 100644 --- a/R/dataset-plankton.R +++ b/R/dataset-plankton.R @@ -56,7 +56,8 @@ whoi_small_plankton_dataset <- torch::dataset( "170921030aee26f9676725a0c55a4420", "332ebc8b822058cd98f778099927c50e", "f0747ae16fc7cd6946ea54c3fe1f30b4"), - size = c(217e6, 396e6, 383e6, 112e6) + size = c(217e6, 396e6, 383e6, 112e6), + stringsAsFactors = FALSE ), initialize = function( @@ -171,7 +172,8 @@ whoi_plankton_dataset <- torch::dataset( "0f4d47f240cd9c30a7dd786171fa40ca", "db827a7de8790cdcae67b174c7b8ea5e", "d3181d9ffaed43d0c01f59455924edca"), - size = c(rep(450e6, 4), rep(490e6, 13), rep(450e6, 2)) + size = c(rep(450e6, 4), rep(490e6, 13), rep(450e6, 2)), + stringsAsFactors = FALSE ) ) @@ -201,6 +203,7 @@ whoi_small_coralnet_dataset <- torch::dataset( "f4dd2d2effc1f9c02918e3ee614b85d3", "d66ec691a4c5c63878a9cfff164a6aaf", "7ea146b9b2f7b6cee99092bd44182d06"), - size = c(430e6, rep(380e6, 4), 192e6) + size = c(430e6, rep(380e6, 4), 192e6), + stringsAsFactors = FALSE ) ) diff --git a/R/dataset-vggface2.R b/R/dataset-vggface2.R index 0ca3fe83..02cb745c 100644 --- a/R/dataset-vggface2.R +++ b/R/dataset-vggface2.R @@ -49,7 +49,8 @@ vggface2_dataset <- torch::dataset( "bb7a323824d1004e14e00c23974facd3", "d08b10f12bc9889509364ef56d73c621", "d315386c7e8e166c4f60e27d9cc61acc"), - size = c(37.9e9, 62e6, 2.03e9, 3e6, 335e3) + size = c(37.9e9, 62e6, 2.03e9, 3e6, 335e3), + stringsAsFactors = FALSE ), initialize = function( root = tempdir(), diff --git a/R/item-transforms-geometry.R b/R/item-transforms-geometry.R index 3d7f9346..f6568888 100644 --- a/R/item-transforms-geometry.R +++ b/R/item-transforms-geometry.R @@ -34,12 +34,11 @@ #' #' after <- item_transform_rotate(before, angle = angle) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) #' -#' grid <- vision_make_grid( -#' torch_stack(list(transform_resize(before_plot, c(600, 600)), -#' transform_resize(after_plot, c(600, 600)))), +#' grid <- vision_make_grid(transform_resize(before_plot, c(600, 600)), +#' transform_resize(after_plot, c(600, 600)), #' scale = TRUE #' ) #' tensor_image_browse(grid) @@ -241,14 +240,14 @@ item_transform_rotate.image_with_rotated_box <- function(x, angle, interpolation #' #' after <- item_transform_crop(before, top = 100, left = 200, height = 500, width = 800) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) -#' -#' before_plot <- transform_resize(before_plot, c(2000, 3000)) -#' after_plot <- transform_resize(after_plot, c(2000, 3000)) +#' before_plot <- before %>% +#' draw_bounding_boxes(colors = "blue", width = 10) %>% +#' transform_resize(c(2000, 3000)) +#' after_plot <- after %>% +#' draw_bounding_boxes(colors = "red", width = 10) %>% +#' transform_resize(c(2000, 3000)) #' -#' grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -#' tensor_image_browse(grid) +#' tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) #' } #' #' @family item_unitary_transforms @@ -399,10 +398,10 @@ item_transform_crop.image_with_rotated_box <- function(x, top, left, height, wid #' #' after <- item_transform_hflip(before) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) #' -#' grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) +#' grid <- vision_make_grid(before_plot, after_plot, scale = TRUE) #' tensor_image_browse(grid) #' } #' @@ -502,11 +501,10 @@ item_transform_hflip.image_with_rotated_box <- function(x) { #' #' after <- item_transform_vflip(before) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) #' -#' grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -#' tensor_image_browse(grid) +#' tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) #' } #' #' @family item_unitary_transforms @@ -613,14 +611,13 @@ item_transform_vflip.image_with_rotated_box <- function(x) { #' #' after <- item_transform_center_crop(before, size = c(1500, 1500)) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) #' #' before_plot <- transform_resize(before_plot, c(1500, 1500)) #' after_plot <- transform_resize(after_plot, c(1500, 1500)) #' -#' grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -#' tensor_image_browse(grid) +#' tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) #' } #' #' @family item_unitary_transforms @@ -917,14 +914,13 @@ item_transform_affine.dataset <- function(x, angle = 0, translate = c(0, 0), #' #' after <- item_transform_pad(before, padding = 100) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) #' #' before_plot <- transform_resize(before_plot, c(2000, 3000)) #' after_plot <- transform_resize(after_plot, c(2000, 3000)) #' -#' grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -#' tensor_image_browse(grid) +#' tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) #' } #' #' @family item_unitary_transforms @@ -1087,11 +1083,10 @@ item_transform_pad.image_with_rotated_box <- function(x, padding, fill = 0, padd #' #' after <- item_transform_perspective(before, startpoints, endpoints) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) #' -#' grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -#' tensor_image_browse(grid) +#' tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) #' } #' #' @family item_unitary_transforms diff --git a/R/item-transforms-random-geometry.R b/R/item-transforms-random-geometry.R index 911bb9d7..b6038d80 100644 --- a/R/item-transforms-random-geometry.R +++ b/R/item-transforms-random-geometry.R @@ -25,11 +25,10 @@ #' #' after <- item_transform_random_horizontal_flip(before) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) #' -#' grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -#' tensor_image_browse(grid) +#' tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) #' } #' #' @family item_random_transforms @@ -111,11 +110,10 @@ item_transform_random_horizontal_flip.image_with_rotated_box <- function(x, p = #' #' after <- item_transform_random_vertical_flip(before) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) #' -#' grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), num_rows = 1, scale = TRUE) -#' tensor_image_browse(grid) +#' tensor_image_browse(vision_make_grid(before_plot, after_plot, per_row = 1, scale = TRUE)) #' } #' #' @family item_random_transforms @@ -365,16 +363,15 @@ rescale_box_angle <- function(angle_deg, scale_w, scale_h) { #' #' after <- item_transform_random_crop(before, size = c(800, 1200)) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +#' # the crop will changes the image size, so resize is needed before vision_make_grid +#' before_plot <- before %>% +#' draw_bounding_boxes(colors = "blue", width = 10) %>% +#' transform_resize(c(600, 600)) +#' after_plot <- after %>% +#' draw_bounding_boxes(colors = "red", width = 10) %>% +#' transform_resize(c(600, 600)) #' -#' # the crop changes the image size, so resize before stacking into the grid -#' grid <- vision_make_grid( -#' torch_stack(list(transform_resize(before_plot, c(600, 600)), -#' transform_resize(after_plot, c(600, 600)))), -#' scale = TRUE -#' ) -#' tensor_image_browse(grid) +#' tensor_image_browse(vision_make_grid(before_plot, after_plot)) #' } #' #' @family item_random_transforms @@ -777,11 +774,10 @@ item_transform_random_rotation.image_with_segmentation_mask <- item_transform_ra #' #' after <- item_transform_random_erasing(before) #' -#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +#' before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +#' after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) #' -#' grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -#' tensor_image_browse(grid) +#' tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) #' } #' #' @family item_random_transforms diff --git a/R/models-faster_rcnn.R b/R/models-faster_rcnn.R index 5663e708..48692def 100644 --- a/R/models-faster_rcnn.R +++ b/R/models-faster_rcnn.R @@ -980,7 +980,7 @@ mobilenet_v3_320_fpn_backbone <- function(pretrained = TRUE) { #' batch_normalized <- input$unsqueeze(1) # Add batch dimension (1, 3, H, W) #' #' # ResNet-50 FPN V2 -#' model <- model_fasterrcnn_resnet50_fpn_v2(pretrained = TRUE, , detections_per_img = 5 ) +#' model <- model_fasterrcnn_resnet50_fpn_v2(pretrained = TRUE, detections_per_img = 5 ) #' model$eval() #' torch::with_no_grad({pred <- model(batch_normalized)$detections[[1]]}) #' labels <- coco_classes(as.integer(pred$labels)) diff --git a/R/models-mask_rcnn.R b/R/models-mask_rcnn.R index 5c5c8197..660eefc4 100644 --- a/R/models-mask_rcnn.R +++ b/R/models-mask_rcnn.R @@ -562,7 +562,7 @@ maskrcnn_model_v2 <- torch::nn_module( #' batch <- input$unsqueeze(1) #' #' # Mask R-CNN ResNet-50 FPN -#' model <- model_maskrcnn_resnet50_fpn(pretrained = TRUE, , detections_per_img = 5) +#' model <- model_maskrcnn_resnet50_fpn(pretrained = TRUE, detections_per_img = 5) #' model$eval() #' #' torch::with_no_grad({pred <- model(batch)$detections[[1]]}) diff --git a/R/models-vgg.R b/R/models-vgg.R index 31f0bde7..75c62137 100644 --- a/R/models-vgg.R +++ b/R/models-vgg.R @@ -10,7 +10,7 @@ VGG <- torch::nn_module( torch::nn_linear(4096, 4096), torch::nn_relu(TRUE), torch::nn_dropout(), - torch::nn_linear(4096, num_classes), + torch::nn_linear(4096, num_classes) ) if (init_weights) diff --git a/R/target-transforms-detection.R b/R/target-transforms-detection.R index 17cbaaaa..cef3391b 100644 --- a/R/target-transforms-detection.R +++ b/R/target-transforms-detection.R @@ -277,11 +277,7 @@ target_transform_sahi_crop.object_detection_target <- function(y, sahi_split, mi #' # Rotated boxes (red, drawn as polygons) #' after_plot <- draw_bounding_boxes(img, boxes = rotated_target$boxes, colors = "red", width = 10) #' -#' grid <- vision_make_grid( -#' torch_stack(list(before_plot, after_plot))$to(torch_float32()), -#' scale = TRUE -#' ) -#' tensor_image_browse(grid) +#' tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) #' } #' #' @family target_transforms_detection diff --git a/R/transforms-generics.R b/R/transforms-generics.R index c3c39af8..911305ac 100644 --- a/R/transforms-generics.R +++ b/R/transforms-generics.R @@ -358,7 +358,7 @@ transform_ten_crop <- function(img, size, vertical_flip = FALSE) { #' class(item) <- "image_with_bounding_box" #' draw_bounding_boxes(item, colors = "red") #' }) -#' grid <- vision_make_grid(torch_stack(preview), scale = FALSE, num_rows = 3) +#' grid <- vision_make_grid(torch_stack(preview), scale = FALSE, per_row = 3) #' tensor_image_browse(grid) #' } #' diff --git a/R/vision_utils.R b/R/vision_utils.R index 2e118d14..0a926992 100644 --- a/R/vision_utils.R +++ b/R/vision_utils.R @@ -2,37 +2,96 @@ #' @importFrom torch torch_uint8 NULL +.min_max_scale <- function(x) { + min <- x$min()$item() + max <- x$max()$item() + x$clamp_(min = min, max = max) + x$add_(-min)$div_(max - min + 1e-5) + x +} + #' A simplified version of torchvision.utils.make_grid #' -#' Arranges a batch B of (image) tensors in a grid, with optional padding between -#' images. Expects a 4d mini-batch tensor of shape (B x C x H x W). +#' Arranges images in a grid, with optional padding between images. +#' +#' For `torch_tensor` input, accepts either: +#' - one or more 3D tensors of shape (C x H x W) passed as separate arguments, or +#' - one or more 4D batch tensor of shape (B x C x H x W). +#' +#' For `magick-image` input, arranges frames using `magick::image_montage()`. #' -#' @param tensor tensor of shape (B x C x H x W) to arrange in grid. -#' @param scale whether to normalize (min-max-scale) the input tensor. -#' @param num_rows number of rows making up the grid (default 8). -#' @param padding amount of padding between batch images (default 2). -#' @param pad_value pixel value to use for padding. +#' @param tensor A 4D `torch_tensor` of shape (B x C x H x W), a 3D `torch_tensor` +#' of shape (C x H x W), or a `magick-image` object. +#' @param ... Additional 3D `torch_tensor` objects (when first argument is a 3D +#' tensor), or additional `magick-image` objects (when first argument is a +#' `magick-image`). +#' @param scale whether to normalize (min-max-scale) the input tensor. Only +#' applied for `torch_tensor` input. +#' @param per_row maximum number of images per row (i.e., number of columns); +#' remaining images wrap to the next row. Default 8. +#' @param num_rows Deprecated. Use `per_row` instead. +#' @param padding amount of padding between images in pixels (default 2). +#' @param pad_value pixel value (0–1) to use for padding background. +#' +#' @return a 3D `torch_tensor` of shape +#' \eqn{\approx(C , n\_rows \times H , per\_row \times W)} and of dtype `torch_float()`. #' -#' @return a 3d torch_tensor of shape \eqn{\approx(C , num\_rows \times H , num\_cols \times W)} of all images arranged in a grid. #' @family image display #' @export -vision_make_grid <- function(tensor, - scale = TRUE, - num_rows = 8, - padding = 2, - pad_value = 0) { +vision_make_grid <- function(tensor, ..., scale = TRUE, per_row = 8, padding = 2, pad_value = 0, num_rows=NULL) { + if (!is.null(num_rows)) { + deprecated("'num_rows' is deprecated, use 'per_row' instead.") + per_row <- num_rows + } + dots <- list(...) + if (length(dots) > 0) { + primary_class <- class(tensor)[1] + non_matching <- Filter(function(x) !inherits(x, primary_class), dots) + if (length(non_matching) > 0) + cli_abort(c( + "All arguments in {.arg ...} must be {.cls {primary_class}} objects.", + "x" = "Got {.cls {class(non_matching[[1]])[1]}}." + )) + } + UseMethod("vision_make_grid") +} + +#' @rdname vision_make_grid +#' @export +vision_make_grid.default <- function(tensor, ..., scale = TRUE, per_row = 8, padding = 2, pad_value = 0, num_rows=NULL) { + cli_abort("The provided {.var tensor} class {.cls {class(tensor)}} is not supported by {.fn vision_make_grid}") +} + +#' @rdname vision_make_grid +#' @export +vision_make_grid.torch_tensor <- function(tensor, ..., scale = TRUE, per_row = 8, padding = 2, pad_value = 0, num_rows=NULL) { + extra_tensors <- list(...) + + if (!tensor$ndim %in% c(3L, 4L)) + value_error("tensor must be 3D (C x H x W) or 4D (B x C x H x W)") + + if (length(extra_tensors) > 0) { + all_ndims <- c(tensor$ndim, vapply(extra_tensors, function(x) x$ndim, integer(1))) + if (length(unique(all_ndims)) > 1) + value_error("All tensors must have the same number of dimensions (all 3D or all 4D)") + } - min_max_scale <- function(x) { - min <- x$min()$item() - max <- x$max()$item() - x$clamp_(min = min, max = max) - x$add_(-min)$div_(max - min + 1e-5) - x + to_float_unit <- function(t) { + if (t$dtype == torch::torch_uint8()) t$to(dtype = torch::torch_float32())$div(255) else t } - if(scale) tensor <- min_max_scale(tensor) + float_tensors <- lapply(c(list(tensor), extra_tensors), to_float_unit) + all_tensors <- if (scale) lapply(float_tensors, .min_max_scale) else float_tensors + + + if (tensor$ndim == 3) { + tensor <- torch::torch_stack(all_tensors) + } else { + tensor <- if (length(all_tensors) > 1) torch::torch_cat(all_tensors, dim = 1) else all_tensors[[1]] + } + nmaps <- tensor$size(1) - xmaps <- min(num_rows, nmaps) + xmaps <- min(per_row, nmaps) ymaps <- ceiling(nmaps / xmaps) height <- floor(tensor$size(3) + padding) width <- floor(tensor$size(4) + padding) @@ -62,6 +121,29 @@ vision_make_grid <- function(tensor, grid } +#' @rdname vision_make_grid +#' @export +`vision_make_grid.magick-image` <- function(tensor, ..., scale = TRUE, per_row = 8, padding = 2, pad_value = 0, num_rows=NULL) { + rlang::check_installed("magick") + + imgs <- tensor + extra_imgs <- list(...) + if (length(extra_imgs) > 0) { + imgs <- do.call(c, c(list(imgs), extra_imgs)) + } + + # transform_to_tensor already normalises to float32 [0,1]; apply per-frame + # scaling here so behaviour matches the 3D-tensor path (per-image, not global). + frame_tensors <- lapply(seq_along(imgs), function(i) { + t <- transform_to_tensor(imgs[i]) + if (scale) .min_max_scale(t) else t + }) + + batch <- torch::torch_stack(frame_tensors) + vision_make_grid.torch_tensor(batch, scale = FALSE, per_row = per_row, + padding = padding, pad_value = pad_value) +} + #' Draws bounding boxes on image. #' diff --git a/man/item_transform_center_crop.Rd b/man/item_transform_center_crop.Rd index 80e6bb08..3cacb80f 100644 --- a/man/item_transform_center_crop.Rd +++ b/man/item_transform_center_crop.Rd @@ -39,14 +39,13 @@ class(before) <- c("image_with_bounding_box", "list") after <- item_transform_center_crop(before, size = c(1500, 1500)) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) before_plot <- transform_resize(before_plot, c(1500, 1500)) after_plot <- transform_resize(after_plot, c(1500, 1500)) -grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -tensor_image_browse(grid) +tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) } } diff --git a/man/item_transform_crop.Rd b/man/item_transform_crop.Rd index 33fce70c..c0b0b410 100644 --- a/man/item_transform_crop.Rd +++ b/man/item_transform_crop.Rd @@ -45,14 +45,14 @@ class(before) <- c("image_with_bounding_box", "list") after <- item_transform_crop(before, top = 100, left = 200, height = 500, width = 800) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +before_plot <- before \%>\% + draw_bounding_boxes(colors = "blue", width = 10) \%>\% + transform_resize(c(2000, 3000)) +after_plot <- after \%>\% + draw_bounding_boxes(colors = "red", width = 10) \%>\% + transform_resize(c(2000, 3000)) -before_plot <- transform_resize(before_plot, c(2000, 3000)) -after_plot <- transform_resize(after_plot, c(2000, 3000)) - -grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -tensor_image_browse(grid) +tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) } } diff --git a/man/item_transform_hflip.Rd b/man/item_transform_hflip.Rd index 65b2b11f..878775cd 100644 --- a/man/item_transform_hflip.Rd +++ b/man/item_transform_hflip.Rd @@ -32,10 +32,10 @@ class(before) <- c("image_with_bounding_box", "list") after <- item_transform_hflip(before) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) -grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) +grid <- vision_make_grid(before_plot, after_plot, scale = TRUE) tensor_image_browse(grid) } diff --git a/man/item_transform_pad.Rd b/man/item_transform_pad.Rd index 074927fe..00f8a7aa 100644 --- a/man/item_transform_pad.Rd +++ b/man/item_transform_pad.Rd @@ -59,14 +59,13 @@ class(before) <- c("image_with_bounding_box", "list") after <- item_transform_pad(before, padding = 100) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) before_plot <- transform_resize(before_plot, c(2000, 3000)) after_plot <- transform_resize(after_plot, c(2000, 3000)) -grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -tensor_image_browse(grid) +tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) } } diff --git a/man/item_transform_perspective.Rd b/man/item_transform_perspective.Rd index d41dfc7d..30f3327d 100644 --- a/man/item_transform_perspective.Rd +++ b/man/item_transform_perspective.Rd @@ -62,11 +62,10 @@ endpoints <- list(c(100, 50), c(2700, 0), c(2800, 1800), c(80, 1900)) after <- item_transform_perspective(before, startpoints, endpoints) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) -grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -tensor_image_browse(grid) +tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) } } diff --git a/man/item_transform_random_crop.Rd b/man/item_transform_random_crop.Rd index 98581d04..0fb0e283 100644 --- a/man/item_transform_random_crop.Rd +++ b/man/item_transform_random_crop.Rd @@ -67,16 +67,15 @@ class(before) <- c("image_with_bounding_box", "list") after <- item_transform_random_crop(before, size = c(800, 1200)) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +# the crop will changes the image size, so resize is needed before vision_make_grid +before_plot <- before \%>\% + draw_bounding_boxes(colors = "blue", width = 10) \%>\% + transform_resize(c(600, 600)) +after_plot <- after \%>\% + draw_bounding_boxes(colors = "red", width = 10) \%>\% + transform_resize(c(600, 600)) -# the crop changes the image size, so resize before stacking into the grid -grid <- vision_make_grid( - torch_stack(list(transform_resize(before_plot, c(600, 600)), - transform_resize(after_plot, c(600, 600)))), - scale = TRUE -) -tensor_image_browse(grid) +tensor_image_browse(vision_make_grid(before_plot, after_plot)) } } diff --git a/man/item_transform_random_erasing.Rd b/man/item_transform_random_erasing.Rd index 83d3ceef..df11658f 100644 --- a/man/item_transform_random_erasing.Rd +++ b/man/item_transform_random_erasing.Rd @@ -61,11 +61,10 @@ class(before) <- c("image_with_bounding_box", "list") after <- item_transform_random_erasing(before) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) -grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -tensor_image_browse(grid) +tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) } } diff --git a/man/item_transform_random_horizontal_flip.Rd b/man/item_transform_random_horizontal_flip.Rd index 6661f878..8d0d4e74 100644 --- a/man/item_transform_random_horizontal_flip.Rd +++ b/man/item_transform_random_horizontal_flip.Rd @@ -35,11 +35,10 @@ class(before) <- c("image_with_bounding_box", "list") after <- item_transform_random_horizontal_flip(before) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) -grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -tensor_image_browse(grid) +tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) } } diff --git a/man/item_transform_random_vertical_flip.Rd b/man/item_transform_random_vertical_flip.Rd index 1c9e8e02..4ad77961 100644 --- a/man/item_transform_random_vertical_flip.Rd +++ b/man/item_transform_random_vertical_flip.Rd @@ -35,11 +35,10 @@ class(before) <- c("image_with_bounding_box", "list") after <- item_transform_random_vertical_flip(before) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) -grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), num_rows = 1, scale = TRUE) -tensor_image_browse(grid) +tensor_image_browse(vision_make_grid(before_plot, after_plot, per_row = 1, scale = TRUE)) } } diff --git a/man/item_transform_rotate.Rd b/man/item_transform_rotate.Rd index e6a49bc5..a114f692 100644 --- a/man/item_transform_rotate.Rd +++ b/man/item_transform_rotate.Rd @@ -59,12 +59,11 @@ class(before) <- c("image_with_bounding_box", "list") after <- item_transform_rotate(before, angle = angle) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) -grid <- vision_make_grid( - torch_stack(list(transform_resize(before_plot, c(600, 600)), - transform_resize(after_plot, c(600, 600)))), +grid <- vision_make_grid(transform_resize(before_plot, c(600, 600)), + transform_resize(after_plot, c(600, 600)), scale = TRUE ) tensor_image_browse(grid) diff --git a/man/item_transform_vflip.Rd b/man/item_transform_vflip.Rd index fafc26f7..0089d489 100644 --- a/man/item_transform_vflip.Rd +++ b/man/item_transform_vflip.Rd @@ -32,11 +32,10 @@ class(before) <- c("image_with_bounding_box", "list") after <- item_transform_vflip(before) -before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10)$to(torch_float())$div(255) -after_plot <- draw_bounding_boxes(after, colors = "red", width = 10)$to(torch_float())$div(255) +before_plot <- draw_bounding_boxes(before, colors = "blue", width = 10) +after_plot <- draw_bounding_boxes(after, colors = "red", width = 10) -grid <- vision_make_grid(torch_stack(list(before_plot, after_plot)), scale = TRUE) -tensor_image_browse(grid) +tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) } } diff --git a/man/model_fasterrcnn.Rd b/man/model_fasterrcnn.Rd index 40574ea2..389a2513 100644 --- a/man/model_fasterrcnn.Rd +++ b/man/model_fasterrcnn.Rd @@ -119,7 +119,7 @@ input <- image \%>\% batch_normalized <- input$unsqueeze(1) # Add batch dimension (1, 3, H, W) # ResNet-50 FPN V2 -model <- model_fasterrcnn_resnet50_fpn_v2(pretrained = TRUE, , detections_per_img = 5 ) +model <- model_fasterrcnn_resnet50_fpn_v2(pretrained = TRUE, detections_per_img = 5 ) model$eval() torch::with_no_grad({pred <- model(batch_normalized)$detections[[1]]}) labels <- coco_classes(as.integer(pred$labels)) diff --git a/man/model_maskrcnn.Rd b/man/model_maskrcnn.Rd index e314ad35..a9448a38 100644 --- a/man/model_maskrcnn.Rd +++ b/man/model_maskrcnn.Rd @@ -109,7 +109,7 @@ input <- img \%>\% batch <- input$unsqueeze(1) # Mask R-CNN ResNet-50 FPN -model <- model_maskrcnn_resnet50_fpn(pretrained = TRUE, , detections_per_img = 5) +model <- model_maskrcnn_resnet50_fpn(pretrained = TRUE, detections_per_img = 5) model$eval() torch::with_no_grad({pred <- model(batch)$detections[[1]]}) diff --git a/man/target_transform_rotate.Rd b/man/target_transform_rotate.Rd index fd061bdb..7dbce8ae 100644 --- a/man/target_transform_rotate.Rd +++ b/man/target_transform_rotate.Rd @@ -63,11 +63,7 @@ rotated_target <- target_transform_rotate(target, angle = 4) # Rotated boxes (red, drawn as polygons) after_plot <- draw_bounding_boxes(img, boxes = rotated_target$boxes, colors = "red", width = 10) -grid <- vision_make_grid( - torch_stack(list(before_plot, after_plot))$to(torch_float32()), - scale = TRUE -) -tensor_image_browse(grid) +tensor_image_browse(vision_make_grid(before_plot, after_plot, scale = TRUE)) } } diff --git a/man/transform_sahi_crop.Rd b/man/transform_sahi_crop.Rd index 538a34e2..b01dc173 100644 --- a/man/transform_sahi_crop.Rd +++ b/man/transform_sahi_crop.Rd @@ -100,7 +100,7 @@ preview <- lapply(1:dim(crops)[1], function(i) { class(item) <- "image_with_bounding_box" draw_bounding_boxes(item, colors = "red") }) -grid <- vision_make_grid(torch_stack(preview), scale = FALSE, num_rows = 3) +grid <- vision_make_grid(torch_stack(preview), scale = FALSE, per_row = 3) tensor_image_browse(grid) } diff --git a/man/vision_make_grid.Rd b/man/vision_make_grid.Rd index e69d7ba7..8fa91cf8 100644 --- a/man/vision_make_grid.Rd +++ b/man/vision_make_grid.Rd @@ -2,33 +2,86 @@ % Please edit documentation in R/vision_utils.R \name{vision_make_grid} \alias{vision_make_grid} +\alias{vision_make_grid.default} +\alias{vision_make_grid.torch_tensor} +\alias{vision_make_grid.magick-image} \title{A simplified version of torchvision.utils.make_grid} \usage{ vision_make_grid( tensor, + ..., scale = TRUE, - num_rows = 8, + per_row = 8, padding = 2, - pad_value = 0 + pad_value = 0, + num_rows = NULL +) + +\method{vision_make_grid}{default}( + tensor, + ..., + scale = TRUE, + per_row = 8, + padding = 2, + pad_value = 0, + num_rows = NULL +) + +\method{vision_make_grid}{torch_tensor}( + tensor, + ..., + scale = TRUE, + per_row = 8, + padding = 2, + pad_value = 0, + num_rows = NULL +) + +\method{vision_make_grid}{`magick-image`}( + tensor, + ..., + scale = TRUE, + per_row = 8, + padding = 2, + pad_value = 0, + num_rows = NULL ) } \arguments{ -\item{tensor}{tensor of shape (B x C x H x W) to arrange in grid.} +\item{tensor}{A 4D \code{torch_tensor} of shape (B x C x H x W), a 3D \code{torch_tensor} +of shape (C x H x W), or a \code{magick-image} object.} -\item{scale}{whether to normalize (min-max-scale) the input tensor.} +\item{...}{Additional 3D \code{torch_tensor} objects (when first argument is a 3D +tensor), or additional \code{magick-image} objects (when first argument is a +\code{magick-image}).} -\item{num_rows}{number of rows making up the grid (default 8).} +\item{scale}{whether to normalize (min-max-scale) the input tensor. Only +applied for \code{torch_tensor} input.} -\item{padding}{amount of padding between batch images (default 2).} +\item{per_row}{maximum number of images per row (i.e., number of columns); +remaining images wrap to the next row. Default 8.} -\item{pad_value}{pixel value to use for padding.} +\item{padding}{amount of padding between images in pixels (default 2).} + +\item{pad_value}{pixel value (0–1) to use for padding background.} + +\item{num_rows}{Deprecated. Use \code{per_row} instead.} } \value{ -a 3d torch_tensor of shape \eqn{\approx(C , num\_rows \times H , num\_cols \times W)} of all images arranged in a grid. +a 3D \code{torch_tensor} of shape +\eqn{\approx(C , n\_rows \times H , per\_row \times W)} and of dtype \code{torch_float()}. } \description{ -Arranges a batch B of (image) tensors in a grid, with optional padding between -images. Expects a 4d mini-batch tensor of shape (B x C x H x W). +Arranges images in a grid, with optional padding between images. +} +\details{ +For \code{torch_tensor} input, accepts either: +\itemize{ +\item one or more 3D tensors of shape (C x H x W) passed as separate arguments, or +\item one or more 4D batch tensor of shape (B x C x H x W). +} + +For \code{magick-image} input, arranges frames using \code{magick::image_montage()}. } \seealso{ Other image display: diff --git a/tests/testthat/test-vision-utils.R b/tests/testthat/test-vision-utils.R index ea6d2654..29d41e4f 100644 --- a/tests/testthat/test-vision-utils.R +++ b/tests/testthat/test-vision-utils.R @@ -1,15 +1,75 @@ context("vision-utils") -test_that("vision_make_grid", { - +test_that("vision_make_grid works with 4D batch tensor", { images <- torch::torch_randn(c(4, 3, 16, 16)) + grid <- vision_make_grid(images, per_row = 2, padding = 0) + expect_tensor_shape(grid, c(3, 32, 32)) + expect_equal_to_r(grid$max() - grid$min(), 1, tolerance = 1e-4) +}) + +test_that("vision_make_grid works with multiple 3D tensors in ...", { + imgs <- lapply(1:4, function(i) torch::torch_randn(c(3, 16, 16))) + grid <- vision_make_grid(imgs[[1]], imgs[[2]], imgs[[3]], imgs[[4]], per_row = 2, padding = 0) + expect_tensor_shape(grid, c(3, 32, 32)) + expect_equal_to_r(grid$max() - grid$min(), 1, tolerance = 1e-4) +}) + +test_that("vision_make_grid works with multiple 4D tensors in ...", { + batch1 <- torch::torch_randn(c(2, 3, 16, 16)) + batch2 <- torch::torch_randn(c(2, 3, 16, 16)) + grid <- vision_make_grid(batch1, batch2, per_row = 2, padding = 0) + expect_tensor_shape(grid, c(3, 32, 32)) +}) + +test_that("vision_make_grid normalizes single uint8 4D batch to float [0,1]", { + images <- torch::torch_randint(0L, 256L, size = c(4, 3, 16, 16))$to(torch::torch_uint8()) + grid <- vision_make_grid(images, per_row = 2, padding = 0, scale = FALSE) + expect_tensor_shape(grid, c(3, 32, 32)) + expect_tensor_dtype(grid, torch::torch_float32()) + expect_gte(grid$min()$item(), 0) + expect_lte(grid$max()$item(), 1) +}) - grid <- vision_make_grid(images, num_rows = 2, padding = 0) +test_that("vision_make_grid normalizes mixed uint8/float inputs to float [0,1]", { + img_float <- torch::torch_rand(c(3, 16, 16)) + img_uint8 <- torch::torch_randint(0L, 256L, size = c(3, 16, 16))$to(torch::torch_uint8()) + grid <- vision_make_grid(img_float, img_uint8, per_row = 2, padding = 0, scale = FALSE) + expect_tensor_shape(grid, c(3, 16, 32)) + expect_tensor_dtype(grid, torch::torch_float32()) + expect_gte(grid$min()$item(), 0) + expect_lte(grid$max()$item(), 1) +}) + +test_that("vision_make_grid errors when ... contains non-matching types", { + img <- torch::torch_randn(c(3, 16, 16)) + expect_error(vision_make_grid(img, "not_a_tensor"), "") +}) + +test_that("vision_make_grid errors on mixed 3D/4D tensors in ...", { + t3d <- torch::torch_randn(c(3, 16, 16)) + t4d <- torch::torch_randn(c(2, 3, 16, 16)) + expect_error(vision_make_grid(t3d, t4d), "same number of dimensions") +}) +test_that("vision_make_grid works with magick-image", { + skip_if_not_installed("magick") + imgs <- magick::image_read(rep(system.file("img", "Rlogo.png", package = "png"), 4)) + h <- magick::image_info(imgs[1])$height + w <- magick::image_info(imgs[1])$width + grid <- vision_make_grid(imgs, per_row = 2, padding = 0) + expect_tensor_shape(grid, c(3L, 2L * h, 2L * w)) +}) - expect_equal(grid$size(), c(3, 32, 32)) - expect_equal(as.numeric(grid$max() - grid$min()), 1, tolerance = 1e-4) +test_that("vision_make_grid errors on unsupported type", { + expect_error(vision_make_grid(list(1, 2, 3)), "is not supported by") +}) +test_that("vision_make_grid emits deprecation warning for num_rows", { + images <- torch::torch_randn(c(4, 3, 16, 16)) + expect_warning( + vision_make_grid(images, num_rows = 2), + class = "deprecated" + ) }) test_that("draw_bounding_boxes works", {