Executes f with abstract array arguments and records every primitive operation into
an AnvlGraph.
The resulting graph can be lowered to StableHLO (via stablehlo()) or transformed
(e.g. via transform_gradient()).
Usage
trace_fn(
f,
args = NULL,
desc = NULL,
mode = NULL,
args_flat = NULL,
in_tree = NULL,
optimize = FALSE
)Arguments
- f
(
function)
The function to trace. Must not be aJitFunction(i.e. already jitted).- args
(
listof (AnvlArray|AbstractArray))
The (unflattened) arguments to the function. Mutually exclusive with theargs_flat/in_treepair.- desc
(
NULL|GraphDescriptor)
Optional descriptor. WhenNULL(default), a new descriptor is created.- mode
(
character(1))
How to handle the inputs. Options are:"toplevel": Used for jit(). Default."subgraph": Use for tracing subgraphs in higher-order primitives likeprim_while()."inline": Use for transformations like jit, where the graph is later inlined into the parent graph.
- args_flat
(
list)
Flattened arguments. Must be accompanied byin_tree.- in_tree
(
Node)
Tree structure describing howargs_flatmaps back tof's arguments.- optimize
(
logical(1)|character())
Which graph optimization passes to run on the traced graph before returning it.TRUEruns all passes,FALSE(default) runs none, and a character vector selects a subset by name. The available passes are:"inline_scalars": replace scalar-shaped constants with inline literals."remove_unused_constants": drop constants not referenced by the graph.
jit()always traces with all passes enabled.
Value
An AnvlGraph containing the traced operations.
See also
stablehlo() to lower the graph, jit() for end-to-end
compilation.