EnzymeAD / EnzymeAD/Reactant.jl

Using Global ConcreteRArrays rather than arguments during compilation

Open
#1,954 10 comments 0 reactions 0 assignees View on GitHub
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.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.