Files
2026-07-04 19:27:17 +09:00

115 lines
3.5 KiB
Elixir
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
defmodule Rsh26pool.Voting.Hungarian do
@moduledoc """
Solves the assignment problem — a minimum-cost perfect matching in a weighted
bipartite graph — with the Hungarian (KuhnMunkres) algorithm in O(n^3).
This is the algorithm the proposal calls for when placing members into groups
in grouping mode: given a cost of assigning each member to each group slot, it
finds the assignment with the lowest total cost.
`min_cost_assignment/1` takes an `n x m` cost matrix (a list of `n` rows, each
a list of `m` numbers) with `n <= m`, and returns `{total_cost, assignment}`
where `assignment` is a list of length `n` and `Enum.at(assignment, i)` is the
0-based column matched to row `i`.
"""
@inf 1_000_000_000
@spec min_cost_assignment([[number()]]) :: {number(), [non_neg_integer()]}
def min_cost_assignment([]), do: {0, []}
def min_cost_assignment(cost) when is_list(cost) do
n = length(cost)
m = length(hd(cost))
if n > m do
raise ArgumentError, "cost matrix must have rows <= cols (got #{n} x #{m})"
end
# 1-indexed access into an immutable tuple-of-tuples.
c = cost |> Enum.map(&List.to_tuple/1) |> List.to_tuple()
cost_at = fn i, j -> c |> elem(i - 1) |> elem(j - 1) end
u = fill(n + 1, 0)
v = fill(m + 1, 0)
p = fill(m + 1, 0)
way = fill(m + 1, 0)
{_u, _v, p, _way} =
Enum.reduce(1..n, {u, v, p, way}, fn i, {u, v, p, way} ->
phase(i, m, cost_at, u, v, p, way)
end)
# p[j] holds the row matched to column j; invert to get column-per-row.
row_to_col =
Enum.reduce(1..m, %{}, fn j, acc ->
row = elem(p, j)
if row >= 1 and row <= n, do: Map.put(acc, row, j - 1), else: acc
end)
assignment = Enum.map(1..n, &Map.fetch!(row_to_col, &1))
total =
assignment
|> Enum.with_index(1)
|> Enum.reduce(0, fn {col0, i}, sum -> sum + cost_at.(i, col0 + 1) end)
{total, assignment}
end
# Augmenting phase for row `i` (see the classic e-maxx Hungarian description).
defp phase(i, m, cost_at, u, v, p, way) do
p = put_elem(p, 0, i)
minv = fill(m + 1, @inf)
used = fill(m + 1, false)
walk(0, m, cost_at, u, v, p, way, minv, used)
end
defp walk(j0, m, cost_at, u, v, p, way, minv, used) do
used = put_elem(used, j0, true)
i0 = elem(p, j0)
ui0 = elem(u, i0)
{delta, j1, minv, way} =
Enum.reduce(1..m, {@inf, -1, minv, way}, fn j, {delta, j1, minv, way} ->
if elem(used, j) do
{delta, j1, minv, way}
else
cur = cost_at.(i0, j) - ui0 - elem(v, j)
{minv, way} =
if cur < elem(minv, j),
do: {put_elem(minv, j, cur), put_elem(way, j, j0)},
else: {minv, way}
mvj = elem(minv, j)
if mvj < delta, do: {mvj, j, minv, way}, else: {delta, j1, minv, way}
end
end)
{u, v, minv} =
Enum.reduce(0..m, {u, v, minv}, fn j, {u, v, minv} ->
if elem(used, j) do
pj = elem(p, j)
{put_elem(u, pj, elem(u, pj) + delta), put_elem(v, j, elem(v, j) - delta), minv}
else
{u, v, put_elem(minv, j, elem(minv, j) - delta)}
end
end)
if elem(p, j1) == 0 do
{u, v, reconstruct(j1, p, way), way}
else
walk(j1, m, cost_at, u, v, p, way, minv, used)
end
end
defp reconstruct(j0, p, way) do
j1 = elem(way, j0)
p = put_elem(p, j0, elem(p, j1))
if j1 == 0, do: p, else: reconstruct(j1, p, way)
end
defp fill(size, value), do: value |> List.duplicate(size) |> List.to_tuple()
end