diff --git a/docs/src/api/api.md b/docs/src/api/api.md index 2e25512a57..c61ea4d127 100644 --- a/docs/src/api/api.md +++ b/docs/src/api/api.md @@ -34,6 +34,8 @@ Reactant.to_rarray ```@docs ConcreteRArray ConcreteRNumber +Reactant.TracedEnum +Reactant.ConcreteEnum ``` ## Inspect Generated HLO diff --git a/docs/src/tutorials/control-flow.md b/docs/src/tutorials/control-flow.md index 9889702c42..6e58bcceab 100644 --- a/docs/src/tutorials/control-flow.md +++ b/docs/src/tutorials/control-flow.md @@ -127,6 +127,61 @@ location to write into, so assigning it inside `@trace if` raises an error. Parameterize the field type and initialize it with a traced value, e.g. via `ReactantCore.promote_to_traced`. +### Enum state + +Values created by `@enum` or EnumX's `@enumx` can pass through traced control flow. +A runtime-dependent enum result is a [`Reactant.ConcreteEnum`](@ref), which supports +comparisons with the original enum and conversion back to it after execution. During +compilation, the value is represented by a [`Reactant.TracedEnum`](@ref) holding a traced +integer. Both wrappers are separate from Julia’s `Number` hierarchy. + +```@example control_flow_tutorial +@enum SolverStatus Initial Success + +function choose_status(x) + status = Initial + @trace if sum(x) > 0 + status = Success + end + return status +end + +status = @jit choose_status(Reactant.to_rarray(Float32[1])) +@assert status == Success +@assert SolverStatus(status) === Success +``` + +Enum fields of mutable structs follow the same rule as numeric fields above: +parameterize the field type and make its initial value traced before entering the branch. + +```@example control_flow_tutorial +using ReactantCore: promote_to_traced + +mutable struct SolverState{S} + status::S +end + +function update_status(x) + state = SolverState(promote_to_traced(Initial)) + @trace if sum(x) > 0 + state.status = Success + end + return state.status +end + +@assert @jit(update_status(Reactant.to_rarray(Float32[1]))) == Success +@assert @jit(update_status(Reactant.to_rarray(Float32[-1]))) == Initial +``` + +For enum state carried by `@trace while`, also initialize it with +`promote_to_traced`. To pass an enum as a runtime input, convert it with +`Reactant.to_rarray(value; track_numbers=Number)` before compilation. + +Converted enum fields can also be assigned plain enum values outside compilation, +for example to reset a state object between compiled calls. These assignments keep +the field concrete and available as a runtime input. Enum wrappers behave as scalars +in broadcasting, just like plain enums. + ### Loops In addition to conditional evaluations, [`@trace`](@ref) also supports capturing diff --git a/src/Enums.jl b/src/Enums.jl new file mode 100644 index 0000000000..e9fe811af4 --- /dev/null +++ b/src/Enums.jl @@ -0,0 +1,117 @@ +enum_basetype(::Type{<:Base.Enum{T}}) where {T} = T + +abstract type AbstractReactantEnum{E<:Base.Enum} end + +""" + TracedEnum{E <: Base.Enum} + +An enum of type `E` represented by a traced integer during compilation. Supports enum +comparisons, selection, integer conversion, and scalar broadcasting. Compiled results +are reconstructed as [`ConcreteEnum`](@ref) values. +""" +mutable struct TracedEnum{E<:Base.Enum} <: AbstractReactantEnum{E} + value::TracedRNumber +end + +""" + ConcreteEnum{E <: Base.Enum, N <: AbstractConcreteNumber} + +An enum of type `E` backed by a concrete runtime integer of type `N`. Created by +`to_rarray(enum; track_numbers=Number)` and returned by compiled enum computations. +Supports enum comparisons, integer conversion, conversion back to `E`, hashing, and +scalar broadcasting. Assigning a plain enum to a field of this type creates a concrete +integer of the same runtime type; it does not require an active compilation. +""" +mutable struct ConcreteEnum{E<:Base.Enum,N<:AbstractConcreteNumber} <: + AbstractReactantEnum{E} + value::N +end + +ConcreteEnum{E}(value::N) where {E,N<:AbstractConcreteNumber} = ConcreteEnum{E,N}(value) + +Base.broadcastable(x::AbstractReactantEnum) = Ref(x) + +function ReactantCore.promote_to_traced(x::E) where {E<:Base.Enum} + return TracedEnum{E}(promote_to(TracedRNumber{enum_basetype(E)}, Integer(x))) +end + +_enum_payload(x::AbstractReactantEnum) = getfield(x, :value) +_enum_payload(x::Base.Enum) = Integer(x) + +_payload_integer(v::AbstractConcreteNumber) = Integer(to_number(v)) +_payload_integer(v) = v + +function _payload_to(::Type{T}, v::TracedRNumber) where {T<:Integer} + return promote_to(TracedRNumber{T}, v) +end +_payload_to(::Type{T}, v::AbstractConcreteNumber) where {T<:Integer} = T(to_number(v)) +_payload_to(::Type{T}, v) where {T<:Integer} = T(v) + +function _traced_payload(::Type{I}, x::Base.Enum) where {I} + return promote_to(TracedRNumber{I}, Integer(x)) +end +_traced_payload(::Type{I}, x::TracedEnum) where {I} = _enum_payload(x) + +# Conversions + +Base.Integer(x::AbstractReactantEnum) = _payload_integer(_enum_payload(x)) +(::Type{T})(x::AbstractReactantEnum) where {T<:Integer} = _payload_to(T, _enum_payload(x)) + +# Unlike the host constructor, this does not check that the integer is a valid member. +function (::Type{E})(x::TracedRNumber{<:Integer}) where {E<:Base.Enum} + return TracedEnum{E}(promote_to(TracedRNumber{enum_basetype(E)}, x)) +end + +function Base.convert(::Type{E}, x::ConcreteEnum{E}) where {E<:Base.Enum} + return E(Integer(x)) +end +(::Type{E})(x::ConcreteEnum{E}) where {E<:Base.Enum} = convert(E, x) + +# Concrete enum wrappers have the same equality and hash as the plain enum. +Base.hash(x::ConcreteEnum{E}, h::UInt) where {E} = hash(E(x), h) +Base.hash(x::TracedEnum, h::UInt) = hash(_enum_payload(x), h) + +function Base.convert(::Type{TracedEnum{E}}, x::E) where {E<:Base.Enum} + return ReactantCore.promote_to_traced(x) +end + +function Base.convert(::Type{ConcreteEnum{E,N}}, x::E) where {E<:Base.Enum,N} + return ConcreteEnum{E,N}(convert(N, Integer(x))) +end + +function Base.convert(::Type{ConcreteEnum{E}}, x::E) where {E<:Base.Enum} + return to_rarray(x; track_numbers=Number) +end + +# Comparisons and selection + +for jlop in ( + :(Base.:(==)), + :(Base.:(!=)), + :(Base.:(>=)), + :(Base.:(>)), + :(Base.:(<=)), + :(Base.:(<)), + :(Base.isless), +) + @eval begin + function $(jlop)( + lhs::AbstractReactantEnum{E}, rhs::AbstractReactantEnum{E} + ) where {E} + return $(jlop)(_enum_payload(lhs), _enum_payload(rhs)) + end + function $(jlop)(lhs::AbstractReactantEnum{E}, rhs::E) where {E} + return $(jlop)(_enum_payload(lhs), _enum_payload(rhs)) + end + function $(jlop)(lhs::E, rhs::AbstractReactantEnum{E}) where {E} + return $(jlop)(_enum_payload(lhs), _enum_payload(rhs)) + end + end +end + +function Base.ifelse( + pred::TracedRNumber{Bool}, x::Union{TracedEnum{E},E}, y::Union{TracedEnum{E},E} +) where {E<:Base.Enum} + I = enum_basetype(E) + return TracedEnum{E}(ifelse(pred, _traced_payload(I, x), _traced_payload(I, y))) +end diff --git a/src/Reactant.jl b/src/Reactant.jl index 789f352b27..cbd95f2005 100644 --- a/src/Reactant.jl +++ b/src/Reactant.jl @@ -270,6 +270,7 @@ export StackedBatchDuplicated, StackedBatchDuplicatedNoNeed const TracedType = Union{TracedRArray,TracedRNumber,MissingTracedValue} include("ControlFlow.jl") +include("Enums.jl") include("Tracing.jl") include("compiler/Compiler.jl") @@ -308,7 +309,7 @@ export ConcreteRArray, within_compile @static if VERSION ≥ v"1.11" - @eval $(Expr(:public, :Periodic, :Binomial)) + @eval $(Expr(:public, :Periodic, :Binomial, :TracedEnum, :ConcreteEnum)) end const registry = Ref{Union{Nothing,MLIR.IR.DialectRegistry}}() diff --git a/src/Tracing.jl b/src/Tracing.jl index b5a9061245..45a847560e 100644 --- a/src/Tracing.jl +++ b/src/Tracing.jl @@ -835,6 +835,64 @@ Base.@nospecializeinfer function traced_type_inner( throw(NoFieldMatchError(T, TT2, subTys)) end +Base.@nospecializeinfer function traced_type_inner( + @nospecialize(T::Type{ConcreteEnum{E,N}}), + seen, + @nospecialize(mode::TraceMode), + @nospecialize(track_numbers::Type), + @nospecialize(ndevices), + @nospecialize(runtime) +) where {E,N} + mode == ConcreteToTraced && return TracedEnum{E} + if mode == ArrayToConcrete + N2 = traced_type_inner(N, seen, mode, track_numbers, ndevices, runtime) + return ConcreteEnum{E,N2} + end + return T +end + +Base.@nospecializeinfer function traced_type_inner( + @nospecialize(T::Type{TracedEnum{E}}), + seen, + @nospecialize(mode::TraceMode), + @nospecialize(track_numbers::Type), + @nospecialize(ndevices), + @nospecialize(runtime) +) where {E} + if mode == TracedToConcrete + N = traced_type_inner( + TracedRNumber{enum_basetype(E)}, seen, mode, track_numbers, ndevices, runtime + ) + return ConcreteEnum{E,N} + end + mode == ConcreteToTraced && error("Cannot trace an existing TracedEnum") + return T +end + +Base.@nospecializeinfer function should_track_enum( + @nospecialize(E::Type{<:Base.Enum}), @nospecialize(track_numbers::Type) +) + return E <: track_numbers || enum_basetype(E) <: track_numbers +end + +Base.@nospecializeinfer function traced_type_inner( + @nospecialize(T::Type{<:Base.Enum}), + seen, + @nospecialize(mode::TraceMode), + @nospecialize(track_numbers::Type), + @nospecialize(ndevices), + @nospecialize(runtime) +) + should_track_enum(T, track_numbers) || return T + if mode == ArrayToConcrete + N = traced_type_inner(enum_basetype(T), seen, mode, track_numbers, ndevices, runtime) + return ConcreteEnum{T,N} + elseif mode == NoStopTracedTrack + return TracedEnum{T} + end + return T +end + const traced_type_cache = Dict{Tuple{TraceMode,Type,Any},Dict{Type,Type}}() # function traced_type_generator(world::UInt, source, self, @nospecialize(T::Type), @nospecialize(mode::Type{<:Val}), @nospecialize(track_numbers::Type)) @@ -2475,3 +2533,59 @@ function make_tracer( ) return prev end + +# Both wrappers keep their integer at field 1, so generic struct tracing preserves +# payload paths and aliases while these type mappings select the destination wrapper. +Base.@nospecializeinfer function make_tracer( + seen, + @nospecialize(prev::Base.Enum), + @nospecialize(path), + mode; + @nospecialize(track_numbers::Type = Union{}), + @nospecialize(sharding = Sharding.NoSharding()), + @nospecialize(runtime = nothing), + @nospecialize(device = nothing), + @nospecialize(client = nothing), + kwargs..., +) + if mode == TracedToTypes + push!(path, prev) + return nothing + end + RT = Core.Typeof(prev) + should_track_enum(RT, track_numbers) || return prev + if mode == ArrayToConcrete + runtime isa Val{:PJRT} && return ConcreteEnum{RT}( + ConcretePJRTNumber(Integer(prev); sharding, device, client) + ) + runtime isa Val{:IFRT} && return ConcreteEnum{RT}( + ConcreteIFRTNumber(Integer(prev); sharding, device, client) + ) + error("Unsupported runtime $runtime") + elseif mode == NoStopTracedTrack + # Plain enum branch results need the same constant promotion as numbers. + payload = make_tracer( + seen, + Integer(prev), + append_path(path, 1), + mode; + track_numbers=Number, + sharding, + runtime, + device, + client, + kwargs..., + ) + return TracedEnum{RT}(payload) + elseif mode == TracedToConcrete + throw("Input is not a traced-type: $(RT)") + end + return prev +end + +# Keep wrapper identity, payload aliases, and path handling in the generic struct walker. +Base.@nospecializeinfer function make_tracer( + seen, @nospecialize(prev::AbstractReactantEnum), @nospecialize(path), mode; kwargs... +) + return make_tracer_unknown(seen, prev, path, mode; kwargs...) +end diff --git a/src/compiler/Codegen.jl b/src/compiler/Codegen.jl index 1a773241ba..f3d5b7cd07 100644 --- a/src/compiler/Codegen.jl +++ b/src/compiler/Codegen.jl @@ -24,6 +24,15 @@ end return Base.getfield(obj, field) end +# A path can address the payload inside a `TracedEnum` while the object actually present at +# the enum's position (e.g. an untraced branch counterpart) is still the plain enum; there +# is nothing to descend into, and callers skip untraced targets. +@inline function traced_getfield(@nospecialize(obj::Base.Enum), field) + # Only the synthetic payload field can stand in for a plain enum. + field === 1 && return obj + return Base.getfield(obj, field) +end + @inline function traced_getfield( @nospecialize( obj::AbstractArray{<:Union{ConcretePJRTNumber,ConcreteIFRTNumber,TracedRNumber}} @@ -142,6 +151,39 @@ function traced_setfield_buffer!(runtime::Val, cache_dict, concrete_res, obj, fi ) end +# A captured traced enum must be replaced in its containing field: its typed payload +# cannot accept a concrete number. Other values retain the usual payload write-back. +function traced_setfield_buffer_at_parent!( + runtime, cache_dict, concrete_res, parent, field, payload_field, path +) + obj = traced_getfield(parent, field) + if obj isa Reactant.TracedEnum + if haskey(cache_dict, obj) + concrete_enum = cache_dict[obj] + else + payload = obj.value + if haskey(cache_dict, payload) + concrete_payload = cache_dict[payload] + else + T = Reactant.unwrapped_eltype(payload) + concrete_payload = if runtime isa Val{:PJRT} + ConcretePJRTNumber{T}(concrete_res) + else + ConcreteIFRTNumber{T}(concrete_res) + end + cache_dict[payload] = concrete_payload + end + E = typeof(obj).parameters[1] + concrete_enum = Reactant.ConcreteEnum{E}(concrete_payload) + cache_dict[obj] = concrete_enum + end + return traced_setfield!(parent, field, concrete_enum, path) + end + return traced_setfield_buffer!( + runtime, cache_dict, concrete_res, obj, payload_field, path + ) +end + function traced_setfield_buffer!(::Val, _cache_dict, val, concrete_res, _obj, _field, path) return traced_setfield!(val, :data, concrete_res, path) end @@ -1124,7 +1166,7 @@ function codegen_unflatten!( end path = path[3:end] - for p in path[1:(end - 1)] + for p in path[1:max(0, end - 2)] unflatcode = :(traced_getfield($unflatcode, $(Meta.quot(p)))) end @@ -1148,7 +1190,20 @@ function codegen_unflatten!( concrete_res_name_final = unresharded_arrays_cache[concrete_res_name] end - if length(path) > 0 + if length(path) > 1 + needs_cache_dict = true + unflatcode = quote + traced_setfield_buffer_at_parent!( + $(runtime), + $(cache_dict), + $(concrete_res_name_final), + $(unflatcode), + $(Meta.quot(path[end - 1])), + $(Meta.quot(path[end])), + $(path), + ) + end + elseif length(path) > 0 needs_cache_dict = true # TODO(#2233): we might need to handle sharding here unflatcode = quote @@ -1174,7 +1229,12 @@ function codegen_unflatten!( if needs_cache_dict pushfirst!( unflatten_code, - :($cache_dict = IdDict{Union{TracedRArray,TracedRNumber},$ctypes}()), + :( + $cache_dict = IdDict{ + Union{TracedRArray,TracedRNumber,Reactant.TracedEnum}, + Union{$ctypes,Reactant.ConcreteEnum}, + }() + ), ) end diff --git a/test/Project.toml b/test/Project.toml index 9dc88cf865..c07b54c243 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -9,6 +9,7 @@ DLFP8Types = "f4c16678-4a16-415b-82ef-ed337c5d6c7c" Dates = "ade2ca70-3891-5945-98fb-dc099432e06a" Distributions = "31c24e10-a181-5473-b8eb-7969acd0382f" Enzyme = "7da242da-08ed-463a-9acd-ee780be4f1d9" +EnumX = "4e289a0a-7415-4d19-859d-a7e5c4648b56" ExplicitImports = "7d51a73a-1435-4ff3-83d9-f097790105c7" FFTW = "7a1cc6ca-52ef-59f5-83cd-3a7055c09341" FileCheck = "4e644321-382b-4b05-b0b6-5d23c3d944fb" diff --git a/test/core/enums.jl b/test/core/enums.jl new file mode 100644 index 0000000000..aa64c637e2 --- /dev/null +++ b/test/core/enums.jl @@ -0,0 +1,232 @@ +using EnumX, Reactant, Test +using Reactant: @trace, TracedEnum, ConcreteEnum, ConcreteRNumber + +@enum Fruit apple = 1 banana = 2 cherry = 3 +@enum Small::UInt8 low = 7 high = 200 +@enumx Code Default Success MaxIters + +fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 + +@testset "Enum tracing" begin + @testset "constant results" begin + f_const(u) = (Fruit(1), Code.Success) + @test @jit(f_const(fresh())) == (apple, Code.Success) + end + + @testset "ifelse" begin + f_ifelse(u) = ifelse(sum(u) > 1, Code.Success, Code.MaxIters) + res = @jit f_ifelse(fresh()) + @test res isa ConcreteEnum{Code.T} + @test res == Code.Success + @test Code.Success == res + @test convert(Code.T, res) === Code.Success + @test Code.T(res) === Code.Success + @test Integer(res) === Int32(1) + @test Int(res) === 1 + @test Int32(@jit(f_ifelse(Reactant.to_rarray(Float32[0, 0])))) === + Int32(Code.MaxIters) + + res2 = @jit f_ifelse(fresh()) + @test res2 !== res + @test isequal(res, res2) + @test hash(res) == hash(res2) == hash(Code.Success) + @test length(Set([res, res2])) == 1 + @test Dict(res => 1)[res2] == 1 + end + + @testset "comparisons and conversions inside the kernel" begin + function f_cmp(u) + code = ifelse(sum(u) > 1, Code.Success, Code.MaxIters) + return ( + code == Code.Success, + Code.Success == code, + code != Code.Success, + code < Code.MaxIters, + Code.MaxIters > code, + Int(code), + Integer(code) + Int32(1), + Code.T(Int32(code) + Int32(1)), + ) + end + res = @jit f_cmp(fresh()) + @test res[1] == true + @test res[2] == true + @test res[3] == false + @test res[4] == true + @test res[5] == true + @test res[6] == 1 + @test res[7] == Int32(2) + @test res[8] == Code.MaxIters + end + + @testset "@trace if" begin + function f_two_armed(u, threshold) + code = Code.Default + @trace if sum(u) > threshold + code = Code.Success + else + code = Code.MaxIters + end + return code + end + @test @jit(f_two_armed(fresh(), 1.0f0)) == Code.Success + @test @jit(f_two_armed(fresh(), 3.0f0)) == Code.MaxIters + + function f_one_armed(u, threshold, promote) + code = if promote + Reactant.ReactantCore.promote_to_traced(Code.Default) + else + Code.Default + end + @trace if sum(u) > threshold + code = Code.Success + end + return code + end + for promote in (false, true) + @test @jit(f_one_armed(fresh(), 1.0f0, promote)) == Code.Success + @test @jit(f_one_armed(fresh(), 3.0f0, promote)) == Code.Default + end + end + + @testset "mutable struct field" begin + mutable struct EnumCache{U,C,B} + u::U + code::C + done::B + end + function f_field(u, threshold) + c = EnumCache( + u, + Reactant.ReactantCore.promote_to_traced(Code.Default), + Reactant.ReactantCore.promote_to_traced(false), + ) + @trace if sum(c.u) > threshold + c.code = Code.Success + c.done = true + end + return c.code, c.done + end + @test @jit(f_field(fresh(), 1.0f0)) == (Code.Success, true) + @test @jit(f_field(fresh(), 3.0f0)) == (Code.Default, false) + + function f_update(c, threshold) + @trace if sum(c.u) > threshold + c.code = Code.Success + c.done = true + end + return c.code, c.done + end + for (threshold, expected) in ((1.0f0, Code.Success), (3.0f0, Code.Default)) + c = Reactant.to_rarray( + EnumCache(Float32[1, 1], Code.Default, false); track_numbers=Number + ) + @test @jit(f_update(c, threshold)) == (expected, expected == Code.Success) + @test c.code == expected + end + + reset_cache = Reactant.to_rarray( + EnumCache(Float32[1, 1], Code.Default, false); track_numbers=Number + ) + update = @compile f_update(reset_cache, 3.0f0) + @test update(reset_cache, 3.0f0) == (Code.Default, false) + reset_cache.code = Code.Success + @test reset_cache.code.value isa ConcreteRNumber{Int32} + @test update(reset_cache, 3.0f0) == (Code.Success, false) + reset_cache.code = Code.MaxIters + @test update(reset_cache, 3.0f0) == (Code.MaxIters, false) + + function f_field_untouched(u, threshold) + c = EnumCache(u, Code.Default, Reactant.ReactantCore.promote_to_traced(false)) + @trace if sum(c.u) > threshold + c.done = true + end + return c.code, c.done + end + @test @jit(f_field_untouched(fresh(), 1.0f0)) == (Code.Default, true) + end + + @testset "@trace while carrying an enum" begin + # Loop-carried scalars must already be traced to be written back after the loop, + # the same as for plain numbers. + function f_while(u, threshold) + code = Reactant.ReactantCore.promote_to_traced(Code.Default) + i = Reactant.ReactantCore.promote_to_traced(0) + @trace while (i < 5) & (code == Code.Default) + u = u ./ 2 + i += 1 + code = ifelse(sum(u) < threshold, Code.Success, code) + end + return u, code, i + end + u, code, i = @jit f_while(fresh(), 0.6f0) + @test u ≈ Float32[0.25, 0.25] + @test code == Code.Success + @test i == 2 + u, code, i = @jit f_while(fresh(), 0.0f0) + @test code == Code.Default + @test i == 5 + end + + @testset "non-default base type" begin + f_small(u) = ifelse(sum(u) > 1, high, low) + res = @jit f_small(fresh()) + @test res isa ConcreteEnum{Small} + @test res == high + @test Integer(res) === UInt8(200) + f_small_int(u) = Integer(ifelse(sum(u) > 1, high, low)) + @test @jit(f_small_int(fresh())) isa ConcreteRNumber{UInt8} + @test Small(Reactant.to_rarray(low; track_numbers=Number)) === low + end + + @testset "enum arguments" begin + f_arg(u, fruit) = (fruit == banana, Int(fruit)) + fruit = Reactant.to_rarray(banana; track_numbers=Number) + @test fruit isa ConcreteEnum{Fruit} + @test fruit == banana + @test Reactant.to_rarray(banana) === banana + res = @jit f_arg(fresh(), fruit) + @test res[1] == true + @test res[2] == 2 + end + + @testset "concrete and traced representations" begin + fruit = Reactant.to_rarray(banana; track_numbers=Number) + @test !(fruit isa Number) + @test convert(typeof(fruit), apple) == apple + @test convert(ConcreteEnum{Fruit}, cherry) == cherry + @test convert(typeof(fruit), fruit) === fruit + + function f_representation(fruit) + @assert fruit isa TracedEnum{Fruit} + @assert fruit.value isa Reactant.TracedRNumber{Int32} + return fruit, fruit + end + first, second = @jit f_representation(fruit) + @test first isa ConcreteEnum{Fruit} + @test first.value isa ConcreteRNumber{Int32} + @test first === second + @test first == banana + @test @jit(f_representation(first))[1] == banana + + small = Reactant.to_rarray(low; track_numbers=Number) + converted = convert(typeof(small), high) + @test converted.value isa ConcreteRNumber{UInt8} + @test Small(converted) === high + end + + @testset "scalar broadcasting" begin + fruit = Reactant.to_rarray(banana; track_numbers=Number) + f_broadcast(fruit) = [apple, banana, cherry] .== fruit + @test f_broadcast(fruit) == [false, true, false] + @test @jit(f_broadcast(fruit)) == [false, true, false] + @test @jit(f_broadcast(Reactant.to_rarray(apple; track_numbers=Number))) == + [true, false, false] + end +end + +@testset "enum payload paths" begin + @test Reactant.Compiler.traced_getfield(apple, 1) === apple + @test_throws BoundsError Reactant.Compiler.traced_getfield(apple, 2) + @test_throws Exception Reactant.Compiler.traced_getfield(apple, :unrelated) +end