diff --git a/lib/axon.ex b/lib/axon.ex index 39c0c416..1e7a9cca 100644 --- a/lib/axon.ex +++ b/lib/axon.ex @@ -4367,6 +4367,228 @@ defmodule Axon do end end + @doc """ + Captures a model or subgraph, returning a reusable function. + + This returns an arity-1 function that accepts new inputs and returns + the model with rewired inputs. For single-input models, pass an Axon + graph directly. For multi-input models, pass a map of input names to + Axon graphs. + + This is useful for transfer learning where you want to extract a + pretrained model's feature extractor: + + # Load a pretrained model + resnet = MyModels.resnet50() + + # Capture at the pooling layer (returns a function) + backbone = Axon.capture(resnet, to: "avg_pool") + + # Use with new inputs + new_input = Axon.input("my_features") + features = backbone.(new_input) + + # Add your own head + my_model = + features + |> Axon.dense(256, activation: :relu) + |> Axon.dense(num_classes) + + For models with multiple inputs, pass a map: + + model = Axon.capture(multi_input_model) + output = model.(%{"image" => image_input, "text" => text_input}) + + You can also capture an entire model without truncation: + + encoder = Axon.capture(encoder_model) + encoded = encoder.(my_input) + + `:to` cuts the graph so the named layer becomes the new output. + Layers after that point are dropped from the captured model: + + model = + Axon.input("features", shape: {nil, 10}) + |> Axon.dense(32, name: "hidden1") + |> Axon.relu(name: "relu1") + |> Axon.dense(16, name: "hidden2") + |> Axon.dense(2, name: "output") + + backbone = Axon.capture(model, to: "hidden2") + new_model = backbone.(Axon.input("x", shape: {nil, 10})) + # new_model ends at "hidden2"; "output" is not present + + On multi-output models the same rule applies from the named layer. + Capturing a shared trunk layer before the head split returns that + single tensor (both heads and the container are cut off). Capturing + one head after the split returns only that branch as a single output; + sibling heads are not included: + + shared = + Axon.input("x", shape: {nil, 4}) + |> Axon.dense(8, name: "trunk") + |> Axon.relu(name: "shared") + + model = + Axon.container(%{ + a: Axon.dense(shared, 2, name: "head_a"), + b: Axon.dense(shared, 3, name: "head_b") + }) + + # Single-tensor backbone up to the shared layer + Axon.capture(model, to: "shared") + + # Only head_a; head_b and the container are gone + Axon.capture(model, to: "head_a") + + Layer names can be discovered using `Axon.properties/1`: + + Axon.properties(model) |> Map.keys() + + ## Options + + * `:to` - name of the layer that should become the captured model's + output. Defaults to the model's original output (entire graph). + + """ + @doc type: :graph + def capture(%Axon{} = axon, opts \\ []) when is_list(opts) do + truncated = + case Keyword.fetch(opts, :to) do + {:ok, layer_name} when is_binary(layer_name) -> + name_to_id = build_name_to_id_map(axon) + + target_id = + case Map.fetch(name_to_id, layer_name) do + {:ok, id} -> + id + + :error -> + available = name_to_id |> Map.keys() |> Enum.sort() + + raise ArgumentError, + "layer #{inspect(layer_name)} not found in model. " <> + "Available layers: #{inspect(available)}" + end + + %Axon{axon | output: target_id} + + {:ok, other} -> + raise ArgumentError, + "expected :to option to be a string layer name, got: #{inspect(other)}" + + :error -> + axon + end + + input_name_to_id = get_input_name_to_id_map(truncated) + + fn new_inputs -> + rewire_inputs(truncated, input_name_to_id, new_inputs) + end + end + + defp build_name_to_id_map(%Axon{output: id, nodes: nodes}) do + {name_to_id, _, _} = do_build_name_to_id(id, nodes, {%{}, %{}, %{}}) + name_to_id + end + + defp do_build_name_to_id(id, nodes, {_name_to_id, cache, _op_counts} = acc) do + case cache do + %{^id => _} -> + acc + + %{} -> + %Axon.Node{parent: parents, name: name_fn, op_name: op_name} = nodes[id] + + {name_to_id, cache, op_counts} = + Enum.reduce(parents, acc, fn parent_id, acc -> + do_build_name_to_id(parent_id, nodes, acc) + end) + + name = name_fn.(op_name, op_counts) + op_counts = Map.update(op_counts, op_name, 1, fn x -> x + 1 end) + name_to_id = Map.put(name_to_id, name, id) + + {name_to_id, Map.put(cache, id, name), op_counts} + end + end + + defp get_input_name_to_id_map(%Axon{output: id, nodes: nodes}) do + {inorder_nodes, _} = traverse_nodes(id, nodes, [], MapSet.new()) + + inorder_nodes + |> Enum.filter(fn %Axon.Node{op: op} -> op == :input end) + |> Map.new(fn %Axon.Node{id: id, name: name_fn} -> + {name_fn.(:input, %{}), id} + end) + end + + defp rewire_inputs(%Axon{output: output_id, nodes: nodes}, input_name_to_id, new_inputs) do + new_inputs_map = + case new_inputs do + %Axon{} = single_input -> + if map_size(input_name_to_id) != 1 do + raise ArgumentError, + "model has #{map_size(input_name_to_id)} inputs, expected a map " <> + "with keys: #{inspect(Map.keys(input_name_to_id))}" + end + + {input_name, _} = Enum.fetch!(input_name_to_id, 0) + %{input_name => single_input} + + %{} = inputs_map -> + inputs_map + + other -> + raise ArgumentError, + "expected an Axon graph or a map of input names to Axon graphs, " <> + "got: #{inspect(other)}" + end + + # Validate all expected inputs are provided + expected_names = Map.keys(input_name_to_id) + provided_names = Map.keys(new_inputs_map) + + missing = expected_names -- provided_names + extra = provided_names -- expected_names + + if missing != [] do + raise ArgumentError, "missing inputs: #{inspect(missing)}" + end + + if extra != [] do + raise ArgumentError, "unexpected inputs: #{inspect(extra)}" + end + + id_mapping = + Map.new(input_name_to_id, fn {name, old_id} -> + {old_id, new_inputs_map[name]} + end) + + merged_nodes = + Enum.reduce(new_inputs_map, nodes, fn {_k, %Axon{nodes: new_nodes}}, acc_nodes -> + Map.merge(new_nodes, acc_nodes) + end) + + merged_nodes = Map.drop(merged_nodes, Map.values(input_name_to_id)) + + updated_nodes = + Map.new(merged_nodes, fn {node_id, node} -> + updated_parents = + Enum.map(node.parent, fn parent_id -> + case Map.fetch(id_mapping, parent_id) do + {:ok, %Axon{output: new_output_id}} -> new_output_id + :error -> parent_id + end + end) + + {node_id, %{node | parent: updated_parents}} + end) + + %Axon{output: output_id, nodes: updated_nodes} + end + @doc """ Builds the given model to `{init_fn, predict_fn}`.