Mostafa86

Mostafa86

Does Nx has an equivalent for JAX tree_map jax.tree_util.tree_map — JAX documentation?

Thank you!

/nx

Showing Posts 1 to 9

Mostafa86

Mostafa86 OP

Couldn’t find any so I hacked a quick one that fits my use case. Any feedback or alternatives is more than welcome as I am pretty new to Elixir.

defmodule TreeMap do
  import Nx.Defn
  import Nx, only: [is_tensor: 1]

  defn fun({a, b, c}) do
    Nx.multiply(c, Nx.add(a, b))
  end

  def process_list_of_arbitrary_nested_maps(list_of_arbitrary_nested_maps) do
    list_of_arbitrary_nested_maps
    |> Enum.map(&Map.to_list/1)
    |> Enum.map(&Enum.sort/1)
    |> Enum.zip()
    |> Enum.map(&Tuple.to_list/1)
    |> Enum.map(fn elem ->
      elem |> Enum.map_reduce("", fn x, _ -> {x |> elem(1), x |> elem(0)} end)
    end)
  end

  def traverse(list_of_arbitrary_nested_maps, fun) do
    condition =
      list_of_arbitrary_nested_maps
      |> Enum.map(&Map.values/1)
      |> Enum.map(fn x -> x |> Enum.all?(&is_tensor/1) end)
      |> Enum.all?()

    cond do
      condition ->
        list_of_arbitrary_nested_maps
        |> process_list_of_arbitrary_nested_maps()
        |> Enum.map(fn {l, k} -> %{"#{k}" => fun.(l |> List.to_tuple())} end)
        |> Enum.reduce(&Map.merge/2)

      true ->
        list_of_arbitrary_nested_maps
        |> process_list_of_arbitrary_nested_maps()
        |> Enum.map(fn {l, k} ->
          %{
            "#{k}" => traverse(l, &fun/1)
          }
        end)
        |> Enum.reduce(&Map.merge/2)
    end
  end
end
polvalente

polvalente

Nx Core Team

Maybe Nx.Defn.Composite.traverse?

Mostafa86

Mostafa86 OP

Thank you. I checked polaris/lib/polaris/shared.ex at ad85df596966548b7c38a89e2032263d0a0b4527 · elixir-nx/polaris · GitHub and I can see that it can be used for 2 arguments but not sure how to generalize to more than 2 arguments :thinking:

polvalente

polvalente

Nx Core Team

What do you mean more than 2 args?

polvalente

polvalente

Nx Core Team

Maybe you want to use Nx.Container.traverse together with Nx.Defn.Composite.flatten_list in a way that your accumulator yields the corresponding “zipped” values for each container.

polvalente

polvalente

Nx Core Team

Yet another possibility is to flatten each container and zip them together

Mostafa86

Mostafa86 OP

I mean more than 2 Nx.Containers

Mostafa86

Mostafa86 OP

Thanks !

This immensely simplified my implementation :smiley:

defmodule TreeMap do
  import Nx.Defn

  defn fun({a, b, c}) do
    Nx.multiply(c, Nx.add(a, b))
  end

  def traverse(list_of_arbitrary_nested_maps, fun) when is_list(list_of_arbitrary_nested_maps) do
    list_of_arbitrary_nested_maps
    |> Enum.map(&List.wrap/1)
    |> Enum.map(&Nx.Defn.Composite.flatten_list/1)
    |> Enum.zip()
    |> Enum.map(fn t -> fun.(t) end)
    |> Kernel.then(fn v ->
      {v, _} =
        Nx.Defn.Composite.traverse(list_of_arbitrary_nested_maps |> Enum.at(0), 0, fn _, acc ->
          {v |> Enum.at(acc), acc + 1}
        end)

      v
    end)
  end
end
polvalente

polvalente

Nx Core Team

Here’s my suggestion:

  def traverse([template_container | _] = list_of_arbitrary_nested_maps, fun) when is_list(list_of_arbitrary_nested_maps) do
    zipped_containers = 
      list_of_arbitrary_nested_maps
      |> Enum.map(&Nx.Defn.Composite.flatten_list[&1])
      |> Enum.zip_with(fun)
    
    {v, []} =
        Nx.Defn.Composite.traverse(template_container, zipped_containers, fn _, [h | t] ->
          {h, t}
        end)

   v
  end
— All posts loaded —

Where Next? Top

Trending in Discussions Top

mudasobwa
I am happy to introduce the very α version of the new programming language compiled to BEAM. Welcome Cure. It has literally three kille...
New
budgie
A little off-topic, but I feel like people here have a good head on their shoulders. I used to be quite good at making software. Was luc...
New
axelson
Hi there! :wave: @frigidcode and I (but mostly him) have been running an Elixir Book club, we’re almost done with Designing Elixir Syste...
New
achempion
I’ve been using Emacs as my main code editor for more than a two years. It’s a custom build version although I’ve tried doom emacs and sp...
New
budgie
I love Elixir. It’s one of 2 programming languages I’ve ever fallen in love with. But I don’t use it anymore. Serverless was the promis...
New
jtormey
Lately I’ve been thinking about how to organize components as a LiveView application grows. One of the pain points I’ve found (for myself...
New
Null-logic-0
What IDE or editor are you using for Elixir development? Personally, I use Zed, and I really like it, but sometimes I wish there were a ...
New

Other Trending Topics Top

GenericJam
Edit: 2026 May 15 - This post is archived. Mob is alive!! Main docs: mob v0.7.11 — Documentation A bit of explanation for the slightly c...
New
garrison
Hobbes is a low-level distributed database for the Elixir programming language. Hobbes provides a simple, safe, and scalable storage lay...
New
KristerV
Hey. Is there anyone here who creates agents in their apps? Not talking about using agents, but creating them. I’m finding it pretty diff...
New
mcass19
ExRatatui lets you cook up rich terminal UIs in Elixir, powered by Rust’s ratatui via Rustler NIFs. Build interactive terminal applicatio...
New
webofbits
With AI doing more of the implementation work, I’ve been wondering how much coding I should deliberately keep doing myself. My main conc...
#ai
New
georgeguimaraes
Just published claude-code-elixir, a plugin marketplace for Claude Code with Elixir support. These are the plugins I’ve been using for my...
New

We're in Beta

About us Mission Statement

Options

Thread Display Mode




Thread Preview

Skip Thread Previews