Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions docs/src/api/api.md
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,8 @@ Reactant.to_rarray
```@docs
ConcreteRArray
ConcreteRNumber
Reactant.TracedEnum
Reactant.ConcreteEnum
```

## Inspect Generated HLO
Expand Down
55 changes: 55 additions & 0 deletions docs/src/tutorials/control-flow.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
117 changes: 117 additions & 0 deletions src/Enums.jl
Original file line number Diff line number Diff line change
@@ -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
3 changes: 2 additions & 1 deletion src/Reactant.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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}}()
Expand Down
114 changes: 114 additions & 0 deletions src/Tracing.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down Expand Up @@ -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
Loading
Loading