#!/usr/bin/env python3

# Zuordnung_Ungarisch.py
"""
Kapitel Graphen: Das Zuordnungsproblem, dreifach geloest.

  (1) Ungarischer Algorithmus (scipy.optimize.linear_sum_assignment) - O(n^3)
  (2) als LP OHNE Ganzzahligkeitsforderung -> liefert trotzdem 0/1 (Birkhoff)
  (3) als MILP MIT Ganzzahligkeitsforderung -> gleiches Ergebnis, mehr Aufwand

Zeigt damit die praktische Bedeutung der totalen Unimodularitaet.
"""

import time

import numpy as np
from scipy.optimize import linear_sum_assignment, linprog


def erzeuge_kosten(n, seed=11):
    rng = np.random.default_rng(seed)
    return rng.integers(10, 99, size=(n, n)).astype(float)


def loese_ungarisch(kosten):
    zeilen, spalten = linear_sum_assignment(kosten)
    return kosten[zeilen, spalten].sum(), spalten


def baue_lp(kosten):
    """Gemeinsame LP-Struktur fuer Variante 2 und 3."""
    n = len(kosten)
    c = kosten.flatten()                       # x_ij in Zeilenreihenfolge
    A_eq = np.zeros((2 * n, n * n))
    for i in range(n):                         # jede Person genau eine Aufgabe
        A_eq[i, i * n:(i + 1) * n] = 1.0
    for j in range(n):                         # jede Aufgabe genau einer Person
        A_eq[n + j, j::n] = 1.0
    b_eq = np.ones(2 * n)
    return c, A_eq, b_eq


def loese_lp(kosten, ganzzahlig):
    n = len(kosten)
    c, A_eq, b_eq = baue_lp(kosten)
    ergebnis = linprog(c=c, A_eq=A_eq, b_eq=b_eq, bounds=[(0, 1)] * (n * n),
                       integrality=np.ones(n * n) if ganzzahlig else None,
                       method="highs")
    x = ergebnis.x.reshape(n, n)
    return ergebnis.fun, x


if __name__ == "__main__":
    print("=" * 84)
    print("  ZUORDNUNGSPROBLEM: DREI WEGE ZUM SELBEN ERGEBNIS")
    print("=" * 84)

    # --- Kleines Beispiel zum Nachvollziehen ------------------------------
    kosten = np.array([[82., 83., 69., 92.],
                       [77., 37., 49., 92.],
                       [11., 69., 5., 86.],
                       [8., 9., 98., 23.]])
    namen = ["Anna", "Ben", "Carla", "David"]
    aufgaben = ["Auftrag W", "Auftrag X", "Auftrag Y", "Auftrag Z"]

    print("\nKostenmatrix (wer bearbeitet was zu welchen Kosten?):")
    print(f"{'':<8}" + "".join(f"{a:>12}" for a in aufgaben))
    for i, name in enumerate(namen):
        print(f"{name:<8}" + "".join(f"{kosten[i, j]:>12.0f}" for j in range(4)))

    wert, zuordnung = loese_ungarisch(kosten)
    print(f"\nOptimale Zuordnung (Gesamtkosten {wert:.0f}):")
    for i, j in enumerate(zuordnung):
        print(f"  {namen[i]:<8} -> {aufgaben[j]:<12} ({kosten[i, j]:.0f} EUR)")

    # --- Nachweis: LP ohne Ganzzahligkeit liefert trotzdem 0/1 -----------
    wert_lp, x_lp = loese_lp(kosten, ganzzahlig=False)
    ist_binaer = np.all((np.abs(x_lp) < 1e-9) | (np.abs(x_lp - 1) < 1e-9))
    print(f"\nLP OHNE Ganzzahligkeitsforderung: Kosten {wert_lp:.0f}, "
          f"Loesung ist {'0/1-wertig' if ist_binaer else 'GEBROCHEN'}")
    print("  -> Satz von Birkhoff/von Neumann bestaetigt: Die Ecken sind Permutationen.")

    # --- Laufzeitvergleich bei wachsender Groesse ------------------------
    print("\n" + "-" * 84)
    print(f"{'n':>4} | {'Ungarisch':>12} | {'LP (kontinuierlich)':>21} | "
          f"{'MILP (ganzzahlig)':>19} | {'gleich?':>8}")
    print("-" * 84)
    for n in [10, 25, 50, 100]:
        k = erzeuge_kosten(n)

        t0 = time.perf_counter(); w1, _ = loese_ungarisch(k); t1 = time.perf_counter() - t0
        t0 = time.perf_counter(); w2, _ = loese_lp(k, False);  t2 = time.perf_counter() - t0
        if n <= 50:
            t0 = time.perf_counter(); w3, _ = loese_lp(k, True); t3 = time.perf_counter() - t0
            t3_text, gleich = f"{t3*1000:>16.1f} ms", abs(w1 - w3) < 1e-6
        else:
            t3_text, gleich = f"{'uebersprungen':>19}", abs(w1 - w2) < 1e-6

        print(f"{n:>4} | {t1*1000:>9.1f} ms | {t2*1000:>18.1f} ms | {t3_text} | "
              f"{'ja' if gleich else 'NEIN':>8}")

    print("-" * 84)
    print("Fazit: Der spezialisierte Ungarische Algorithmus ist um Groessenordnungen")
    print("schneller. Nutzen Sie fuer reine Zuordnungen NIE einen MILP-Solver.")
    print("=" * 84)
