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"]
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)
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()
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]])
))
...