EnzymeAD / EnzymeAD/Reactant.jl
Using Global ConcreteRArrays rather than arguments during compilation
- Dominant language
- Julia
- Stars
- 370
- Forks
- 74
- Avg merge
- 18h 47m
- Merged PRs (30d)
- 30
Description
I've been exploring some of the interop between Reactant and JAX and ran into the following error:
```julia
import Reactant as Rx
import Flux as Fx
import CUDA as Cu
import Preferences as Prefs
import PythonCall as Py
const jax = Py.pyimport("jax")
const jnp = Py.pyimport("jax.numpy")
const jax = Py.pyimport("jax")
const jnp = Py.pyimport("jax.numpy")
Rx.set_default_backend("gpu")
batchsize = 128
input_dim = 2
hidden_dim = 128
output_dim = 2
# build a simple model in Flux
model_fx = Fx.Chain(
Fx.Dense(input_dim, hidden_dim, Fx.relu),
Fx.Dense(hidden_dim, hidden_dim, Fx.relu),
Fx.Dense(hidden_dim, hidden_dim, Fx.relu),
Fx.Dense(hidden_dim, hidden_dim, Fx.relu),
Fx.Dense(hidden_dim, output_dim)
) |> Fx.gpu
x_fx = Cu.randn(Float32, input_dim, batchsize)
# convert the model to Reactant, do not compile just yet
model_rx = Rx.to_rarray(model_fx)
x_rx = Rx.to_rarray(x_fx)
result_rx_only = Rx.@jit model_rx(x_rx) # compiling only the julia model works
result_jax_only = Rx.@jit jnp.sum(x_rx) # compiling JAX code with a julia input works
result_rx_jax_mixed = Rx.@jit jnp.sum(model_rx(x_rx)) # compiling both jointly throws `ERROR: conversion to pointer not defined for Reactant.ConcretePJRTArray{Float32, 2, 1}`
```
Stacktrace:
```julia
ERROR: conversion to pointer not defined for Reactant.ConcretePJRTArray{Float32, 2, 1}
Stacktrace:
[1] error(s::String)
@ Base ./error.jl:44
[2] unsafe_convert(::Type{Ptr{Float32}}, a::Reactant.ConcretePJRTArray{Float32, 2, 1})
@ Base ./pointer.jl:66
[3] gemm!(transA::Char, transB::Char, alpha::Float32, A::Reactant.ConcretePJRTArray{…}, B::Reactant.ConcretePJRTArray{…}, beta::Float32, C::Reactant.ConcretePJRTArray{…})
@ LinearAlgebra.BLAS ~/.julia/juliaup/julia-1.12.2+0.x64.linux.gnu/share/julia/stdlib/v1.12/LinearAlgebra/src/blas.jl:1648
[4] gemm_wrapper!(C::Reactant.ConcretePJRTArray{…}, tA::Char, tB::Char, A::Reactant.ConcretePJRTArray{…}, B::Reactant.ConcretePJRTArray{…}, α::Bool, β::Bool)
@ LinearAlgebra ~/.julia/juliaup/julia-1.12.2+0.x64.linux.gnu/share/julia/stdlib/v1.12/LinearAlgebra/src/matmul.jl:808
[5] _syrk_herk_gemm_wrapper!(C::Reactant.ConcretePJRTArray{…}, tA::Char, tB::Char, A::Reactant.ConcretePJRTArray{…}, B::Reactant.ConcretePJRTArray{…}, α::Bool, β::Bool, ::Val{…})
@ LinearAlgebra ~/.julia/juliaup/julia-1.12.2+0.x64.linux.gnu/share/julia/stdlib/v1.12/LinearAlgebra/src/matmul.jl:527
[6] generic_matmatmul_wrapper!(C::Reactant.ConcretePJRTArray{…}, tA::Char, tB::Char, A::Reactant.ConcretePJRTArray{…}, B::Reactant.ConcretePJRTArray{…}, α::Bool, β::Bool, val::Val{…})
@ LinearAlgebra ~/.julia/juliaup/julia-1.12.2+0.x64.linux.gnu/share/julia/stdlib/v1.12/LinearAlgebra/src/matmul.jl:507
[7] _mul!
@ ~/.julia/juliaup/julia-1.12.2+0.x64.linux.gnu/share/julia/stdlib/v1.12/LinearAlgebra/src/matmul.jl:328 [inlined]
[8] mul!
@ ~/.julia/juliaup/julia-1.12.2+0.x64.linux.gnu/share/julia/stdlib/v1.12/LinearAlgebra/src/matmul.jl:297 [inlined]
[9] mul!
@ ~/.julia/juliaup/julia-1.12.2+0.x64.linux.gnu/share/julia/stdlib/v1.12/LinearAlgebra/src/matmul.jl:265 [inlined]
[10] *
@ ~/.julia/juliaup/julia-1.12.2+0.x64.linux.gnu/share/julia/stdlib/v1.12/LinearAlgebra/src/matmul.jl:136 [inlined]
[11] (::Flux.Dense{…})(x::Reactant.ConcretePJRTArray{…})
@ Flux ~/.julia/packages/Flux/WMUyh/src/layers/basic.jl:199
[12] macro expansion
@ ~/.julia/packages/Flux/WMUyh/src/layers/basic.jl:68 [inlined]
[13] _applychain
@ ~/.julia/packages/Flux/WMUyh/src/layers/basic.jl:68 [inlined]
[14] (::Flux.Chain{Tuple{…}})(x::Reactant.ConcretePJRTArray{Float32, 2, 1})
@ Flux ~/.julia/packages/Flux/WMUyh/src/layers/basic.jl:65
[15] top-level scope
@ ~/.julia/packages/Reactant/woboD/src/Compiler.jl:2693
[16] top-level scope
@ REPL:1
Some type information was truncated. Use `show(err)` to see complete types.
```
Contributor guide
No contributing guide indexed for this repository
Assessment
This issue has not been assessed yet.