Skip to content

Latest commit

 

History

History
406 lines (329 loc) · 10.5 KB

File metadata and controls

406 lines (329 loc) · 10.5 KB

NxEinsum

Mix.install([
  {:nx, "~> 0.12.0"},
  {:kino, "~> 0.19.0"},
  {:pythonx, "~> 0.4.2"},
  {:kino_pythonx, "~> 0.1.0"},
  {:exla, "~> 0.12.0"}
])
[project]
name = "project"
version = "0.0.0"
requires-python = "==3.13.*"
dependencies = ["numpy"]

Einsum Macro

defmodule NxEinsum do
  defmacro defeinsum(name, notation) when is_binary(notation) do
    func_name =
      case name do
        {atom, _, _} when is_atom(atom) -> atom
        atom when is_atom(atom) -> atom
      end

    {input_specs, output_spec} = parse_notation(notation)
    arity = length(input_specs)

    arg_names = Enum.take([:a, :b, :c, :d, :e, :f, :g, :h], arity)
    arg_vars = Enum.map(arg_names, fn n -> Macro.var(n, nil) end)

    body = build_body(input_specs, output_spec, arg_vars)

    quote do
      def unquote(func_name)(unquote_splicing(arg_vars)) do
        unquote(body)
      end
    end
  end

  defp parse_notation(notation) do
    [inputs_str, output_str] = String.split(notation, "->")
    inputs = inputs_str |> String.split(",") |> Enum.map(&String.graphemes/1)
    {inputs, String.graphemes(output_str)}
  end

  # ── Single input ────────────────────────────────────────────────────────────

  defp build_body([spec], output_spec, [var]) do
    {var, spec} = reduce_repeated(var, spec)
    repeated = length(spec) != length(Enum.uniq(spec))
    summed = Enum.reject(spec, &(&1 in output_spec))

    cond do
      # trace "ii->"
      repeated and output_spec == [] ->
        quote do: unquote(var) |> Nx.take_diagonal() |> Nx.sum()

      # diagonal "ii->i"
      repeated ->
        quote do: Nx.take_diagonal(unquote(var))

      # sum over axes "ij->i"
      summed != [] ->
        sum_axes = Enum.map(summed, &Enum.find_index(spec, fn x -> x == &1 end))
        remaining = Enum.reject(spec, &(&1 in summed))
        result = quote do: Nx.sum(unquote(var), axes: unquote(sum_axes))

        if remaining == output_spec do
          result
        else
          axes = Enum.map(output_spec, fn idx -> Enum.find_index(remaining, &(&1 == idx)) end)
          quote do: Nx.transpose(unquote(result), axes: unquote(axes))
        end

      # transpose/permute "ij->ji"
      spec != output_spec ->
        axes = Enum.map(output_spec, fn idx -> Enum.find_index(spec, &(&1 == idx)) end)
        quote do: Nx.transpose(unquote(var), axes: unquote(axes))

      # identity "ij->ij"
      true ->
        var
    end
  end

  # ── 2 inputs: fold pairwise ─────────────────────────────────────────────
  defp build_body([spec_a, spec_b], output_spec, [var_a, var_b]) do
    {var_a, spec_a} = reduce_repeated(var_a, spec_a)
    {var_b, spec_b} = reduce_repeated(var_b, spec_b)

    shared = Enum.filter(spec_a, &(&1 in spec_b))
    contracted = Enum.filter(shared, &(&1 not in output_spec))
    batch = Enum.filter(shared, &(&1 in output_spec))

    free_a = Enum.reject(spec_a, &(&1 in shared))
    free_b = Enum.reject(spec_b, &(&1 in shared))

    new_order_a = batch ++ free_a ++ contracted
    new_order_b = batch ++ free_b ++ contracted

    axes_a = Enum.map(new_order_a, &Enum.find_index(spec_a, fn x -> x == &1 end))
    axes_b = Enum.map(new_order_b, &Enum.find_index(spec_b, fn x -> x == &1 end))

    n_batch = length(batch)
    n_free_a = length(free_a)
    n_free_b = length(free_b)
    n_contract = length(contracted)

    safe_batch_axes = if n_batch == 0, do: [], else: Enum.to_list(0..(n_batch - 1))

    safe_contract_axes_a =
      if n_contract == 0,
        do: [],
        else: Enum.to_list((n_batch + n_free_a)..(n_batch + n_free_a + n_contract - 1))

    safe_contract_axes_b =
      if n_contract == 0,
        do: [],
        else: Enum.to_list((n_batch + n_free_b)..(n_batch + n_free_b + n_contract - 1))

    ta =
      if axes_a == Enum.to_list(0..max(length(axes_a) - 1, 0)) && length(axes_a) > 0 do
        var_a
      else
        quote do: Nx.transpose(unquote(var_a), axes: unquote(axes_a))
      end

    tb =
      if axes_b == Enum.to_list(0..max(length(axes_b) - 1, 0)) && length(axes_b) > 0 do
        var_b
      else
        quote do: Nx.transpose(unquote(var_b), axes: unquote(axes_b))
      end

    # Output of Nx.dot/6 is: [batch..., free_a..., free_b...]
    dot_spec = batch ++ free_a ++ free_b

    dot =
      quote do
        Nx.dot(
          unquote(ta),
          unquote(safe_contract_axes_a),
          unquote(safe_batch_axes),
          unquote(tb),
          unquote(safe_contract_axes_b),
          unquote(safe_batch_axes)
        )
      end

    # Indices in dot_spec not in output_spec need to be summed
    summed = Enum.reject(dot_spec, &(&1 in output_spec))

    cond do
      output_spec == [] ->
        quote do: Nx.sum(unquote(dot))

      summed != [] ->
        sum_axes = Enum.map(summed, &Enum.find_index(dot_spec, fn x -> x == &1 end))
        remaining = Enum.reject(dot_spec, &(&1 in summed))
        result = quote do: Nx.sum(unquote(dot), axes: unquote(sum_axes))

        if remaining == output_spec do
          result
        else
          axes = Enum.map(output_spec, fn idx -> Enum.find_index(remaining, &(&1 == idx)) end)
          quote do: Nx.transpose(unquote(result), axes: unquote(axes))
        end

      dot_spec == output_spec ->
        dot

      true ->
        axes = Enum.map(output_spec, fn idx -> Enum.find_index(dot_spec, &(&1 == idx)) end)
        quote do: Nx.transpose(unquote(dot), axes: unquote(axes))
    end
  end

  # ── >=3 inputs: fold pairwise ─────────────────────────────────────────────
  defp build_body(specs, output_spec, vars) do
    [{first_spec, first_var} | rest] = Enum.zip(specs, vars)

    {stmts, final_spec, final_var} =
      rest
      |> Enum.with_index(2)
      |> Enum.reduce({[], first_spec, first_var}, fn {{spec_b, var_b}, i},
                                                     {stmts, acc_spec, acc_var} ->
        future_indices = specs |> Enum.drop(i) |> List.flatten()
        needed = Enum.uniq(output_spec ++ future_indices)

        intermediate_spec =
          (acc_spec ++ spec_b)
          |> Enum.uniq()
          |> Enum.filter(&(&1 in needed))

        expr = build_body([acc_spec, spec_b], intermediate_spec, [acc_var, var_b])
        tmp = Macro.unique_var(:tmp, __MODULE__)

        {stmts ++ [{:=, [], [tmp, expr]}], intermediate_spec, tmp}
      end)

    result =
      cond do
        # scalar output: sum everything remaining
        output_spec == [] ->
          quote do: Nx.sum(unquote(final_var))

        # already correct order
        final_spec == output_spec ->
          final_var

        # needs transpose
        true ->
          axes = Enum.map(output_spec, fn idx -> Enum.find_index(final_spec, &(&1 == idx)) end)
          quote do: Nx.transpose(unquote(final_var), axes: unquote(axes))
      end

    {:__block__, [], stmts ++ [result]}
  end

  # Reduces repeated indices in a spec by taking diagonal/trace first
  # "ii" -> take_diagonal -> "i"
  # "iij" -> take_diagonal -> "ij" (diagonal over first two)
  defp reduce_repeated(var, spec) do
    duplicates =
      spec
      |> Enum.group_by(& &1)
      |> Enum.filter(fn {_, v} -> length(v) > 1 end)
      |> Enum.map(fn {k, _} -> k end)

    Enum.reduce(duplicates, {var, spec}, fn idx, {cur_var, cur_spec} ->
      # Find the two positions of this index
      positions =
        cur_spec
        |> Enum.with_index()
        |> Enum.filter(fn {x, _} -> x == idx end)
        |> Enum.map(fn {_, i} -> i end)

      [pos1, pos2 | _] = positions

      # Bring the two axes together via transpose, take diagonal, update spec
      other_axes = Enum.reject(Enum.to_list(0..(length(cur_spec) - 1)), &(&1 in [pos1, pos2]))
      new_order = [pos1, pos2] ++ other_axes
      new_spec = [idx] ++ Enum.map(other_axes, &Enum.at(cur_spec, &1))

      new_var =
        if new_order == Enum.to_list(0..(length(new_order) - 1)) do
          quote do: Nx.take_diagonal(unquote(cur_var))
        else
          quote do: Nx.take_diagonal(Nx.transpose(unquote(cur_var), axes: unquote(new_order)))
        end

      {new_var, new_spec}
    end)
  end
end

Kino.nothing()
Nx.global_default_backend(EXLA.Backend)

Test Cases

defmodule MyOps do
  import NxEinsum

  defeinsum(matmul, "ij,ji->i")
  defeinsum(bar, "ij,ji,j->")
  defeinsum(fuu, "ij,ji->")
  defeinsum(baz, "ij,ji->ji")
  defeinsum(diag, "ii,jj->")
  defeinsum(diag1, "ii,jj->i")
  defeinsum(diag2, "ii,jj->j")
  defeinsum(diag3, "ii,jj->ij")
  defeinsum(weighted_trace, "ij,ji,j->")
end

MyOps.matmul(
  Nx.tensor([[3, 2], [4, 5]]),
  Nx.tensor([[1, -1], [1, 3]])
)
|> Nx.to_list()
|> IO.inspect()

MyOps.bar(
  Nx.tensor([[3, 2], [4, 5]]),
  Nx.tensor([[1, -1], [1, 3]]),
  Nx.tensor([1, 2])
)
|> Nx.to_number()
|> IO.inspect()

MyOps.fuu(
  Nx.tensor([[3, 2], [4, 5]]),
  Nx.tensor([[1, -1], [1, 3]])
)
|> Nx.to_number()
|> IO.inspect()

MyOps.baz(
  Nx.tensor([[3, 2], [4, 5]]),
  Nx.tensor([[1, -1], [1, 3]])
)
|> Nx.to_list()
|> IO.inspect()

MyOps.diag(
  Nx.tensor([[3, 2], [4, 5]]),
  Nx.tensor([[1, -1], [1, 3]])
)
|> Nx.to_number()
|> IO.inspect()

MyOps.diag1(
  Nx.tensor([[3, 2], [4, 5]]),
  Nx.tensor([[1, -1], [1, 3]])
)
|> Nx.to_list()
|> IO.inspect()

MyOps.diag2(
  Nx.tensor([[3, 2], [4, 5]]),
  Nx.tensor([[1, -1], [1, 3]])
)
|> Nx.to_list()
|> IO.inspect()

MyOps.diag3(
  Nx.tensor([[3, 2], [4, 5]]),
  Nx.tensor([[1, -1], [1, 3]])
)
|> Nx.to_list()
|> IO.inspect()

Kino.nothing()
MyOps.weighted_trace(
  Nx.tensor([[3, 2], [4, 5]]),
  Nx.tensor([[1, -1], [1, 3]]),
  Nx.tensor([1, 2])
)
|> Nx.to_number()
|> IO.inspect()

Python Numpy reference

import numpy as np

print(np.einsum("ij,ji->i",
  np.array([[3,2],[4,5]]),
  np.array([[1,-1],[1,3]])
))

print(np.einsum(
  "ij,ji,j->",
  np.array([[3,2],[4,5]]),
  np.array([[1,-1],[1,3]]),
  np.array([1,2])
))

print(np.einsum("ij,ji->",
  np.array([[3,2],[4,5]]),
  np.array([[1,-1],[1,3]])
))

print(np.einsum("ij,ji->ji",
  np.array([[3,2],[4,5]]),
  np.array([[1,-1],[1,3]])
))

print(np.einsum("ii,jj->",
  np.array([[3,2],[4,5]]),
  np.array([[1,-1],[1,3]])
))

print(np.einsum("ii,jj->i",
  np.array([[3,2],[4,5]]),
  np.array([[1,-1],[1,3]])
))

print(np.einsum("ii,jj->j",
  np.array([[3,2],[4,5]]),
  np.array([[1,-1],[1,3]])
))

print(np.einsum("ii,jj->ij",
  np.array([[3,2],[4,5]]),
  np.array([[1,-1],[1,3]])
))

...