munkres 2.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- munkres/__init__.py +104 -0
- munkres/__main__.py +64 -0
- munkres/_analysis.py +185 -0
- munkres/_api.py +298 -0
- munkres/_core.py +558 -0
- munkres/_extras.py +236 -0
- munkres/_trace.py +91 -0
- munkres/cli.py +158 -0
- munkres/py.typed +0 -0
- munkres-2.0.0.dist-info/METADATA +308 -0
- munkres-2.0.0.dist-info/RECORD +16 -0
- munkres-2.0.0.dist-info/WHEEL +5 -0
- munkres-2.0.0.dist-info/entry_points.txt +2 -0
- munkres-2.0.0.dist-info/licenses/LICENSE.md +14 -0
- munkres-2.0.0.dist-info/licenses/NOTICE +10 -0
- munkres-2.0.0.dist-info/top_level.txt +1 -0
munkres/__init__.py
ADDED
|
@@ -0,0 +1,104 @@
|
|
|
1
|
+
# Copyright (c) 2008-2020 Brian M. Clapper (original author)
|
|
2
|
+
# Copyright (c) 2026 Eishit Nigam (modifications for 2.x)
|
|
3
|
+
# Licensed under the Apache License, Version 2.0. See LICENSE.md and NOTICE.
|
|
4
|
+
"""
|
|
5
|
+
Munkres: the Hungarian (Kuhn-Munkres) algorithm for the assignment problem.
|
|
6
|
+
|
|
7
|
+
Given a cost for every (row, column) pairing, find the one-to-one assignment
|
|
8
|
+
with the lowest total cost::
|
|
9
|
+
|
|
10
|
+
>>> from munkres import Munkres
|
|
11
|
+
>>> cost = [[4, 1, 3],
|
|
12
|
+
... [2, 0, 5],
|
|
13
|
+
... [3, 2, 2]]
|
|
14
|
+
>>> Munkres().compute(cost)
|
|
15
|
+
[(0, 1), (1, 0), (2, 2)]
|
|
16
|
+
|
|
17
|
+
Forbid a pairing with `DISALLOWED`; turn a profit matrix into a cost matrix
|
|
18
|
+
with `make_cost_matrix`. If the forbidden cells make a complete assignment
|
|
19
|
+
impossible, `UnsolvableMatrix` is raised at once (it never hangs).
|
|
20
|
+
|
|
21
|
+
The package has no runtime dependencies. numpy arrays and pandas DataFrames
|
|
22
|
+
are accepted as input if you have them, and are never modified.
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
from munkres._analysis import (
|
|
26
|
+
Prices,
|
|
27
|
+
bottleneck,
|
|
28
|
+
counterfactual,
|
|
29
|
+
k_best,
|
|
30
|
+
shadow_prices,
|
|
31
|
+
tolerance,
|
|
32
|
+
)
|
|
33
|
+
from munkres._api import (
|
|
34
|
+
Assignment,
|
|
35
|
+
Diagnosis,
|
|
36
|
+
build_cost_matrix,
|
|
37
|
+
diagnose,
|
|
38
|
+
linear_sum_assignment,
|
|
39
|
+
solve,
|
|
40
|
+
)
|
|
41
|
+
from munkres._core import (
|
|
42
|
+
DISALLOWED,
|
|
43
|
+
DISALLOWED_OBJ,
|
|
44
|
+
DISALLOWED_PRINTVAL,
|
|
45
|
+
AnyNum,
|
|
46
|
+
Cell,
|
|
47
|
+
Matrix,
|
|
48
|
+
MatrixLike,
|
|
49
|
+
Munkres,
|
|
50
|
+
Number,
|
|
51
|
+
UnsolvableMatrix,
|
|
52
|
+
make_cost_matrix,
|
|
53
|
+
print_matrix,
|
|
54
|
+
)
|
|
55
|
+
from munkres._extras import (
|
|
56
|
+
SinkhornResult,
|
|
57
|
+
Transport,
|
|
58
|
+
sinkhorn,
|
|
59
|
+
soft_assignment,
|
|
60
|
+
stable_matching,
|
|
61
|
+
transport,
|
|
62
|
+
)
|
|
63
|
+
from munkres._trace import Trace
|
|
64
|
+
|
|
65
|
+
__all__ = [
|
|
66
|
+
"DISALLOWED",
|
|
67
|
+
"DISALLOWED_OBJ",
|
|
68
|
+
"DISALLOWED_PRINTVAL",
|
|
69
|
+
"AnyNum",
|
|
70
|
+
"Assignment",
|
|
71
|
+
"Cell",
|
|
72
|
+
"Diagnosis",
|
|
73
|
+
"Matrix",
|
|
74
|
+
"MatrixLike",
|
|
75
|
+
"Munkres",
|
|
76
|
+
"Number",
|
|
77
|
+
"Prices",
|
|
78
|
+
"SinkhornResult",
|
|
79
|
+
"Trace",
|
|
80
|
+
"Transport",
|
|
81
|
+
"UnsolvableMatrix",
|
|
82
|
+
"bottleneck",
|
|
83
|
+
"build_cost_matrix",
|
|
84
|
+
"counterfactual",
|
|
85
|
+
"diagnose",
|
|
86
|
+
"k_best",
|
|
87
|
+
"linear_sum_assignment",
|
|
88
|
+
"make_cost_matrix",
|
|
89
|
+
"print_matrix",
|
|
90
|
+
"shadow_prices",
|
|
91
|
+
"sinkhorn",
|
|
92
|
+
"soft_assignment",
|
|
93
|
+
"solve",
|
|
94
|
+
"stable_matching",
|
|
95
|
+
"tolerance",
|
|
96
|
+
"transport",
|
|
97
|
+
]
|
|
98
|
+
|
|
99
|
+
__version__ = "2.0.0"
|
|
100
|
+
__author__ = "Brian Clapper, bmc@clapper.org"
|
|
101
|
+
__maintainer__ = "Eishit Nigam"
|
|
102
|
+
__url__ = "https://github.com/iameishit/Paldita-munkres"
|
|
103
|
+
__copyright__ = "(c) 2008-2020 Brian M. Clapper; (c) 2026 Eishit Nigam"
|
|
104
|
+
__license__ = "Apache-2.0"
|
munkres/__main__.py
ADDED
|
@@ -0,0 +1,64 @@
|
|
|
1
|
+
# Copyright (c) 2008-2020 Brian M. Clapper; (c) 2026 Eishit Nigam
|
|
2
|
+
# Licensed under the Apache License, Version 2.0. See LICENSE.md and NOTICE.
|
|
3
|
+
"""``python -m munkres`` -- with no arguments, solve a set of built-in example
|
|
4
|
+
matrices and check each answer against its known optimum (an installation smoke
|
|
5
|
+
test); with arguments, the command line interface (see ``python -m munkres -h``)."""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
from munkres import DISALLOWED, Munkres, print_matrix
|
|
12
|
+
|
|
13
|
+
D = DISALLOWED
|
|
14
|
+
|
|
15
|
+
EXAMPLES: list[tuple[list[list[Any]], float]] = [
|
|
16
|
+
# Square
|
|
17
|
+
([[400, 150, 400], [400, 450, 600], [300, 225, 300]], 850),
|
|
18
|
+
# Rectangular variant
|
|
19
|
+
([[400, 150, 400, 1], [400, 450, 600, 2], [300, 225, 300, 3]], 452),
|
|
20
|
+
# Square
|
|
21
|
+
([[10, 10, 8], [9, 8, 1], [9, 7, 4]], 18),
|
|
22
|
+
# Square variant with floating point value
|
|
23
|
+
([[10.1, 10.2, 8.3], [9.4, 8.5, 1.6], [9.7, 7.8, 4.9]], 19.5),
|
|
24
|
+
# Rectangular variant
|
|
25
|
+
([[10, 10, 8, 11], [9, 8, 1, 1], [9, 7, 4, 10]], 15),
|
|
26
|
+
# Rectangular variant with floating point value
|
|
27
|
+
([[10.01, 10.02, 8.03, 11.04], [9.05, 8.06, 1.07, 1.08], [9.09, 7.1, 4.11, 10.12]], 15.2),
|
|
28
|
+
# Rectangular with DISALLOWED
|
|
29
|
+
([[4, 5, 6, D], [1, 9, 12, 11], [D, 5, 4, D], [12, 12, 12, 10]], 20),
|
|
30
|
+
# Rectangular variant with DISALLOWED and floating point value
|
|
31
|
+
(
|
|
32
|
+
[
|
|
33
|
+
[4.001, 5.002, 6.003, D],
|
|
34
|
+
[1.004, 9.005, 12.006, 11.007],
|
|
35
|
+
[D, 5.008, 4.009, D],
|
|
36
|
+
[12.01, 12.011, 12.012, 10.013],
|
|
37
|
+
],
|
|
38
|
+
20.028,
|
|
39
|
+
),
|
|
40
|
+
# DISALLOWED to force pairings
|
|
41
|
+
([[1, D, D, D], [D, 2, D, D], [D, D, 3, D], [D, D, D, 4]], 10),
|
|
42
|
+
# DISALLOWED to force pairings with floating point value
|
|
43
|
+
([[1.1, D, D, D], [D, 2.2, D, D], [D, D, 3.3, D], [D, D, D, 4.4]], 11.0),
|
|
44
|
+
]
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def main() -> None:
|
|
48
|
+
solver = Munkres()
|
|
49
|
+
for cost_matrix, expected_total in EXAMPLES:
|
|
50
|
+
print_matrix(cost_matrix, msg="cost matrix")
|
|
51
|
+
total_cost: Any = 0
|
|
52
|
+
for r, c in solver.compute(cost_matrix):
|
|
53
|
+
value = cost_matrix[r][c]
|
|
54
|
+
total_cost += value
|
|
55
|
+
print(f"({r}, {c}) -> {value}")
|
|
56
|
+
print(f"lowest cost={total_cost}")
|
|
57
|
+
if expected_total != total_cost:
|
|
58
|
+
raise SystemExit(f"self-check failed: expected {expected_total}, got {total_cost}")
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
if __name__ == "__main__":
|
|
62
|
+
from munkres.cli import main as cli_main
|
|
63
|
+
|
|
64
|
+
raise SystemExit(cli_main())
|
munkres/_analysis.py
ADDED
|
@@ -0,0 +1,185 @@
|
|
|
1
|
+
# Copyright (c) 2026 Eishit Nigam. Licensed under the Apache License, Version 2.0.
|
|
2
|
+
"""Analysis of an assignment problem: prices, what-ifs, runners-up, bottleneck."""
|
|
3
|
+
|
|
4
|
+
from __future__ import annotations
|
|
5
|
+
|
|
6
|
+
import heapq
|
|
7
|
+
import itertools
|
|
8
|
+
from dataclasses import dataclass
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
from munkres._api import Assignment, _make_assignment, _negated, solve
|
|
12
|
+
from munkres._core import (
|
|
13
|
+
DISALLOWED,
|
|
14
|
+
MatrixLike,
|
|
15
|
+
UnsolvableMatrix,
|
|
16
|
+
_as_lists,
|
|
17
|
+
_assign,
|
|
18
|
+
_validated_edges,
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
__all__ = ["Prices", "bottleneck", "counterfactual", "k_best", "shadow_prices", "tolerance"]
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
@dataclass(frozen=True)
|
|
25
|
+
class Prices:
|
|
26
|
+
"""
|
|
27
|
+
LP dual variables of the assignment problem ("shadow prices").
|
|
28
|
+
|
|
29
|
+
For every allowed cell, `cost[i][j] >= row_prices[i] + col_prices[j]`, with
|
|
30
|
+
equality on every pair of the optimal assignment, and
|
|
31
|
+
`total == sum(row_prices) + sum(col_prices) ==` the optimal cost. That
|
|
32
|
+
equality is a *proof of optimality* anyone can re-check in O(rows * columns).
|
|
33
|
+
A row price says how much one more unit of "demand" at that row would cost.
|
|
34
|
+
"""
|
|
35
|
+
|
|
36
|
+
row_prices: tuple[Any, ...]
|
|
37
|
+
col_prices: tuple[Any, ...]
|
|
38
|
+
total: Any
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def shadow_prices(matrix: MatrixLike) -> Prices:
|
|
42
|
+
"""Solve the problem and return its dual variables (see `Prices`)."""
|
|
43
|
+
n_rows, n_cols, edges = _validated_edges(matrix)
|
|
44
|
+
duals: list[Any] = []
|
|
45
|
+
_assign(n_rows, n_cols, edges, duals=duals)
|
|
46
|
+
u, v = duals
|
|
47
|
+
rows_p, cols_p = (u, v) if n_rows <= n_cols else (v, u)
|
|
48
|
+
return Prices(tuple(rows_p), tuple(cols_p), sum(rows_p) + sum(cols_p))
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _check_cell(matrix: list[list[Any]], row: int, col: int) -> None:
|
|
52
|
+
n_rows, n_cols, edges = _validated_edges(matrix)
|
|
53
|
+
if not (0 <= row < n_rows and 0 <= col < n_cols):
|
|
54
|
+
raise ValueError(f"cell ({row}, {col}) is outside the {n_rows}x{n_cols} matrix")
|
|
55
|
+
if all(c != col for c, _ in edges[row]):
|
|
56
|
+
raise ValueError(f"cell ({row}, {col}) is forbidden")
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def counterfactual(matrix: MatrixLike, row: int, col: int) -> Any:
|
|
60
|
+
"""
|
|
61
|
+
"What would it cost to force `row` onto `col`?" Returns how much the optimal
|
|
62
|
+
total rises if that pair is made mandatory (0 if some optimal assignment
|
|
63
|
+
already contains it). Raises `UnsolvableMatrix` if forcing it makes the
|
|
64
|
+
problem impossible, and `ValueError` for a forbidden or out-of-range cell.
|
|
65
|
+
"""
|
|
66
|
+
rows = _as_lists(matrix)
|
|
67
|
+
_check_cell(rows, row, col)
|
|
68
|
+
best = solve(rows).total
|
|
69
|
+
forced = [list(r) for r in rows]
|
|
70
|
+
for j in range(len(rows[0])):
|
|
71
|
+
if j != col:
|
|
72
|
+
forced[row][j] = DISALLOWED
|
|
73
|
+
for i in range(len(rows)):
|
|
74
|
+
if i != row:
|
|
75
|
+
forced[i][col] = DISALLOWED
|
|
76
|
+
return solve(forced).total - best
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def tolerance(matrix: MatrixLike, row: int, col: int) -> Any:
|
|
80
|
+
"""
|
|
81
|
+
How much can the cost of a pair in the optimal assignment *rise* before the
|
|
82
|
+
assignment has to change? (the optimum without that pair, minus the optimum).
|
|
83
|
+
`float('inf')` if the pair can never be replaced. `ValueError` if the pair is
|
|
84
|
+
not in the assignment `solve` returns.
|
|
85
|
+
"""
|
|
86
|
+
rows = _as_lists(matrix)
|
|
87
|
+
_check_cell(rows, row, col)
|
|
88
|
+
chosen = solve(rows)
|
|
89
|
+
if (row, col) not in chosen.pairs:
|
|
90
|
+
raise ValueError(f"({row}, {col}) is not in the optimal assignment")
|
|
91
|
+
without = [list(r) for r in rows]
|
|
92
|
+
without[row][col] = DISALLOWED
|
|
93
|
+
try:
|
|
94
|
+
return solve(without).total - chosen.total
|
|
95
|
+
except UnsolvableMatrix:
|
|
96
|
+
return float("inf")
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def k_best(matrix: MatrixLike, k: int, *, maximize: bool = False) -> list[Assignment]:
|
|
100
|
+
"""
|
|
101
|
+
The `k` best complete assignments, best first (Murty's algorithm). Fewer are
|
|
102
|
+
returned if fewer exist. Each is a distinct set of pairs; ties are ordered
|
|
103
|
+
arbitrarily.
|
|
104
|
+
"""
|
|
105
|
+
if k < 0:
|
|
106
|
+
raise ValueError("k must be >= 0")
|
|
107
|
+
original = _as_lists(matrix)
|
|
108
|
+
work = _negated(original) if maximize else original
|
|
109
|
+
n_rows, n_cols, _ = _validated_edges(work)
|
|
110
|
+
found: list[Assignment] = []
|
|
111
|
+
if k == 0:
|
|
112
|
+
return found
|
|
113
|
+
|
|
114
|
+
def constrained(forbid: tuple[tuple[int, int], ...], force: tuple[tuple[int, int], ...]) -> Any:
|
|
115
|
+
grid = [list(r) for r in work]
|
|
116
|
+
for i, j in forbid:
|
|
117
|
+
grid[i][j] = DISALLOWED
|
|
118
|
+
for i, j in force:
|
|
119
|
+
for jj in range(n_cols):
|
|
120
|
+
if jj != j:
|
|
121
|
+
grid[i][jj] = DISALLOWED
|
|
122
|
+
for ii in range(n_rows):
|
|
123
|
+
if ii != i:
|
|
124
|
+
grid[ii][j] = DISALLOWED
|
|
125
|
+
try:
|
|
126
|
+
return solve(grid)
|
|
127
|
+
except UnsolvableMatrix:
|
|
128
|
+
return None
|
|
129
|
+
|
|
130
|
+
counter = itertools.count()
|
|
131
|
+
heap: list[Any] = []
|
|
132
|
+
first = constrained((), ())
|
|
133
|
+
if first is not None:
|
|
134
|
+
heapq.heappush(heap, (first.total, next(counter), first, (), ()))
|
|
135
|
+
while heap and len(found) < k:
|
|
136
|
+
_, _, node, forbid, force = heapq.heappop(heap)
|
|
137
|
+
pairs = list(node.pairs)
|
|
138
|
+
found.append(_make_assignment(original, pairs, (n_rows, n_cols), maximize=maximize))
|
|
139
|
+
for t, pair in enumerate(pairs):
|
|
140
|
+
child_force = force + tuple(pairs[:t])
|
|
141
|
+
child_forbid = (*forbid, pair)
|
|
142
|
+
child = constrained(child_forbid, child_force)
|
|
143
|
+
if child is not None:
|
|
144
|
+
heapq.heappush(heap, (child.total, next(counter), child, child_forbid, child_force))
|
|
145
|
+
return found
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def bottleneck(matrix: MatrixLike, *, maximize: bool = False) -> Assignment:
|
|
149
|
+
"""
|
|
150
|
+
The assignment whose *single worst pair* is as good as possible (minimise the
|
|
151
|
+
largest cost; with `maximize=True`, maximise the smallest profit). Among
|
|
152
|
+
assignments that tie on the worst pair, the lowest total (highest, when
|
|
153
|
+
maximising) wins. Good for fairness: "nobody gets a terrible job".
|
|
154
|
+
Raises `UnsolvableMatrix` if no complete assignment exists.
|
|
155
|
+
"""
|
|
156
|
+
original = _as_lists(matrix)
|
|
157
|
+
work = _negated(original) if maximize else original
|
|
158
|
+
n_rows, n_cols, edges = _validated_edges(work)
|
|
159
|
+
shape = (n_rows, n_cols)
|
|
160
|
+
values = sorted({c for allowed in edges for _, c in allowed})
|
|
161
|
+
if not values:
|
|
162
|
+
_assign(n_rows, n_cols, edges) # raises if rows exist but nothing is allowed
|
|
163
|
+
return _make_assignment(original, [], shape, maximize=maximize)
|
|
164
|
+
|
|
165
|
+
def capped(limit: Any) -> list[list[tuple[int, Any]]]:
|
|
166
|
+
return [[(j, c) for j, c in allowed if c <= limit] for allowed in edges]
|
|
167
|
+
|
|
168
|
+
def feasible(limit: Any) -> bool:
|
|
169
|
+
try:
|
|
170
|
+
_assign(n_rows, n_cols, capped(limit))
|
|
171
|
+
except UnsolvableMatrix:
|
|
172
|
+
return False
|
|
173
|
+
return True
|
|
174
|
+
|
|
175
|
+
if not feasible(values[-1]):
|
|
176
|
+
_assign(n_rows, n_cols, edges) # raise the real, explanatory error
|
|
177
|
+
lo, hi = 0, len(values) - 1
|
|
178
|
+
while lo < hi:
|
|
179
|
+
mid = (lo + hi) // 2
|
|
180
|
+
if feasible(values[mid]):
|
|
181
|
+
hi = mid
|
|
182
|
+
else:
|
|
183
|
+
lo = mid + 1
|
|
184
|
+
pairs = _assign(n_rows, n_cols, capped(values[lo]))
|
|
185
|
+
return _make_assignment(original, pairs, shape, maximize=maximize)
|
munkres/_api.py
ADDED
|
@@ -0,0 +1,298 @@
|
|
|
1
|
+
# Copyright (c) 2026 Eishit Nigam. Licensed under the Apache License, Version 2.0.
|
|
2
|
+
"""The 2.x high-level API: `solve`, `Assignment`, `diagnose`, `linear_sum_assignment`."""
|
|
3
|
+
|
|
4
|
+
from __future__ import annotations
|
|
5
|
+
|
|
6
|
+
import importlib
|
|
7
|
+
from collections.abc import Callable, Iterable, Iterator
|
|
8
|
+
from dataclasses import dataclass
|
|
9
|
+
from typing import Any, TypeVar
|
|
10
|
+
|
|
11
|
+
from munkres._core import (
|
|
12
|
+
_INF,
|
|
13
|
+
_NUMBER_TYPES,
|
|
14
|
+
DISALLOWED,
|
|
15
|
+
MatrixLike,
|
|
16
|
+
UnsolvableMatrix,
|
|
17
|
+
_as_lists,
|
|
18
|
+
_assign,
|
|
19
|
+
_validated_edges,
|
|
20
|
+
)
|
|
21
|
+
from munkres._trace import Trace
|
|
22
|
+
|
|
23
|
+
__all__ = [
|
|
24
|
+
"Assignment",
|
|
25
|
+
"Diagnosis",
|
|
26
|
+
"build_cost_matrix",
|
|
27
|
+
"diagnose",
|
|
28
|
+
"linear_sum_assignment",
|
|
29
|
+
"solve",
|
|
30
|
+
]
|
|
31
|
+
|
|
32
|
+
_A = TypeVar("_A")
|
|
33
|
+
_B = TypeVar("_B")
|
|
34
|
+
_NO_PAIRS: Any = object() # sentinel: a threshold so strict that nothing can be matched
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
@dataclass(frozen=True)
|
|
38
|
+
class Assignment:
|
|
39
|
+
"""
|
|
40
|
+
The result of `solve`.
|
|
41
|
+
|
|
42
|
+
- `pairs`: the matched `(row, column)` index pairs, sorted by row
|
|
43
|
+
- `total`: the sum of *your original* matrix values over `pairs`
|
|
44
|
+
(the total cost, or the total profit when `maximize=True`)
|
|
45
|
+
- `unmatched_rows` / `unmatched_cols`: indexes left without a partner
|
|
46
|
+
(always the surplus side of a rectangular matrix, plus anything gated out)
|
|
47
|
+
- `shape`: `(rows, columns)` of the input
|
|
48
|
+
- `row_labels` / `col_labels`: the DataFrame's index / columns, if you passed one
|
|
49
|
+
- `trace`: a `Trace` if you asked for one
|
|
50
|
+
|
|
51
|
+
Iterating an `Assignment` yields its pairs, so it can stand in for the list
|
|
52
|
+
`Munkres().compute()` returns. Note that an empty one is falsy.
|
|
53
|
+
"""
|
|
54
|
+
|
|
55
|
+
pairs: tuple[tuple[int, int], ...]
|
|
56
|
+
total: Any
|
|
57
|
+
unmatched_rows: tuple[int, ...]
|
|
58
|
+
unmatched_cols: tuple[int, ...]
|
|
59
|
+
shape: tuple[int, int]
|
|
60
|
+
maximize: bool = False
|
|
61
|
+
row_labels: tuple[Any, ...] | None = None
|
|
62
|
+
col_labels: tuple[Any, ...] | None = None
|
|
63
|
+
trace: Trace | None = None
|
|
64
|
+
|
|
65
|
+
def __iter__(self) -> Iterator[tuple[int, int]]:
|
|
66
|
+
return iter(self.pairs)
|
|
67
|
+
|
|
68
|
+
def __len__(self) -> int:
|
|
69
|
+
return len(self.pairs)
|
|
70
|
+
|
|
71
|
+
@property
|
|
72
|
+
def rows(self) -> tuple[int, ...]:
|
|
73
|
+
"""Matched row indexes, in order (like SciPy's `row_ind`)."""
|
|
74
|
+
return tuple(r for r, _ in self.pairs)
|
|
75
|
+
|
|
76
|
+
@property
|
|
77
|
+
def cols(self) -> tuple[int, ...]:
|
|
78
|
+
"""Matched column indexes, aligned with `rows` (like SciPy's `col_ind`)."""
|
|
79
|
+
return tuple(c for _, c in self.pairs)
|
|
80
|
+
|
|
81
|
+
def as_dict(self) -> dict[int, int]:
|
|
82
|
+
"""`{row: column}` for every matched row."""
|
|
83
|
+
return dict(self.pairs)
|
|
84
|
+
|
|
85
|
+
def labelled(self) -> list[tuple[Any, Any]]:
|
|
86
|
+
"""The pairs as `(row label, column label)`; plain indexes if there are no labels."""
|
|
87
|
+
return [
|
|
88
|
+
(
|
|
89
|
+
self.row_labels[r] if self.row_labels is not None else r,
|
|
90
|
+
self.col_labels[c] if self.col_labels is not None else c,
|
|
91
|
+
)
|
|
92
|
+
for r, c in self.pairs
|
|
93
|
+
]
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
@dataclass(frozen=True)
|
|
97
|
+
class Diagnosis:
|
|
98
|
+
"""Why a matrix has no complete assignment (a violation of Hall's condition)."""
|
|
99
|
+
|
|
100
|
+
message: str
|
|
101
|
+
rows: tuple[int, ...]
|
|
102
|
+
cols: tuple[int, ...]
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def _negated(rows: list[list[Any]]) -> list[list[Any]]:
|
|
106
|
+
"""Profit -> cost by negation (exact for int, Fraction and Decimal)."""
|
|
107
|
+
out: list[list[Any]] = []
|
|
108
|
+
for i, row in enumerate(rows):
|
|
109
|
+
new: list[Any] = []
|
|
110
|
+
for j, value in enumerate(row):
|
|
111
|
+
if value is DISALLOWED or not isinstance(value, _NUMBER_TYPES):
|
|
112
|
+
new.append(value) # forbidden, or a TypeError for validation to report
|
|
113
|
+
continue
|
|
114
|
+
try:
|
|
115
|
+
infinite_profit = value == _INF
|
|
116
|
+
except ArithmeticError: # Decimal('sNaN'): validation reports it
|
|
117
|
+
new.append(value)
|
|
118
|
+
continue
|
|
119
|
+
if infinite_profit:
|
|
120
|
+
raise ValueError(
|
|
121
|
+
f"cell [{i}][{j}] is +infinity, which has no defined assignment value "
|
|
122
|
+
"(use -inf or DISALLOWED to forbid a pairing)"
|
|
123
|
+
)
|
|
124
|
+
new.append(-value)
|
|
125
|
+
out.append(new)
|
|
126
|
+
return out
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
def _gate(threshold: Any, *, maximize: bool) -> Any:
|
|
130
|
+
"""Turn the user's threshold into a cost-space gate (None = no gating)."""
|
|
131
|
+
if threshold is None:
|
|
132
|
+
return None
|
|
133
|
+
name = "min_profit" if maximize else "max_cost"
|
|
134
|
+
if not isinstance(threshold, _NUMBER_TYPES):
|
|
135
|
+
raise TypeError(f"{name} must be a number, got {threshold!r}")
|
|
136
|
+
if threshold != threshold: # noqa: PLR0124
|
|
137
|
+
raise ValueError(f"{name} is NaN")
|
|
138
|
+
gate = -threshold if maximize else threshold
|
|
139
|
+
if gate == _INF:
|
|
140
|
+
return None
|
|
141
|
+
if gate == -_INF:
|
|
142
|
+
return _NO_PAIRS
|
|
143
|
+
return gate
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def _labels(matrix: Any) -> tuple[tuple[Any, ...] | None, tuple[Any, ...] | None]:
|
|
147
|
+
index: Any = getattr(matrix, "index", None)
|
|
148
|
+
columns: Any = getattr(matrix, "columns", None)
|
|
149
|
+
if hasattr(index, "tolist") and hasattr(columns, "tolist"): # a DataFrame, not a list
|
|
150
|
+
return tuple(index.tolist()), tuple(columns.tolist())
|
|
151
|
+
return None, None
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def solve(
|
|
155
|
+
matrix: MatrixLike,
|
|
156
|
+
*,
|
|
157
|
+
maximize: bool = False,
|
|
158
|
+
max_cost: Any = None,
|
|
159
|
+
min_profit: Any = None,
|
|
160
|
+
trace: bool = False,
|
|
161
|
+
) -> Assignment:
|
|
162
|
+
"""
|
|
163
|
+
Solve an assignment problem and describe the answer.
|
|
164
|
+
|
|
165
|
+
**Parameters**
|
|
166
|
+
|
|
167
|
+
- `matrix`: costs (or profits, with `maximize=True`) as nested sequences, a
|
|
168
|
+
numpy array or a pandas DataFrame. `DISALLOWED` and `+inf` (`-inf` when
|
|
169
|
+
maximising) forbid a pairing.
|
|
170
|
+
- `maximize`: maximise the total instead of minimising it. Exact for int,
|
|
171
|
+
`Fraction` and `Decimal`, because values are negated rather than subtracted
|
|
172
|
+
from a maximum.
|
|
173
|
+
- `max_cost`: **gating** (minimising only). Pairs costing more than this are
|
|
174
|
+
never made, and a row may be left unmatched instead. The result minimises
|
|
175
|
+
the cost of the pairs made plus `max_cost` for every row left unmatched, so
|
|
176
|
+
a pair is only worth making if it costs less than `max_cost`. A pair
|
|
177
|
+
costing exactly `max_cost` is a tie with leaving the row unmatched. With
|
|
178
|
+
gating a matrix can never be unsolvable.
|
|
179
|
+
- `min_profit`: the same gate when `maximize=True`: pairs earning less than
|
|
180
|
+
this are never made.
|
|
181
|
+
- `trace`: record the solver's steps in `Assignment.trace`.
|
|
182
|
+
|
|
183
|
+
**Raises**
|
|
184
|
+
|
|
185
|
+
- `UnsolvableMatrix`: forbidden cells make a complete assignment impossible
|
|
186
|
+
(not possible with gating)
|
|
187
|
+
- `ValueError`: ragged matrix, `NaN`, an infinite value with no meaning, a
|
|
188
|
+
threshold that does not fit the mode
|
|
189
|
+
- `TypeError`: a non-numeric cell or threshold
|
|
190
|
+
"""
|
|
191
|
+
if maximize and max_cost is not None:
|
|
192
|
+
raise ValueError("max_cost is for minimising; use min_profit with maximize=True")
|
|
193
|
+
if not maximize and min_profit is not None:
|
|
194
|
+
raise ValueError("min_profit needs maximize=True; use max_cost when minimising")
|
|
195
|
+
gate = _gate(min_profit if maximize else max_cost, maximize=maximize)
|
|
196
|
+
|
|
197
|
+
row_labels, col_labels = _labels(matrix)
|
|
198
|
+
original = _as_lists(matrix)
|
|
199
|
+
n_rows, n_cols, edges = _validated_edges(_negated(original) if maximize else original)
|
|
200
|
+
|
|
201
|
+
recorder = Trace(shape=(n_rows, n_cols)) if trace else None
|
|
202
|
+
if gate is _NO_PAIRS:
|
|
203
|
+
pairs: list[tuple[int, int]] = []
|
|
204
|
+
if recorder is not None:
|
|
205
|
+
recorder.events.append({"event": "done", "pairs": []})
|
|
206
|
+
else:
|
|
207
|
+
pairs = _assign(n_rows, n_cols, edges, gate=gate, trace=recorder)
|
|
208
|
+
|
|
209
|
+
return _make_assignment(
|
|
210
|
+
original, pairs, (n_rows, n_cols), maximize=maximize,
|
|
211
|
+
labels=(row_labels, col_labels), trace=recorder,
|
|
212
|
+
) # fmt: skip
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
def _make_assignment(
|
|
216
|
+
original: list[list[Any]],
|
|
217
|
+
pairs: list[tuple[int, int]],
|
|
218
|
+
shape: tuple[int, int],
|
|
219
|
+
*,
|
|
220
|
+
maximize: bool,
|
|
221
|
+
labels: tuple[tuple[Any, ...] | None, tuple[Any, ...] | None] = (None, None),
|
|
222
|
+
trace: Trace | None = None,
|
|
223
|
+
) -> Assignment:
|
|
224
|
+
n_rows, n_cols = shape
|
|
225
|
+
matched_rows = {r for r, _ in pairs}
|
|
226
|
+
matched_cols = {c for _, c in pairs}
|
|
227
|
+
return Assignment(
|
|
228
|
+
pairs=tuple(pairs),
|
|
229
|
+
total=sum(original[r][c] for r, c in pairs),
|
|
230
|
+
unmatched_rows=tuple(i for i in range(n_rows) if i not in matched_rows),
|
|
231
|
+
unmatched_cols=tuple(j for j in range(n_cols) if j not in matched_cols),
|
|
232
|
+
shape=shape,
|
|
233
|
+
maximize=maximize,
|
|
234
|
+
row_labels=labels[0],
|
|
235
|
+
col_labels=labels[1],
|
|
236
|
+
trace=trace,
|
|
237
|
+
)
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
def diagnose(matrix: MatrixLike, *, maximize: bool = False) -> Diagnosis | None:
|
|
241
|
+
"""
|
|
242
|
+
Explain why a matrix cannot be solved, without raising.
|
|
243
|
+
|
|
244
|
+
Returns `None` if a complete assignment exists. Otherwise returns a
|
|
245
|
+
`Diagnosis` naming rows that can only use too few columns (or the reverse for
|
|
246
|
+
a tall matrix), which is exactly Hall's marriage-theorem violation.
|
|
247
|
+
"""
|
|
248
|
+
original = _as_lists(matrix)
|
|
249
|
+
n_rows, n_cols, edges = _validated_edges(_negated(original) if maximize else original)
|
|
250
|
+
try:
|
|
251
|
+
_assign(n_rows, n_cols, edges)
|
|
252
|
+
except UnsolvableMatrix as bad:
|
|
253
|
+
return Diagnosis(str(bad), bad.rows, bad.cols)
|
|
254
|
+
return None
|
|
255
|
+
|
|
256
|
+
|
|
257
|
+
def linear_sum_assignment(cost_matrix: MatrixLike, maximize: bool = False) -> tuple[Any, Any]:
|
|
258
|
+
"""
|
|
259
|
+
A drop-in for `scipy.optimize.linear_sum_assignment`: returns
|
|
260
|
+
`(row_ind, col_ind)`, sorted by row.
|
|
261
|
+
|
|
262
|
+
The return type follows the input: numpy arrays (and DataFrames) get numpy
|
|
263
|
+
index arrays back; anything else gets lists. As in SciPy, `+inf` forbids a
|
|
264
|
+
pairing, an infeasible matrix raises `ValueError` (here the subclass
|
|
265
|
+
`UnsolvableMatrix`), and when several optimal answers exist the one returned
|
|
266
|
+
may differ from SciPy's (the total cost is always the same).
|
|
267
|
+
"""
|
|
268
|
+
try:
|
|
269
|
+
result = solve(cost_matrix, maximize=maximize)
|
|
270
|
+
except TypeError as bad:
|
|
271
|
+
raise ValueError(str(bad)) from None
|
|
272
|
+
rows, cols = list(result.rows), list(result.cols)
|
|
273
|
+
if hasattr(cost_matrix, "tolist") or hasattr(cost_matrix, "to_numpy"):
|
|
274
|
+
try:
|
|
275
|
+
np = importlib.import_module("numpy")
|
|
276
|
+
except ImportError:
|
|
277
|
+
return rows, cols
|
|
278
|
+
return np.array(rows, dtype=np.intp), np.array(cols, dtype=np.intp)
|
|
279
|
+
return rows, cols
|
|
280
|
+
|
|
281
|
+
|
|
282
|
+
def build_cost_matrix(
|
|
283
|
+
rows: Iterable[_A], cols: Iterable[_B], cost: Callable[[_A, _B], Any]
|
|
284
|
+
) -> list[list[Any]]:
|
|
285
|
+
"""
|
|
286
|
+
Build a matrix by calling `cost(row_item, col_item)` for every pair, e.g. a
|
|
287
|
+
distance between points. If `cost` returns `None` that pairing is
|
|
288
|
+
`DISALLOWED`.
|
|
289
|
+
"""
|
|
290
|
+
col_items = list(cols)
|
|
291
|
+
matrix: list[list[Any]] = []
|
|
292
|
+
for a in rows:
|
|
293
|
+
row: list[Any] = []
|
|
294
|
+
for b in col_items:
|
|
295
|
+
value = cost(a, b)
|
|
296
|
+
row.append(DISALLOWED if value is None else value)
|
|
297
|
+
matrix.append(row)
|
|
298
|
+
return matrix
|