From 3fb34c52ce32c8f5623abd9858d8567661e35727 Mon Sep 17 00:00:00 2001 From: ChrisRackauckas-Claude Date: Mon, 31 Aug 2026 13:08:50 -0400 Subject: [PATCH 01/10] Trace Base.Enum values as an enum wrapper around the traced base integer MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Base.Enum values (@enum, EnumX's @enumx, ...) were silently dropped by @trace if/while instead of being traced. Represent a traced enum as TracedEnum{E}, a mutable wrapper holding the traced base integer, per the review direction in #3231 — not as TracedRNumber{E} — so the payload is an ordinary traced number and the wrapper rides the generic struct tracing and result reconstruction. The wrapper carries the enum operations (comparisons, ifelse, integer conversion, conversion back to the plain enum once concrete), and traced_getfield tolerates a plain enum standing at a path that addresses the wrapper payload. Fixes #3231. Co-Authored-By: Chris Rackauckas Co-Authored-By: Claude Agent-Harness: Claude Code 2.1.251 Agent-Model: claude-fable-5 Agent-Session: https://claude.ai/code/session_016LsC6pp9z6s5EABX9DnVjE --- docs/src/api/api.md | 1 + src/Enums.jl | 182 ++++++++++++++++++++++++++++++++++++++++ src/Reactant.jl | 1 + src/compiler/Codegen.jl | 5 ++ test/Project.toml | 1 + test/core/enums.jl | 144 +++++++++++++++++++++++++++++++ 6 files changed, 334 insertions(+) create mode 100644 src/Enums.jl create mode 100644 test/core/enums.jl diff --git a/docs/src/api/api.md b/docs/src/api/api.md index 2e25512a57..7a83852b23 100644 --- a/docs/src/api/api.md +++ b/docs/src/api/api.md @@ -34,6 +34,7 @@ Reactant.to_rarray ```@docs ConcreteRArray ConcreteRNumber +Reactant.TracedEnum ``` ## Inspect Generated HLO diff --git a/src/Enums.jl b/src/Enums.jl new file mode 100644 index 0000000000..adc6524c1f --- /dev/null +++ b/src/Enums.jl @@ -0,0 +1,182 @@ +# `Base.Enum` values (`@enum`, EnumX's `@enumx`, ...) are isbits reinterpretations of an +# integer. A traced enum is an enum-typed wrapper around the traced base integer — not a +# `TracedRNumber` with an enum element type — so the payload is an ordinary traced number +# and the wrapper rides the generic struct tracing and result reconstruction. The wrapper +# is mutable with an untyped slot because the result write-back machinery replaces traced +# payloads with concrete numbers via `setfield!`. + +const enum_basetype = Base.Enums.basetype + +""" + TracedEnum{E <: Base.Enum} + +A value of the enum type `E` carried through a Reactant compilation. `value` holds the +enum's base integer: a `TracedRNumber` while tracing, a concrete number in the result of a +compiled call. Supports the enum operations (`==`, `!=`, ordered comparisons, `ifelse`) +against other `TracedEnum`s and plain `E` values, `Integer`/integer-type conversion, and — +once the payload is concrete — conversion back to `E`. +""" +mutable struct TracedEnum{E<:Base.Enum} + value +end + +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::TracedEnum) = 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::TracedEnum) = _payload_integer(_enum_payload(x)) +(::Type{T})(x::TracedEnum) 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::TracedEnum{E}) where {E<:Base.Enum} + v = _enum_payload(x) + v isa TracedRNumber && + error("cannot convert a traced $E back to the plain enum during tracing") + return E(Integer(_payload_integer(v))) +end +(::Type{E})(x::TracedEnum{E}) where {E<:Base.Enum} = convert(E, x) + +function Base.convert(::Type{TracedEnum{E}}, x::E) where {E<:Base.Enum} + return ReactantCore.promote_to_traced(x) +end + +# Comparisons and selection + +for jlop in ( + :(Base.:(==)), + :(Base.:(!=)), + :(Base.:(>=)), + :(Base.:(>)), + :(Base.:(<=)), + :(Base.:(<)), + :(Base.isless), +) + @eval begin + function $(jlop)(lhs::TracedEnum{E}, rhs::TracedEnum{E}) where {E} + return $(jlop)(_enum_payload(lhs), _enum_payload(rhs)) + end + function $(jlop)(lhs::TracedEnum{E}, rhs::E) where {E} + return $(jlop)(_enum_payload(lhs), _enum_payload(rhs)) + end + function $(jlop)(lhs::E, rhs::TracedEnum{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 + +# Tracing. The wrapper itself is an ordinary struct handled by the generic machinery; only +# the plain `Base.Enum` value needs entry points, and they place the payload at the +# wrapper-relative path (`value` is field 1). + +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 || + mode == NoStopTracedTrack || + mode == TracedTrack || + mode == TracedSetPath + return TracedEnum{T} + end + return T +end + +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 TracedEnum{RT}( + ConcretePJRTNumber(Integer(prev); sharding, device, client) + ) + runtime isa Val{:IFRT} && return TracedEnum{RT}( + ConcreteIFRTNumber(Integer(prev); sharding, device, client) + ) + error("Unsupported runtime $runtime") + elseif mode == NoStopTracedTrack + payload = TracedRNumber{enum_basetype(RT)}( + (append_path(path, 1),), @opcall(constant(Integer(prev))).mlir_data + ) + seen[gensym("enum")] = payload + return TracedEnum{RT}(payload) + elseif mode == TracedToConcrete + throw("Input is not a traced-type: $(RT)") + end + return prev +end + +@inline function to_rarray_internal( + @nospecialize(x::Base.Enum), + @nospecialize(track_numbers::Type), + @nospecialize(sharding), + runtime, + @nospecialize(device), + @nospecialize(client) +) + should_track_enum(typeof(x), track_numbers) || return x + if runtime isa Val{:PJRT} + return TracedEnum{typeof(x)}( + ConcretePJRTNumber(Integer(x); sharding, device, client) + ) + elseif runtime isa Val{:IFRT} + return TracedEnum{typeof(x)}( + ConcreteIFRTNumber(Integer(x); sharding, device, client) + ) + end + return error("Unsupported runtime $runtime") +end diff --git a/src/Reactant.jl b/src/Reactant.jl index 789f352b27..c6aea0c8da 100644 --- a/src/Reactant.jl +++ b/src/Reactant.jl @@ -271,6 +271,7 @@ const TracedType = Union{TracedRArray,TracedRNumber,MissingTracedValue} include("ControlFlow.jl") include("Tracing.jl") +include("Enums.jl") include("compiler/Compiler.jl") diff --git a/src/compiler/Codegen.jl b/src/compiler/Codegen.jl index 0e794b108c..a8a40a2bdd 100644 --- a/src/compiler/Codegen.jl +++ b/src/compiler/Codegen.jl @@ -24,6 +24,11 @@ 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 traced_getfield(@nospecialize(obj::Base.Enum), field) = obj + @inline function traced_getfield( @nospecialize( obj::AbstractArray{<:Union{ConcretePJRTNumber,ConcreteIFRTNumber,TracedRNumber}} 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..4340bda200 --- /dev/null +++ b/test/core/enums.jl @@ -0,0 +1,144 @@ +using EnumX, Reactant, Test +using Reactant: @trace, TracedEnum, 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 TracedEnum{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) + 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) + code = Reactant.ReactantCore.promote_to_traced(Code.Default) + @trace if sum(u) > threshold + code = Code.Success + end + return code + end + @test @jit(f_one_armed(fresh(), 1.0f0)) == Code.Success + @test @jit(f_one_armed(fresh(), 3.0f0)) == Code.Default + 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) + 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 TracedEnum{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 TracedEnum{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 +end From f6422be886fc12a950ad1bf430dc9f1622ec539a Mon Sep 17 00:00:00 2001 From: ChrisRackauckas-Claude Date: Fri, 4 Sep 2026 07:37:26 -0400 Subject: [PATCH 02/10] Hash TracedEnum values consistently with == A wrapper with a concrete payload compares equal to other wrappers and to plain enum values, so it must hash like the enum it converts to; with the default identity hash two compiled results for the same enum ended up as distinct Set members. A traced payload keeps the payload's hash, like a bare TracedRNumber. Co-Authored-By: Chris Rackauckas Co-Authored-By: Claude Agent-Harness: Claude Code 2.1.251 Agent-Model: claude-fable-5 Agent-Session: https://claude.ai/code/session_016LsC6pp9z6s5EABX9DnVjE --- src/Enums.jl | 7 +++++++ test/core/enums.jl | 7 +++++++ 2 files changed, 14 insertions(+) diff --git a/src/Enums.jl b/src/Enums.jl index adc6524c1f..222f4721ca 100644 --- a/src/Enums.jl +++ b/src/Enums.jl @@ -59,6 +59,13 @@ function Base.convert(::Type{E}, x::TracedEnum{E}) where {E<:Base.Enum} end (::Type{E})(x::TracedEnum{E}) where {E<:Base.Enum} = convert(E, x) +# With a concrete payload the wrapper is value-equal to `E`, so it hashes like `E` too. +function Base.hash(x::TracedEnum{E}, h::UInt) where {E<:Base.Enum} + v = _enum_payload(x) + v isa TracedRNumber && return hash(v, h) + return hash(E(Integer(_payload_integer(v))), h) +end + function Base.convert(::Type{TracedEnum{E}}, x::E) where {E<:Base.Enum} return ReactantCore.promote_to_traced(x) end diff --git a/test/core/enums.jl b/test/core/enums.jl index 4340bda200..2932b7a88e 100644 --- a/test/core/enums.jl +++ b/test/core/enums.jl @@ -25,6 +25,13 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 @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 From d4594e2e9bece3172612886ed9174a5f35873981 Mon Sep 17 00:00:00 2001 From: ChrisRackauckas-Claude Date: Fri, 4 Sep 2026 07:39:30 -0400 Subject: [PATCH 03/10] Error when an if branch assigns to an untraced struct field `if_condition` writes branch results back into mutable arguments only when the target is already a Concrete/Traced value; any other location (a struct field holding a plain number or enum) was skipped silently, so the branch assignment was dropped without a message. Now a branch that leaves a different value at such a location errors and tells the user to trace the initial value first. Fields the branches did not touch keep working. Co-Authored-By: Chris Rackauckas Co-Authored-By: Claude Agent-Harness: Claude Code 2.1.251 Agent-Model: claude-fable-5 Agent-Session: https://claude.ai/code/session_016LsC6pp9z6s5EABX9DnVjE --- src/Ops.jl | 24 ++++++++++++++++++++++++ test/core/control_flow.jl | 29 +++++++++++++++++++++++++++++ test/core/enums.jl | 21 +++++++++++++++++++++ 3 files changed, 74 insertions(+) diff --git a/src/Ops.jl b/src/Ops.jl index 696e588719..89198333d7 100644 --- a/src/Ops.jl +++ b/src/Ops.jl @@ -2874,6 +2874,10 @@ end Reactant.TracedUtils.set!( args, path[2:end], MLIR.IR.result(if_compiled, residx) ) + else + check_untraced_branch_state( + target, tb_traced_args, fb_traced_args, path[2:end] + ) end end end @@ -2881,6 +2885,26 @@ end return corrected_traced_results end +# An untraced location in a mutable argument (e.g. a struct field holding a plain number or +# enum) has nothing the `if` result can be written back into, so a branch that assigned it +# a new value would be a silent no-op. +function check_untraced_branch_state(target, tb_traced_args, fb_traced_args, path) + for branch_args in (tb_traced_args, fb_traced_args) + leaf = branch_args + for p in path + leaf = Reactant.Compiler.traced_getfield(leaf, p) + end + leaf === target && continue + error( + "if_condition: a branch assigned a value of type $(typeof(leaf)) to an untraced \ + location holding $(repr(target)) (path $(path)); the assignment cannot be \ + carried out of the branch. Make the initial value traced before the `if`, e.g. \ + with `Reactant.ReactantCore.promote_to_traced`.", + ) + end + return nothing +end + """ case( index::TracedRNumber{<:Integer}, branch_fns::Vector, args...; diff --git a/test/core/control_flow.jl b/test/core/control_flow.jl index 16e5ab6108..b3aaaea917 100644 --- a/test/core/control_flow.jl +++ b/test/core/control_flow.jl @@ -1168,6 +1168,35 @@ end @test simulation.stop_iteration == 3 end +mutable struct PlainFieldCache{U,C,B} + u::U + count::C + done::B +end + +function plain_field_assigned(u, threshold) + c = PlainFieldCache(u, 0, ReactantCore.promote_to_traced(false)) + @trace if sum(u) > threshold + c.count = 1 + c.done = true + end + return c.count, c.done +end + +function plain_field_untouched(u, threshold) + c = PlainFieldCache(u, 0, ReactantCore.promote_to_traced(false)) + @trace if sum(u) > threshold + c.done = true + end + return c.count, c.done +end + +@testset "if: assignment to an untraced struct field" begin + u = Reactant.to_rarray(Float32[1, 1]) + @test_throws "untraced location" @jit plain_field_assigned(u, 1.0f0) + @test @jit(plain_field_untouched(u, 1.0f0)) == (0, true) +end + function ternary_max(x, y) @trace result = x > y ? x : y return result diff --git a/test/core/enums.jl b/test/core/enums.jl index 2932b7a88e..f0b6cb9cbd 100644 --- a/test/core/enums.jl +++ b/test/core/enums.jl @@ -103,6 +103,27 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 end @test @jit(f_field(fresh(), 1.0f0)) == (Code.Success, true) @test @jit(f_field(fresh(), 3.0f0)) == (Code.Default, false) + + # A plain enum in the field has nothing the branch result can be written into, the + # same as a plain number; the assignment must error rather than be dropped. + function f_field_plain(u, threshold) + c = EnumCache(u, 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_throws "untraced location" @jit f_field_plain(fresh(), 1.0f0) + + 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 From 71763c1d0e4733b771fb483082f5c94908478e3d Mon Sep 17 00:00:00 2001 From: ChrisRackauckas-Claude Date: Sat, 5 Sep 2026 06:15:25 -0400 Subject: [PATCH 04/10] Complete the documented enum tracing interface Cover unpromoted local enum state and mutation of enum fields converted with to_rarray, and document the required field and loop initialization. Declare TracedEnum public and avoid the private Base.Enums.basetype helper. Co-Authored-By: Chris Rackauckas Co-Authored-By: Codex Agent-Harness: Codex CLI 0.153.4 Agent-Model: gpt-6-astra Agent-Session: local session 01a070df-69eb-7e00-b6bc-a09c892048ab --- docs/src/tutorials/control-flow.md | 49 ++++++++++++++++++++++++++++++ src/Enums.jl | 2 +- src/Reactant.jl | 2 +- test/core/enums.jl | 29 +++++++++++++++--- 4 files changed, 76 insertions(+), 6 deletions(-) diff --git a/docs/src/tutorials/control-flow.md b/docs/src/tutorials/control-flow.md index e906269ced..82919efb09 100644 --- a/docs/src/tutorials/control-flow.md +++ b/docs/src/tutorials/control-flow.md @@ -120,6 +120,55 @@ In our simple example, the condition is passed directly as an argument but the same mechanism is applied to conditions which are computed from within a function from traced arguments, leading to a traced condition. +### Enum state + +Values created by `@enum` or EnumX's `@enumx` can pass through traced control flow. +A runtime-dependent enum result is a [`Reactant.TracedEnum`](@ref), which supports +comparisons with the original enum and conversion back to it after execution. + +```@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 +``` + +For mutable struct fields, parameterize the field type and make its initial value +traced before entering the branch. This requirement also applies to numeric fields. +Assigning a different value to an untraced field inside `@trace if` raises an error. + +```@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. + ### Loops In addition to conditional evaluations, [`@trace`](@ref) also supports capturing diff --git a/src/Enums.jl b/src/Enums.jl index 222f4721ca..79cee4bfd6 100644 --- a/src/Enums.jl +++ b/src/Enums.jl @@ -5,7 +5,7 @@ # is mutable with an untyped slot because the result write-back machinery replaces traced # payloads with concrete numbers via `setfield!`. -const enum_basetype = Base.Enums.basetype +enum_basetype(::Type{<:Base.Enum{T}}) where {T} = T """ TracedEnum{E <: Base.Enum} diff --git a/src/Reactant.jl b/src/Reactant.jl index c6aea0c8da..095c8388e4 100644 --- a/src/Reactant.jl +++ b/src/Reactant.jl @@ -309,7 +309,7 @@ export ConcreteRArray, within_compile @static if VERSION ≥ v"1.11" - @eval $(Expr(:public, :Periodic, :Binomial)) + @eval $(Expr(:public, :Periodic, :Binomial, :TracedEnum)) end const registry = Ref{Union{Nothing,MLIR.IR.DialectRegistry}}() diff --git a/test/core/enums.jl b/test/core/enums.jl index f0b6cb9cbd..dbb0bbcaa8 100644 --- a/test/core/enums.jl +++ b/test/core/enums.jl @@ -72,15 +72,21 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 @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) - code = Reactant.ReactantCore.promote_to_traced(Code.Default) + 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 - @test @jit(f_one_armed(fresh(), 1.0f0)) == Code.Success - @test @jit(f_one_armed(fresh(), 3.0f0)) == Code.Default + 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 @@ -104,6 +110,21 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 @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 + # A plain enum in the field has nothing the branch result can be written into, the # same as a plain number; the assignment must error rather than be dropped. function f_field_plain(u, threshold) From dbf10279c002e09b1ebc08b2ddd6c05870df358a Mon Sep 17 00:00:00 2001 From: ChrisRackauckas-Claude Date: Fri, 11 Sep 2026 07:03:37 -0400 Subject: [PATCH 05/10] Split out untraced-field diagnostic and drop redundant enum to_rarray Per review on #3232: the if_condition check for assignments to untraced mutable fields applies to ordinary numbers too, so it moves to its own PR. The enum-specific to_rarray_internal is dropped: the generic fallback already reaches the new make_tracer(ArrayToConcrete) method. Co-Authored-By: Chris Rackauckas Co-Authored-By: Devin <158243242+devin-ai-integration[bot]@users.noreply.github.com> Agent-Harness: Devin CLI Agent-Model: SWE-2 High Agent-Session: local CLI session (no shareable URL) --- docs/src/tutorials/control-flow.md | 5 +++-- src/Enums.jl | 21 --------------------- src/Ops.jl | 24 ------------------------ test/core/control_flow.jl | 29 ----------------------------- test/core/enums.jl | 12 ------------ 5 files changed, 3 insertions(+), 88 deletions(-) diff --git a/docs/src/tutorials/control-flow.md b/docs/src/tutorials/control-flow.md index 82919efb09..8d27a28ebd 100644 --- a/docs/src/tutorials/control-flow.md +++ b/docs/src/tutorials/control-flow.md @@ -143,8 +143,9 @@ status = @jit choose_status(Reactant.to_rarray(Float32[1])) ``` For mutable struct fields, parameterize the field type and make its initial value -traced before entering the branch. This requirement also applies to numeric fields. -Assigning a different value to an untraced field inside `@trace if` raises an error. +traced before entering the branch. This requirement also applies to numeric fields: +an untraced field has no location the `if` result can be written back into, so an +assignment to it cannot be carried out of the branch. ```@example control_flow_tutorial using ReactantCore: promote_to_traced diff --git a/src/Enums.jl b/src/Enums.jl index 79cee4bfd6..08a68d56e7 100644 --- a/src/Enums.jl +++ b/src/Enums.jl @@ -166,24 +166,3 @@ Base.@nospecializeinfer function make_tracer( end return prev end - -@inline function to_rarray_internal( - @nospecialize(x::Base.Enum), - @nospecialize(track_numbers::Type), - @nospecialize(sharding), - runtime, - @nospecialize(device), - @nospecialize(client) -) - should_track_enum(typeof(x), track_numbers) || return x - if runtime isa Val{:PJRT} - return TracedEnum{typeof(x)}( - ConcretePJRTNumber(Integer(x); sharding, device, client) - ) - elseif runtime isa Val{:IFRT} - return TracedEnum{typeof(x)}( - ConcreteIFRTNumber(Integer(x); sharding, device, client) - ) - end - return error("Unsupported runtime $runtime") -end diff --git a/src/Ops.jl b/src/Ops.jl index 89198333d7..696e588719 100644 --- a/src/Ops.jl +++ b/src/Ops.jl @@ -2874,10 +2874,6 @@ end Reactant.TracedUtils.set!( args, path[2:end], MLIR.IR.result(if_compiled, residx) ) - else - check_untraced_branch_state( - target, tb_traced_args, fb_traced_args, path[2:end] - ) end end end @@ -2885,26 +2881,6 @@ end return corrected_traced_results end -# An untraced location in a mutable argument (e.g. a struct field holding a plain number or -# enum) has nothing the `if` result can be written back into, so a branch that assigned it -# a new value would be a silent no-op. -function check_untraced_branch_state(target, tb_traced_args, fb_traced_args, path) - for branch_args in (tb_traced_args, fb_traced_args) - leaf = branch_args - for p in path - leaf = Reactant.Compiler.traced_getfield(leaf, p) - end - leaf === target && continue - error( - "if_condition: a branch assigned a value of type $(typeof(leaf)) to an untraced \ - location holding $(repr(target)) (path $(path)); the assignment cannot be \ - carried out of the branch. Make the initial value traced before the `if`, e.g. \ - with `Reactant.ReactantCore.promote_to_traced`.", - ) - end - return nothing -end - """ case( index::TracedRNumber{<:Integer}, branch_fns::Vector, args...; diff --git a/test/core/control_flow.jl b/test/core/control_flow.jl index aa1fa116f4..70db3d0296 100644 --- a/test/core/control_flow.jl +++ b/test/core/control_flow.jl @@ -1168,35 +1168,6 @@ end @test simulation.stop_iteration == 3 end -mutable struct PlainFieldCache{U,C,B} - u::U - count::C - done::B -end - -function plain_field_assigned(u, threshold) - c = PlainFieldCache(u, 0, ReactantCore.promote_to_traced(false)) - @trace if sum(u) > threshold - c.count = 1 - c.done = true - end - return c.count, c.done -end - -function plain_field_untouched(u, threshold) - c = PlainFieldCache(u, 0, ReactantCore.promote_to_traced(false)) - @trace if sum(u) > threshold - c.done = true - end - return c.count, c.done -end - -@testset "if: assignment to an untraced struct field" begin - u = Reactant.to_rarray(Float32[1, 1]) - @test_throws "untraced location" @jit plain_field_assigned(u, 1.0f0) - @test @jit(plain_field_untouched(u, 1.0f0)) == (0, true) -end - function ternary_max(x, y) @trace result = x > y ? x : y return result diff --git a/test/core/enums.jl b/test/core/enums.jl index dbb0bbcaa8..1fd4f90c81 100644 --- a/test/core/enums.jl +++ b/test/core/enums.jl @@ -125,18 +125,6 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 @test c.code == expected end - # A plain enum in the field has nothing the branch result can be written into, the - # same as a plain number; the assignment must error rather than be dropped. - function f_field_plain(u, threshold) - c = EnumCache(u, 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_throws "untraced location" @jit f_field_plain(fresh(), 1.0f0) - function f_field_untouched(u, threshold) c = EnumCache(u, Code.Default, Reactant.ReactantCore.promote_to_traced(false)) @trace if sum(c.u) > threshold From 9cf37a2e797d0aa1812378e584d2baada54941d0 Mon Sep 17 00:00:00 2001 From: Gabriel Baraldi Date: Mon, 14 Sep 2026 09:41:47 -0300 Subject: [PATCH 06/10] Fix concrete enum conversion and scalar broadcasting --- docs/src/tutorials/control-flow.md | 5 +++++ src/Enums.jl | 10 +++++++++- test/core/enums.jl | 20 ++++++++++++++++++++ 3 files changed, 34 insertions(+), 1 deletion(-) diff --git a/docs/src/tutorials/control-flow.md b/docs/src/tutorials/control-flow.md index 8d27a28ebd..510582cd95 100644 --- a/docs/src/tutorials/control-flow.md +++ b/docs/src/tutorials/control-flow.md @@ -170,6 +170,11 @@ 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 index 08a68d56e7..196219f051 100644 --- a/src/Enums.jl +++ b/src/Enums.jl @@ -15,11 +15,16 @@ enum's base integer: a `TracedRNumber` while tracing, a concrete number in the r compiled call. Supports the enum operations (`==`, `!=`, ordered comparisons, `ifelse`) against other `TracedEnum`s and plain `E` values, `Integer`/integer-type conversion, and — once the payload is concrete — conversion back to `E`. +Like plain enums, these values behave as scalars in broadcasting. Converting a plain enum +to `TracedEnum{E}` creates a concrete payload outside compilation and a traced payload +inside compilation. """ mutable struct TracedEnum{E<:Base.Enum} value end +Base.broadcastable(x::TracedEnum) = 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 @@ -67,7 +72,10 @@ function Base.hash(x::TracedEnum{E}, h::UInt) where {E<:Base.Enum} end function Base.convert(::Type{TracedEnum{E}}, x::E) where {E<:Base.Enum} - return ReactantCore.promote_to_traced(x) + if ReactantCore.within_compile() + return ReactantCore.promote_to_traced(x) + end + return to_rarray(x; track_numbers=Number) end # Comparisons and selection diff --git a/test/core/enums.jl b/test/core/enums.jl index 1fd4f90c81..92a976a769 100644 --- a/test/core/enums.jl +++ b/test/core/enums.jl @@ -125,6 +125,17 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 @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 @@ -178,4 +189,13 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 @test res[1] == true @test res[2] == 2 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 From 98332abd343e2895ee8f8c269f34504a92359cc5 Mon Sep 17 00:00:00 2001 From: Gabriel Baraldi Date: Mon, 14 Sep 2026 10:00:31 -0300 Subject: [PATCH 07/10] Separate concrete and traced enum representations --- docs/src/api/api.md | 1 + docs/src/tutorials/control-flow.md | 6 +- src/Enums.jl | 136 +++++++++++++++++++---------- src/Reactant.jl | 2 +- src/compiler/Codegen.jl | 57 +++++++++++- test/core/enums.jl | 33 ++++++- 6 files changed, 177 insertions(+), 58 deletions(-) diff --git a/docs/src/api/api.md b/docs/src/api/api.md index 7a83852b23..c61ea4d127 100644 --- a/docs/src/api/api.md +++ b/docs/src/api/api.md @@ -35,6 +35,7 @@ Reactant.to_rarray 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 510582cd95..435fa7da58 100644 --- a/docs/src/tutorials/control-flow.md +++ b/docs/src/tutorials/control-flow.md @@ -123,8 +123,10 @@ a function from traced arguments, leading to a traced condition. ### Enum state Values created by `@enum` or EnumX's `@enumx` can pass through traced control flow. -A runtime-dependent enum result is a [`Reactant.TracedEnum`](@ref), which supports -comparisons with the original enum and conversion back to it after execution. +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 diff --git a/src/Enums.jl b/src/Enums.jl index 196219f051..0f558fa534 100644 --- a/src/Enums.jl +++ b/src/Enums.jl @@ -1,35 +1,41 @@ -# `Base.Enum` values (`@enum`, EnumX's `@enumx`, ...) are isbits reinterpretations of an -# integer. A traced enum is an enum-typed wrapper around the traced base integer — not a -# `TracedRNumber` with an enum element type — so the payload is an ordinary traced number -# and the wrapper rides the generic struct tracing and result reconstruction. The wrapper -# is mutable with an untyped slot because the result write-back machinery replaces traced -# payloads with concrete numbers via `setfield!`. - enum_basetype(::Type{<:Base.Enum{T}}) where {T} = T +abstract type AbstractReactantEnum{E<:Base.Enum} end + """ TracedEnum{E <: Base.Enum} -A value of the enum type `E` carried through a Reactant compilation. `value` holds the -enum's base integer: a `TracedRNumber` while tracing, a concrete number in the result of a -compiled call. Supports the enum operations (`==`, `!=`, ordered comparisons, `ifelse`) -against other `TracedEnum`s and plain `E` values, `Integer`/integer-type conversion, and — -once the payload is concrete — conversion back to `E`. -Like plain enums, these values behave as scalars in broadcasting. Converting a plain enum -to `TracedEnum{E}` creates a concrete payload outside compilation and a traced payload -inside compilation. +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 + """ -mutable struct TracedEnum{E<:Base.Enum} - value + 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 -Base.broadcastable(x::TracedEnum) = Ref(x) +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::TracedEnum) = getfield(x, :value) +_enum_payload(x::AbstractReactantEnum) = getfield(x, :value) _enum_payload(x::Base.Enum) = Integer(x) _payload_integer(v::AbstractConcreteNumber) = Integer(to_number(v)) @@ -48,33 +54,32 @@ _traced_payload(::Type{I}, x::TracedEnum) where {I} = _enum_payload(x) # Conversions -Base.Integer(x::TracedEnum) = _payload_integer(_enum_payload(x)) -(::Type{T})(x::TracedEnum) where {T<:Integer} = _payload_to(T, _enum_payload(x)) +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::TracedEnum{E}) where {E<:Base.Enum} - v = _enum_payload(x) - v isa TracedRNumber && - error("cannot convert a traced $E back to the plain enum during tracing") - return E(Integer(_payload_integer(v))) +function Base.convert(::Type{E}, x::ConcreteEnum{E}) where {E<:Base.Enum} + return E(Integer(x)) end -(::Type{E})(x::TracedEnum{E}) where {E<:Base.Enum} = convert(E, x) +(::Type{E})(x::ConcreteEnum{E}) where {E<:Base.Enum} = convert(E, x) -# With a concrete payload the wrapper is value-equal to `E`, so it hashes like `E` too. -function Base.hash(x::TracedEnum{E}, h::UInt) where {E<:Base.Enum} - v = _enum_payload(x) - v isa TracedRNumber && return hash(v, h) - return hash(E(Integer(_payload_integer(v))), h) -end +# 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} - if ReactantCore.within_compile() - return ReactantCore.promote_to_traced(x) - end + 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 @@ -90,13 +95,15 @@ for jlop in ( :(Base.isless), ) @eval begin - function $(jlop)(lhs::TracedEnum{E}, rhs::TracedEnum{E}) where {E} + function $(jlop)( + lhs::AbstractReactantEnum{E}, rhs::AbstractReactantEnum{E} + ) where {E} return $(jlop)(_enum_payload(lhs), _enum_payload(rhs)) end - function $(jlop)(lhs::TracedEnum{E}, rhs::E) where {E} + function $(jlop)(lhs::AbstractReactantEnum{E}, rhs::E) where {E} return $(jlop)(_enum_payload(lhs), _enum_payload(rhs)) end - function $(jlop)(lhs::E, rhs::TracedEnum{E}) where {E} + function $(jlop)(lhs::E, rhs::AbstractReactantEnum{E}) where {E} return $(jlop)(_enum_payload(lhs), _enum_payload(rhs)) end end @@ -109,9 +116,42 @@ function Base.ifelse( return TracedEnum{E}(ifelse(pred, _traced_payload(I, x), _traced_payload(I, y))) end -# Tracing. The wrapper itself is an ordinary struct handled by the generic machinery; only -# the plain `Base.Enum` value needs entry points, and they place the payload at the -# wrapper-relative path (`value` is field 1). +# 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 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) @@ -128,10 +168,10 @@ Base.@nospecializeinfer function traced_type_inner( @nospecialize(runtime) ) should_track_enum(T, track_numbers) || return T - if mode == ArrayToConcrete || - mode == NoStopTracedTrack || - mode == TracedTrack || - mode == TracedSetPath + if mode == ArrayToConcrete + N = traced_type_inner(enum_basetype(T), seen, mode, Number, ndevices, runtime) + return ConcreteEnum{T,N} + elseif mode == NoStopTracedTrack || mode == TracedTrack || mode == TracedSetPath return TracedEnum{T} end return T @@ -156,10 +196,10 @@ Base.@nospecializeinfer function make_tracer( RT = Core.Typeof(prev) should_track_enum(RT, track_numbers) || return prev if mode == ArrayToConcrete - runtime isa Val{:PJRT} && return TracedEnum{RT}( + runtime isa Val{:PJRT} && return ConcreteEnum{RT}( ConcretePJRTNumber(Integer(prev); sharding, device, client) ) - runtime isa Val{:IFRT} && return TracedEnum{RT}( + runtime isa Val{:IFRT} && return ConcreteEnum{RT}( ConcreteIFRTNumber(Integer(prev); sharding, device, client) ) error("Unsupported runtime $runtime") diff --git a/src/Reactant.jl b/src/Reactant.jl index 095c8388e4..fe8efa7755 100644 --- a/src/Reactant.jl +++ b/src/Reactant.jl @@ -309,7 +309,7 @@ export ConcreteRArray, within_compile @static if VERSION ≥ v"1.11" - @eval $(Expr(:public, :Periodic, :Binomial, :TracedEnum)) + @eval $(Expr(:public, :Periodic, :Binomial, :TracedEnum, :ConcreteEnum)) end const registry = Ref{Union{Nothing,MLIR.IR.DialectRegistry}}() diff --git a/src/compiler/Codegen.jl b/src/compiler/Codegen.jl index 613724eef5..83dcdfd8e8 100644 --- a/src/compiler/Codegen.jl +++ b/src/compiler/Codegen.jl @@ -147,6 +147,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 @@ -1129,7 +1162,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 @@ -1153,7 +1186,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 @@ -1179,7 +1225,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/core/enums.jl b/test/core/enums.jl index 92a976a769..aa39493ba9 100644 --- a/test/core/enums.jl +++ b/test/core/enums.jl @@ -1,5 +1,5 @@ using EnumX, Reactant, Test -using Reactant: @trace, TracedEnum, ConcreteRNumber +using Reactant: @trace, TracedEnum, ConcreteEnum, ConcreteRNumber @enum Fruit apple = 1 banana = 2 cherry = 3 @enum Small::UInt8 low = 7 high = 200 @@ -16,7 +16,7 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 @testset "ifelse" begin f_ifelse(u) = ifelse(sum(u) > 1, Code.Success, Code.MaxIters) res = @jit f_ifelse(fresh()) - @test res isa TracedEnum{Code.T} + @test res isa ConcreteEnum{Code.T} @test res == Code.Success @test Code.Success == res @test convert(Code.T, res) === Code.Success @@ -171,7 +171,7 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 @testset "non-default base type" begin f_small(u) = ifelse(sum(u) > 1, high, low) res = @jit f_small(fresh()) - @test res isa TracedEnum{Small} + @test res isa ConcreteEnum{Small} @test res == high @test Integer(res) === UInt8(200) f_small_int(u) = Integer(ifelse(sum(u) > 1, high, low)) @@ -182,7 +182,7 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 @testset "enum arguments" begin f_arg(u, fruit) = (fruit == banana, Int(fruit)) fruit = Reactant.to_rarray(banana; track_numbers=Number) - @test fruit isa TracedEnum{Fruit} + @test fruit isa ConcreteEnum{Fruit} @test fruit == banana @test Reactant.to_rarray(banana) === banana res = @jit f_arg(fresh(), fruit) @@ -190,6 +190,31 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 @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 From 16ee7759ceda90678ef8fa286168cc5f2e1d8dcc Mon Sep 17 00:00:00 2001 From: Gabriel Baraldi Date: Mon, 21 Sep 2026 19:17:25 +0000 Subject: [PATCH 08/10] Clarify enum tracing and validate payload paths Move enum tracing hooks into Tracing.jl and make wrapper dispatch explicit. Reuse numeric tracing for plain enum branch results while keeping tracking modes from promoting plain enums. Restrict the plain-enum field-access adaptation to payload index 1 and add regressions for invalid fields. Validation: 71 enum and 147 control-flow assertions passed on CPU/PJRT with Julia 1.12.6; JuliaFormatter 1 and git diff --check passed. --- src/Enums.jl | 99 ---------------------------------- src/Reactant.jl | 2 +- src/Tracing.jl | 115 ++++++++++++++++++++++++++++++++++++++++ src/compiler/Codegen.jl | 6 ++- test/core/enums.jl | 6 +++ 5 files changed, 127 insertions(+), 101 deletions(-) diff --git a/src/Enums.jl b/src/Enums.jl index 0f558fa534..e9fe811af4 100644 --- a/src/Enums.jl +++ b/src/Enums.jl @@ -115,102 +115,3 @@ function Base.ifelse( I = enum_basetype(E) return TracedEnum{E}(ifelse(pred, _traced_payload(I, x), _traced_payload(I, y))) 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 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, Number, ndevices, runtime) - return ConcreteEnum{T,N} - elseif mode == NoStopTracedTrack || mode == TracedTrack || mode == TracedSetPath - return TracedEnum{T} - end - return T -end - -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 - payload = TracedRNumber{enum_basetype(RT)}( - (append_path(path, 1),), @opcall(constant(Integer(prev))).mlir_data - ) - seen[gensym("enum")] = payload - return TracedEnum{RT}(payload) - elseif mode == TracedToConcrete - throw("Input is not a traced-type: $(RT)") - end - return prev -end diff --git a/src/Reactant.jl b/src/Reactant.jl index fe8efa7755..cbd95f2005 100644 --- a/src/Reactant.jl +++ b/src/Reactant.jl @@ -270,8 +270,8 @@ export StackedBatchDuplicated, StackedBatchDuplicatedNoNeed const TracedType = Union{TracedRArray,TracedRNumber,MissingTracedValue} include("ControlFlow.jl") -include("Tracing.jl") include("Enums.jl") +include("Tracing.jl") include("compiler/Compiler.jl") diff --git a/src/Tracing.jl b/src/Tracing.jl index 288020cfa9..b7eaad0a83 100644 --- a/src/Tracing.jl +++ b/src/Tracing.jl @@ -2475,3 +2475,118 @@ 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 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, Number, ndevices, runtime) + return ConcreteEnum{T,N} + elseif mode == NoStopTracedTrack + return TracedEnum{T} + end + return T +end + +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 83dcdfd8e8..f3d5b7cd07 100644 --- a/src/compiler/Codegen.jl +++ b/src/compiler/Codegen.jl @@ -27,7 +27,11 @@ 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 traced_getfield(@nospecialize(obj::Base.Enum), field) = obj +@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( diff --git a/test/core/enums.jl b/test/core/enums.jl index aa39493ba9..aa64c637e2 100644 --- a/test/core/enums.jl +++ b/test/core/enums.jl @@ -224,3 +224,9 @@ fresh() = Reactant.to_rarray(Float32[1, 1]) # sum == 2 [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 From 4ecc0c8da8b5947a5f42ddefbf1150e193edc975 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Sergio=20S=C3=A1nchez=20Ram=C3=ADrez?= Date: Wed, 23 Sep 2026 08:36:06 -0500 Subject: [PATCH 09/10] refactor and add support for `TracedTrack`, `TracedSetPath` modes --- src/Tracing.jl | 119 ++++++++++++++++++++++++------------------------- 1 file changed, 59 insertions(+), 60 deletions(-) diff --git a/src/Tracing.jl b/src/Tracing.jl index 27ea652c52..acea21c4d5 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 || mode == TracedTrack || mode == TracedSetPath + 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)) @@ -2478,65 +2536,6 @@ 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 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, Number, ndevices, runtime) - return ConcreteEnum{T,N} - elseif mode == NoStopTracedTrack - return TracedEnum{T} - end - return T -end - Base.@nospecializeinfer function make_tracer( seen, @nospecialize(prev::Base.Enum), @@ -2563,7 +2562,7 @@ Base.@nospecializeinfer function make_tracer( ConcreteIFRTNumber(Integer(prev); sharding, device, client) ) error("Unsupported runtime $runtime") - elseif mode == NoStopTracedTrack + elseif mode == TracedTrack || mode == NoStopTracedTrack || mode == TracedSetPath # Plain enum branch results need the same constant promotion as numbers. payload = make_tracer( seen, From 59e193f50ace2a77c0b9a72bd5a0876ed9142a93 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Sergio=20S=C3=A1nchez=20Ram=C3=ADrez?= Date: Thu, 24 Sep 2026 14:44:24 -0500 Subject: [PATCH 10/10] revert back `TracedTrack`, `TracedSetPath` behavior --- src/Tracing.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/Tracing.jl b/src/Tracing.jl index acea21c4d5..45a847560e 100644 --- a/src/Tracing.jl +++ b/src/Tracing.jl @@ -887,7 +887,7 @@ Base.@nospecializeinfer function traced_type_inner( if mode == ArrayToConcrete N = traced_type_inner(enum_basetype(T), seen, mode, track_numbers, ndevices, runtime) return ConcreteEnum{T,N} - elseif mode == NoStopTracedTrack || mode == TracedTrack || mode == TracedSetPath + elseif mode == NoStopTracedTrack return TracedEnum{T} end return T @@ -2562,7 +2562,7 @@ Base.@nospecializeinfer function make_tracer( ConcreteIFRTNumber(Integer(prev); sharding, device, client) ) error("Unsupported runtime $runtime") - elseif mode == TracedTrack || mode == NoStopTracedTrack || mode == TracedSetPath + elseif mode == NoStopTracedTrack # Plain enum branch results need the same constant promotion as numbers. payload = make_tracer( seen,