General N-D windowed convolution. The axis arguments say which axis of each
operand plays which role, so any layout can be described.
Most users want nv_conv1d() / nv_conv2d() / nv_conv3d() instead.
Usage
prim_convolution(
x,
kernel,
input_batch_axis,
input_feature_axis,
input_spatial_axes,
kernel_input_feature_axis,
kernel_output_feature_axis,
kernel_spatial_axes,
output_batch_axis,
output_feature_axis,
output_spatial_axes,
window_strides,
padding,
x_dilation,
kernel_dilation,
feature_group_count = 1L,
batch_group_count = 1L,
precision = "highest"
)Arguments
- x
(
arrayish)
Input, e.g.[batch, channels, *spatial]. Can be any data type.xandkernelmust have the same data type. An R value among them assumes the data type of the others when it is in its data type category, and its default data type when none of them has one.- kernel
(
arrayish)
Kernel, e.g.[out_ch, in_ch/groups, *spatial]. Sharesx's data type – seex.- input_batch_axis, input_feature_axis
(
integer(1))
Batch and feature axis ofx.- input_spatial_axes
(
integer())
Spatial axes ofx.- kernel_input_feature_axis, kernel_output_feature_axis
(
integer(1))
Input and output feature axis ofkernel.- kernel_spatial_axes
(
integer())
Spatial axes ofkernel.- output_batch_axis, output_feature_axis
(
integer(1))
Batch and feature axis of the output.- output_spatial_axes
(
integer())
Spatial axes of the output.- window_strides
(
integer())
Stride per spatial axis.- padding
(
matrix)[n_spatial, 2]of(low, high)padding.- x_dilation, kernel_dilation
(
integer())
Input/kernel dilation.- feature_group_count, batch_group_count
(
integer(1))
Grouping.- precision
(
character(1))
One of"highest","high","default".
Value
(arrayish)
Has the data type x and kernel agreed on. Its shape is given by the
output axis arguments: the batch axis holds x's batch size divided by
batch_group_count, the feature axis kernel's output feature size, and
each spatial axis the number of window positions along it.
StableHLO
Lowers to hlo_convolution(), specified under
convolution.
The axis numbers are converted on the way down.
Examples
# a 1-D convolution in NCW layout: one batch, one channel, width 5,
# convolved with a width-3 kernel, giving 3 window positions
x <- nv_array(1:5, shape = c(1, 1, 5), dtype = "f32")
kernel <- nv_array(c(1, 0, -1), shape = c(1, 1, 3), dtype = "f32")
prim_convolution(
x, kernel,
input_batch_axis = 1L, input_feature_axis = 2L, input_spatial_axes = 3L,
kernel_output_feature_axis = 1L, kernel_input_feature_axis = 2L,
kernel_spatial_axes = 3L,
output_batch_axis = 1L, output_feature_axis = 2L, output_spatial_axes = 3L,
window_strides = 1L, padding = matrix(0L, nrow = 1, ncol = 2),
x_dilation = 1L, kernel_dilation = 1L
)
#> AnvlArray
#> (1,.,.) =
#> -2 -2 -2
#> [ CPUf32{1,1,3} ]