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 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