Skip to content
Closed
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
31 changes: 22 additions & 9 deletions nx/lib/nx/defn/expr.ex
Original file line number Diff line number Diff line change
Expand Up @@ -99,20 +99,33 @@ defmodule Nx.Defn.Expr do
inspection.
"""
def metadata(expr, metadata) when is_map(metadata) do
case to_container_expr(expr) do
%{data: %{context: context}} = res ->
expr(res, context, :metadata, [Nx.devectorize(expr), metadata])
if expr_container?(expr) do
case to_container_expr(expr) do
%{data: %{context: context}} = res ->
expr(res, context, :metadata, [Nx.devectorize(expr), metadata])

t when is_tuple(t) ->
context = elem(t, 0).data.context
t when is_tuple(t) ->
context = elem(t, 0).data.context

tuple(
expr(tuple_out(tuple_size(t)), context, :metadata, [Nx.devectorize(expr), metadata]),
Tuple.to_list(t)
)
tuple(
expr(tuple_out(tuple_size(t)), context, :metadata, [Nx.devectorize(expr), metadata]),
Tuple.to_list(t)
)
end
else
# BinaryBackend.block/4 (and similar) may re-run callbacks that call
# stop_grad/custom_grad/metadata; concrete tensors must pass through.
expr
end
end

defp expr_container?(container) do
Composite.reduce(container, true, fn
%T{data: %Expr{}}, true -> true
_, _ -> false
end)
end

@doc """
Creates a tuple with elements in `list` that points to tuple
expression `expr`.
Expand Down
Loading