From b04dbdf7ea603594b61c367916e8aed5f08a6345 Mon Sep 17 00:00:00 2001 From: Jitendra Verma Date: Mon, 11 May 2026 00:21:30 +0530 Subject: [PATCH] fix: transfer GridEmbedding grid to input device (fixes CUDA scalar indexing #125) GridEmbedding built the positional grid using CPU range/meshgrid, then called cat(grid, x) where x may be a CuArray. This caused: ERROR: Scalar indexing is disallowed. Invocation of getindex resulted in scalar indexing of a GPU array. Fix: call Lux.get_device(x)(grid) immediately after building the grid, so the array is moved to the same device as the input before the cat. This is a no-op on CPU and transparently transfers to GPU/Metal/etc. Fixes #125 --- src/layers.jl | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/src/layers.jl b/src/layers.jl index d331af4..1b70a44 100644 --- a/src/layers.jl +++ b/src/layers.jl @@ -336,7 +336,7 @@ function (layer::GridEmbedding)(x::AbstractArray{T, N}, ps, st) where {T, N} grid = meshgrid( map(enumerate(layer.grid_boundaries)) do (i, (min, max)) return range(T(min), T(max); length = size(x, i)) - end... + end..., ) grid = repeat( @@ -344,6 +344,10 @@ function (layer::GridEmbedding)(x::AbstractArray{T, N}, ps, st) where {T, N} ntuple(Returns(1), N - 1)..., size(x, N), ) + + # Move the CPU-built grid to the same device as x (fixes CUDA scalar indexing, #125) + grid = Lux.get_device(x)(grid) + return cat(grid, x; dims = N - 1), st end