#!/usr/bin/env python3

# Vektorisierte_Modellgenerierung.py
"""
Kapitel Oekosystem: Warum der Solver oft gar nicht der Engpass ist.

In realen Projekten geht ein grosser Teil der Rechenzeit nicht ins Loesen,
sondern ins AUFBAUEN des Modells. Dieses Programm misst das an einem
Transportproblem wachsender Groesse in vier Stufen:

    A  Modellierungsschicht, Nebenbedingung fuer Nebenbedingung (OR-Tools)
    B  Matrix direkt, aber mit Python-Schleifen ueber die Eintraege (COO)
    C  Matrix vektorisiert ueber Kronecker-Produkte (NumPy/SciPy)
    D  Daten kommen als lange Tabelle, aufbereitet mit Polars

Alle Varianten loesen dasselbe Problem und muessen denselben Zielwert
liefern - das wird am Ende geprueft.

Benoetigt: numpy, scipy, ortools; Variante D zusaetzlich polars (optional).
"""

from __future__ import annotations

import importlib.util
import time

import numpy as np
import scipy.sparse as sp
from ortools.linear_solver import pywraplp
from scipy.optimize import linprog

HAT_POLARS = importlib.util.find_spec("polars") is not None


def erzeuge_daten(m: int, n: int, saat: int = 3
                  ) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
    """Transportproblem: m Werke, n Kunden.

    Liefert (kosten[m, n], angebot[m], bedarf[n]). Das Gesamtangebot liegt
    20 % ueber dem Gesamtbedarf, damit das Modell sicher loesbar ist.
    """
    rng = np.random.default_rng(saat)
    kosten = rng.uniform(1.0, 20.0, size=(m, n))
    bedarf = rng.uniform(10.0, 50.0, size=n)
    angebot = np.full(m, 1.2 * bedarf.sum() / m)
    return kosten, angebot, bedarf


# --- Variante A: Modellierungsschicht, Bedingung fuer Bedingung -------------

def loese_mit_modellierungsschicht(kosten: np.ndarray, angebot: np.ndarray,
                                   bedarf: np.ndarray) -> tuple[float, float, float]:
    """So schreibt man ein Transportproblem zuerst hin - gut lesbar, nah an
    der mathematischen Formulierung, jede Nebenbedingung ein eigener Aufruf.

    Jedes `s.Add(sum(...))` baut in Python einen Ausdrucksbaum aus m bzw. n
    Termen auf und uebergibt ihn einzeln an die C++-Schicht. Das ist der
    Preis der Bequemlichkeit - und er waechst linear mit der Modellgroesse.
    """
    m, n = kosten.shape

    t0 = time.perf_counter()
    s = pywraplp.Solver.CreateSolver("GLOP")
    x = [[s.NumVar(0, s.infinity(), f"x_{i}_{j}") for j in range(n)]
         for i in range(m)]
    for i in range(m):
        s.Add(sum(x[i][j] for j in range(n)) <= angebot[i])
    for j in range(n):
        s.Add(sum(x[i][j] for i in range(m)) >= bedarf[j])
    s.Minimize(sum(kosten[i][j] * x[i][j] for i in range(m) for j in range(n)))
    t_aufbau = time.perf_counter() - t0

    t0 = time.perf_counter()
    status = s.Solve()
    t_loesen = time.perf_counter() - t0
    if status != pywraplp.Solver.OPTIMAL:
        raise RuntimeError(f"Solver-Status: {status}")
    return t_aufbau, t_loesen, s.Objective().Value()


# --- Varianten B bis D: Matrix selbst bauen, dann SciPy/HiGHS ---------------

def baue_mit_schleifen(kosten: np.ndarray, angebot: np.ndarray,
                       bedarf: np.ndarray):
    """Die Nebenbedingungsmatrix als COO-Tripel (Zeile, Spalte, Wert), erzeugt
    in verschachtelten Python-Schleifen.

    Schon deutlich naeher am Blech als Variante A - es entsteht kein
    Ausdrucksbaum mehr. Die Schleife selbst bleibt aber Python.
    """
    m, n = kosten.shape
    zeilen: list[int] = []
    spalten: list[int] = []
    werte: list[float] = []
    rechte_seite: list[float] = []

    for i in range(m):                       # Angebot je Werk
        for j in range(n):
            zeilen.append(i)
            spalten.append(i * n + j)
            werte.append(1.0)
        rechte_seite.append(float(angebot[i]))

    for j in range(n):                       # Bedarf je Kunde, als -x <= -bedarf
        for i in range(m):
            zeilen.append(m + j)
            spalten.append(i * n + j)
            werte.append(-1.0)
        rechte_seite.append(-float(bedarf[j]))

    A_ub = sp.csr_matrix((werte, (zeilen, spalten)), shape=(m + n, m * n))
    return A_ub, np.array(rechte_seite), kosten.ravel()


def baue_vektorisiert(kosten: np.ndarray, angebot: np.ndarray,
                      bedarf: np.ndarray):
    """Dieselbe Matrix ohne eine einzige Schleife - ueber Kronecker-Produkte.

    Die Angebotsmatrix ist  kron(I_m, 1_n^T):  je Werk eine Zeile mit Einsen
    an genau den n Spalten dieses Werks. Die Bedarfsmatrix ist
    kron(1_m^T, I_n). Beide entstehen in je einem Aufruf und sind sofort
    duennbesetzt.
    """
    m, n = kosten.shape
    angebots_matrix = sp.kron(sp.identity(m, format="csr"), np.ones((1, n)))
    bedarfs_matrix = sp.kron(np.ones((1, m)), sp.identity(n, format="csr"))

    A_ub = sp.vstack([angebots_matrix, -bedarfs_matrix], format="csr")
    b_ub = np.concatenate([angebot, -bedarf])
    return A_ub, b_ub, kosten.ravel()


def baue_mit_polars(kosten: np.ndarray, angebot: np.ndarray, bedarf: np.ndarray):
    """Der realistische Fall: Die Kosten kommen als LANGE Tabelle
    (werk, kunde, kosten) aus Datenbank, Data Lake oder CSV-Datei.

    Polars berechnet den Spaltenindex jeder Variablen in einem einzigen
    Spaltenausdruck - ohne Python-Schleife ueber die Zeilen. Genau so baut man
    Modelle aus Millionen Tabellenzeilen.
    """
    import polars as pl

    m, n = kosten.shape
    tabelle = pl.DataFrame({
        "werk": np.repeat(np.arange(m), n),
        "kunde": np.tile(np.arange(n), m),
        "kosten": kosten.ravel(),
    }).with_columns(
        (pl.col("werk") * n + pl.col("kunde")).alias("var_index")
    )

    var_index = tabelle["var_index"].to_numpy()
    werk = tabelle["werk"].to_numpy()
    kunde = tabelle["kunde"].to_numpy()
    eins = np.ones(var_index.size)

    angebots_matrix = sp.csr_matrix((eins, (werk, var_index)), shape=(m, m * n))
    bedarfs_matrix = sp.csr_matrix((eins, (kunde, var_index)), shape=(n, m * n))

    A_ub = sp.vstack([angebots_matrix, -bedarfs_matrix], format="csr")
    b_ub = np.concatenate([angebot, -bedarf])
    return A_ub, b_ub, tabelle["kosten"].to_numpy()


def messe_matrixvariante(bauer, kosten, angebot, bedarf
                         ) -> tuple[float, float, float]:
    """Liefert (Aufbauzeit, Loesezeit, Zielwert) fuer die Varianten B bis D."""
    t0 = time.perf_counter()
    A_ub, b_ub, c = bauer(kosten, angebot, bedarf)
    t_aufbau = time.perf_counter() - t0

    t0 = time.perf_counter()
    ergebnis = linprog(c, A_ub=A_ub, b_ub=b_ub, bounds=(0, None), method="highs")
    t_loesen = time.perf_counter() - t0

    if not ergebnis.success:
        raise RuntimeError(f"Solver-Status: {ergebnis.message}")
    return t_aufbau, t_loesen, float(ergebnis.fun)


if __name__ == "__main__":
    print("=" * 88)
    print("  MODELLAUFBAU: WO DIE ZEIT WIRKLICH HINGEHT")
    print("=" * 88)
    print("Transportproblem mit m Werken und n Kunden -> m*n Variablen.\n")

    varianten: list[tuple[str, object]] = [
        ("A: OR-Tools, Add() je NB", None),        # Sonderfall, eigener Messpfad
        ("B: COO in Schleifen", baue_mit_schleifen),
        ("C: NumPy vektorisiert", baue_vektorisiert),
    ]
    if HAT_POLARS:
        varianten.append(("D: Polars-Tabelle", baue_mit_polars))
    else:
        print("Hinweis: polars nicht installiert - Variante D wird uebersprungen.\n")

    # Aufwaermlauf: Der erste Aufruf bezahlt Importe und einmalige
    # Initialisierungen. Wer den mitmisst, vergleicht Startkosten statt
    # Rechenarbeit - ein klassischer Benchmark-Fehler.
    aufwaerm = erzeuge_daten(10, 10)
    loese_mit_modellierungsschicht(*aufwaerm)
    for _, bauer in varianten[1:]:
        bauer(*aufwaerm)

    print(f"{'Groesse':<26} {'Variante':<26} {'Aufbau':>9} {'Loesen':>9} "
          f"{'Aufbauanteil':>13}")
    print("-" * 88)

    for m, n in [(40, 40), (120, 120), (250, 250)]:
        kosten, angebot, bedarf = erzeuge_daten(m, n)
        zielwerte = []
        for nummer, (name, bauer) in enumerate(varianten):
            if bauer is None:
                t_aufbau, t_loesen, ziel = loese_mit_modellierungsschicht(
                    kosten, angebot, bedarf)
            else:
                t_aufbau, t_loesen, ziel = messe_matrixvariante(
                    bauer, kosten, angebot, bedarf)
            zielwerte.append(ziel)
            anteil = 100.0 * t_aufbau / (t_aufbau + t_loesen)
            groesse = f"{m}x{n} = {m*n:,} Variablen" if nummer == 0 else ""
            print(f"{groesse:<26} {name:<26} {t_aufbau:>8.3f}s {t_loesen:>8.3f}s "
                  f"{anteil:>12.0f} %")

        # Alle Varianten muessen dasselbe Problem beschreiben.
        spanne = max(zielwerte) - min(zielwerte)
        assert spanne < 1e-6 * max(abs(z) for z in zielwerte), \
            f"Varianten widersprechen sich: {zielwerte}"
        print(f"{'':<26} {'-> Zielwert (alle gleich)':<26} {zielwerte[0]:>9.2f}"
              f"   Spanne {spanne:.1e}")
        print("-" * 88)

    print("\nZwei Lehren aus der Tabelle:")
    print("1. Bei Variante A geht mehr Zeit in den AUFBAU als ins Loesen. Wer hier")
    print("   einen schnelleren Solver kauft, beschleunigt den kleineren Teil.")
    print("2. Zwischen B und C liegt keine andere Mathematik, nur eine andere")
    print("   Schreibweise derselben Matrix - Schleife gegen Kronecker-Produkt.")
    print("\nDie Lesbarkeit von Variante A ist trotzdem viel wert: Fangen Sie dort an,")
    print("und vektorisieren Sie erst, wenn die Messung es verlangt.")
    print("=" * 88)
