For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).
Mojo function
dispatch_rdna_conv2d
def dispatch_rdna_conv2d[input_type: DType, filter_type: DType, output_type: DType, filter_is_fcrs: Bool, maybe_epilogue_func: Optional[def[dtype: DType, rank: Int, width: SIMDLength, alignment: Int = Int(1)](IndexList[rank], SIMD[dtype, width]) capturing thin -> None] = None, has_residual: Bool = False](input: TileTensor[input_type, Storage=input.Storage, address_space=input.address_space, linear_idx_type=input.linear_idx_type], filter: TileTensor[filter_type, Storage=filter.Storage, address_space=filter.address_space, linear_idx_type=filter.linear_idx_type], output: TileTensor[output_type, Storage=output.Storage, address_space=output.address_space, linear_idx_type=output.linear_idx_type], stride: IndexList[Int(2)], dilation: IndexList[Int(2)], symmetric_padding: IndexList[Int(2)], num_groups: Int, ctx: DeviceContext, source_ptr: Pointer[Scalar[output_type], MutAnyOrigin] = Pointer.unsafe_dangling(), beta: Float32 = 0) -> Bool
Try to dispatch Conv2D on RDNA via implicit GEMM (im2col fused into WMMA).
Returns True if the convolution was handled, False if the caller should fall back to another implementation (e.g. MIOpen).
Uses the implicit GEMM kernel when C_in is aligned to BLOCK_K (covers all FLUX VAE shapes), falling back to explicit im2col + matmul otherwise.
When has_residual=True, folds output = conv + beta * source into the
conv epilogue (the RDNA implicit-GEMM/im2col kernels have no native
residual path). source_ptr is NHWC-contiguous, same shape as output:
e.g. ResNet skip connections that the graph compiler fuses into the conv.
Parameters:
- input_type (
DType):DTypeof the input tensor elements. - filter_type (
DType):DTypeof the filter tensor elements. - output_type (
DType):DTypeof the output tensor elements. - filter_is_fcrs (
Bool): True iffilteris laid out as FCRS, False for RSCF. - maybe_epilogue_func (
Optional[def[dtype: DType, rank: Int, width: SIMDLength, alignment: Int = Int(1)](IndexList[rank], SIMD[dtype, width]) capturing thin -> None]): Optional elementwise epilogue applied to each output element (defaults toNone). - has_residual (
Bool): True to fold a scaled residualbeta * source_ptrinto the conv epilogue (defaults toFalse).
Args:
- input (
TileTensor[input_type, Storage=input.Storage, address_space=input.address_space, linear_idx_type=input.linear_idx_type]): Rank-4 NHWC input tensor. - filter (
TileTensor[filter_type, Storage=filter.Storage, address_space=filter.address_space, linear_idx_type=filter.linear_idx_type]): Rank-4 filter tensor in FCRS or RSCF layout perfilter_is_fcrs. - output (
TileTensor[output_type, Storage=output.Storage, address_space=output.address_space, linear_idx_type=output.linear_idx_type]): Rank-4 NHWC output tensor written by the convolution. - stride (
IndexList[Int(2)]): Spatial stride[stride_h, stride_w]; only[1, 1]is supported. - dilation (
IndexList[Int(2)]): Spatial dilation[dilation_h, dilation_w]; only[1, 1]is supported. - symmetric_padding (
IndexList[Int(2)]): Symmetric spatial padding[pad_h, pad_w]applied to the input. - num_groups (
Int): Number of convolution groups; only1is supported. - ctx (
DeviceContext):DeviceContextused to enqueue kernels and synchronize. - source_ptr (
Pointer[Scalar[output_type], MutAnyOrigin]): NHWC-contiguous residual source with the same shape asoutput; read only whenhas_residual=True(defaults to a dangling pointer). - beta (
Float32): Scale factor applied to the residual source whenhas_residual=True(defaults to0.0).
Returns: