# versucht, LIN(0/1) fuer quadratische Matrix mit einem greedy-Verfahren zu loesen.
# nimmt als Basis eine reelle Loesung des GLS, wo x(i) auch < 0.0 oder > 1.0
# sein koennen. Setzt dann alle Variablen erst auf 0 und dann auf 1, alle anderen Variablen
# werden bzgl. des Betrages der Korrelation durchlaufen und so gesetzt, dass in
# allen Gleichungen >= rechte Seite versucht wird zu erreichen/halten, ohne sich zu verschlechtern
# Von einer numerisch reellen Loesung wird polynomiell versucht, eine diskrete Loesung in {0,1} zu finden.
# J. Gamenik August 2019
import math
import random
import lin01transform

anzZeilen = 13
anzSpalten = 13
abbruchSofort = True
nurEindeutig = True  # True: es soll nur eine Loesung geben
sucheTransformiert = False  # True: danach Suche mit einem transformierten System
random = random.Random()


# gibt Tupel von Matrix, rechter Seite und reeler Loesung zurueck
def holeBeispiel():
  global anzZeilen, anzSpalten
  matrix = [[-1, -1, -1, 7], 
            [-1, 4, -1, -1], 
            [-1, -1, 5, -1], 
            [6, -1, -1, -1]
           ]
  rechts = [1, 1, 1, 1]
  anzZeilen = 4
  anzSpalten = 4
  return (matrix, rechts, loeseReell(matrix,rechts))


# holt die Loesung im rellen
def loeseReell(matrix, rechts):
  return [0.034, 0.054, 0.039, 0.045]  #  [0.2, 0.2, 0.2, 0.2]   # TODO mit Cramer loesen


# gibt Tupel von Matrix, rechter Seite und reeler Loesung zurueck
def holeBeispiel2():
  global anzZeilen, anzSpalten
  matrix = [[-8, -1, 0, 6], 
            [-9, 5, -8, -5], 
            [3, -5, -9, -10], 
            [4, 0, -9, -7]
           ]
  rechts = [-8, -26, -34, -22]
  anzZeilen = 4
  anzSpalten = 4
  return (matrix, rechts, loeseReell2(matrix,rechts))


# holt die Loesung im rellen
def loeseReell2(matrix, rechts):
  return [1.22, 1.62, 2.28, 0.75]  #  [0.2, 0.2, 0.2, 0.2]   # TODO mit Cramer loesen


# prueft, ob Loesung im 0/1 existiert mit vollstaendiger Suche. Gibt 0,1 oder 2 zurueck
def existiert01(matrix, rechts):
  bis = 2 ** anzSpalten  # exponentiell !
  gefunden = 0
  for a in range(0,bis):
    loesung = []
    lauf = a
    for b in range(0,anzSpalten):
      loesung.append(lauf % 2)
      lauf = lauf // 2
    # alle Bedingungen mit dieser Loesung durchrechnen
    allesErfuellt = True
    for b in range(0,anzZeilen):
      summe = 0
      for c in range(0,anzSpalten):
        summe += matrix[b][c] * loesung[c]
      if summe < rechts[b]:
        allesErfuellt = False
        break
    if allesErfuellt:
      # print("Existierende Loesung {0}".format(loesung))
      gefunden += 1
      if gefunden == 2: return gefunden
  return gefunden


# definiert Zufallswerte
def holeZufallsbeispiel():
  loesung = []
  fak = 6
  # reeler Loesungsvektor im Intervall [-fak/2;+fak/2]
  for a in range(0,anzSpalten):
    loesung.append( random.random() * fak - fak/2)
  # Matrix mit ganzahligen Koeffizienten zwischen [-10;+10]
  matrix = []
  for z in range(0,anzZeilen):
    zeile = []
    for s in range(0,anzSpalten):
      zeile.append(random.randint(-10,+10))
    matrix.append(zeile)
  # rechte Seite, so dass eine Loesung existiert im reellen
  rechts = []
  for a in range(0,anzZeilen):
    summe = 0
    for b in range(0,anzSpalten):
      summe += matrix[a][b] * loesung[b]
    rechts.append(int(summe) - 2)   # damit die init. Loesung aufgeht, hier reell ?
  return(matrix,rechts,loesung)


# definiert Zufallswerte in Matrix mit einem dominanten Element pro Spalte
def holeZufallsbeispiel2():
  loesung = []
  for a in range(0,anzSpalten):
    loesung.append(random.random() * 0.2 - 0.1)
  print("Loesung reell: {0}".format(loesung))
  # Diagonalmatrix mit ganzahligen Koeffizienten zwischen [1;+10]; Rest -1
  matrix = []
  for z in range(0,anzZeilen):
    zeile = []
    for s in range(0,anzSpalten):
      zeile.append(random.randint(anzSpalten,+10) if (z % 2 == 0 and s == z+1) or (z % 2 == 1 and s == z-1) else -1 )
    matrix.append(zeile)
  print("Matrix: {0}".format(matrix))
  rechts = []
  for a in range(0,anzZeilen):
    rechts.append(0.1)   # damit 0-Loesung Variablen nicht geht
  print("rechts: {0}".format(rechts))
  return(matrix,rechts,loesung)


# holt Loesung in 0/1. Falls matrixO und rechtsO uebergeben, ist matrix und rechts verrechnet, dann ist
# aber nicht jede Loesung davon eine von matrixO und rechtsO. Es wird nur dann beenendet, wenn in beiden eine Loesung
def loese01(matrix, rechts, teilloesung, vorverarbeitet, matrixO, rechtsO):
  diff = holeAbweichung(matrix, rechts, teilloesung)
  print("Initiale Abweichung ", diff)

  erledigt = vorverarbeitung(matrix, teilloesung, vorverarbeitet)
  if len(erledigt) > 0:
    diff = holeAbweichung(matrix, rechts, teilloesung)
    print("Danach Abweichung ", diff)
  # Sonderfall keine freien Variablen
  if len(erledigt) == anzSpalten:
    if (pruefeErweitert(diff, teilloesung,matrixO,rechtsO)):
      return teilloesung 
  # TODO Spezialfall: die ganze rechte Seite >= 0, dann alle Variablen 0
  # fuer jede der n Variablen 0 und 1 testen, dann gemaess Korrelation weiter
  for a in range(0,anzSpalten):
    if a in erledigt:
      continue
    # Korrelationen betragsmaessig absteigend
    korr = holeKorrelationen(a,matrix,erledigt)
    #print("Korrelationsliste zu x{0}: ".format(a+1), zeigeKorr(korr))
    for test in range(0,2):
      teilloesung2 = teilloesung.copy()
      teilloesung2[a] = test
      #print("Setze x{0} auf {1}".format(a+1, test), sep=" ")
      diff2 = holeAbweichung(matrix, rechts, teilloesung2)
      # die anderen Variablen, ggf. mit einer Alternative
      alternativen = []
      b = 0
      while b < len(korr):
        # setze Loesung in der Zeile mit der groessten Differenz
        wo = setzeLoesungsteil(diff2,matrix,korr,teilloesung2,b,alternativen)
        diff2 = holeAbweichung(matrix, rechts, teilloesung2)
        #print("Nachher x{0}, {1} => {2}".format(wo+1, teilloesung2[wo], diff2))
        b += 1
        if b == len(korr):
          if (pruefeErweitert(diff2, teilloesung2,matrixO,rechtsO)):
            return teilloesung2
          while len(alternativen) > 0:
            # es gibt noch Alternativen, zu letzten zurueck
            #print("Anzahl Alternativen", len(alternativen))
            b = alternativen[len(alternativen)-1]
            alternativen.remove(b)
            #print("Hole Alternative x{0}".format(korr[b][1]+1))
            teilloesung2[korr[b][1]] = 1
            diff2 = holeAbweichung(matrix, rechts, teilloesung2)
            b += 1
            if b == len(korr):
              if (pruefeErweitert(diff2, teilloesung2,matrixO,rechtsO)):
                return teilloesung2
            if b < len(korr):
              break
      pass
      if len(alternativen) > 0:
        raise RuntimeError("noch bestehende Alternativen")
      if len(korr) == 0:
        if (pruefeErweitert(diff2, teilloesung2,matrixO,rechtsO)):
          return teilloesung2
  pass
  return None


# prueft, ob die Differenz passst und ggf in weiterem UGLS
def pruefeErweitert(diff, teilloesung,matrixO,rechtsO):
  if testeFertig(diff) == True:
    print("Gefundene Loesung", teilloesung)
    if abbruchSofort: 
      if matrixO != None and rechtsO != None:
        diffO = holeAbweichung(matrixO, rechtsO, teilloesung)
        print(diffO)
        if testeFertig(diffO) == True:
          print("Diese loest auch das Ursprungssystem")
          return True
      else:
        return True 
  return False

# gibt Tupel zu Korrelationswert auf Zielvariable aus
def zeigeKorr(korr):
  rueck = ""
  for a in korr:
    rueck += "{0:.3f} -> x{1}, ".format(a[0],a[1]+1)
  return rueck


# setzt einen Loesungsteil
def setzeLoesungsteil(diff2,matrix,korr,teilloesung2,b,alternativen):
  # Zeile mit der groessten Differenz finden. Es koennen aber Korrelationen
  # mit dem gleichen Wert vorkommen, das hier beachten 
  wertEx = 1e8
  indexEx = -1
  for c in range(0,anzZeilen):
    if (diff2[c] <= wertEx):
      wertEx = diff2[c]
      indexEx = c
  #print("Minimum {0} in Zeile {1}".format(wertEx, c))
  # naechste Korrelationen untersuchen, bei gleichen ueber Koeffizient auswerten
  wo,koeff = holeKorrelationsvariable(b, indexEx, korr, matrix)       
  # print("Vorher x{0}, {1} => {2}".format(wo+1, teilloesung2[wo], diff2))
  if teilloesung2[wo] <= 0 and koeff >= 0:
    teilloesung2[wo] = 0
    #print("Dazu Alternative bei x", korr[b][1] +1, sep='')
    alternativen.append(b)
  elif koeff >= 0:
    teilloesung2[wo] = 1    # eindeutige Verbesserung      
  elif teilloesung2[wo] >= 1 and koeff <= 0:
    teilloesung2[wo] = 0
    #print("Dazu Alternative bei x", korr[b][1] + 1, sep='')
    alternativen.append(b)
  elif koeff <= 0:
    teilloesung2[wo] = 0   # eindeutige Verbesserung
  else:
    raise RuntimeError("unbekannter Fall")
  return wo


# hole die naechste Variable anhand der Korrelation
def holeKorrelationsvariable(b, indexEx, korr, matrix):
  aktuelleKorr = korr[b][0]
  wo = korr[b][1]
  neu = b
  koeff = matrix[indexEx][wo]
  for k in range(b+1,len(korr)):
    if korr[k][0] < aktuelleKorr - 0.001:
      break    
    wo2 = korr[k][1]
    koeff2 = matrix[indexEx][wo2]
    if abs(koeff2) > abs(koeff):
      # print("{0} mit {1} => {2} mit {3}".format(b, koeff, k, koeff2))
      koeff = koeff2  # Vorsicht !
      neu = k
  if b != neu:
    # die Elemente in der Korrelationsliste tauschen
    #print ("Tausche gleiche Korrelationen x{0} und x{1}".format(korr[b][1]+1, korr[neu][1]+1))
    t2 = korr[neu]
    korr.remove(t2)
    t1 = korr[b]
    korr.remove(t1)
    korr.insert(b, t1)
    korr.insert(neu, t2) 
  return (wo, koeff)


# gib Korrelationen zur Variablen a zurueck
def holeKorrelationen(a,matrix,erledigt):
  korr = []
  for b in range(0,anzSpalten):
    if b in erledigt or a == b:
      continue
    kneu = vxy(matrix,a,b)
    tupelneu = (abs(kneu), b)
    korr.append(tupelneu)
  korr.sort()
  korr.reverse()
  return korr


# prueft triviale Bedingungen, also alle Koeffizienten einer Var. mit gl. Vorzeichen
def vorverarbeitung(matrix, teilloesung, vorverarbeitet):
  erledigt = {}  # fest gesetzte Variablen
  for a in range(0,anzSpalten):
    if vorverarbeitet != None and a in vorverarbeitet:
      print("Variable {0} nicht optimiert".format(a+1))
      continue
    ch = holeVariablentyp(matrix, a)
    if (ch[0] == 0):
      print("Variable {0} kann negativ gesetzt werden".format(a+1))
      erledigt[a] = True
      teilloesung[a] = 0
    if (ch[1] == 0):
      print("Variable {0} kann positiv gesetzt werden".format(a+1))
      erledigt[a] = True
      teilloesung[a] = 1
  return erledigt


# testet, ob alles positiv ist
def testeFertig(diff):
  for a in range(0,len(diff)):
    if diff[a] < 0.0:
      return False
  return True


# bestimmt Differenzvektor. Positiver Eintrag heisst, die Gleichung ist erfuellt
def holeAbweichung(matrix, rechts, teilloesung):
  abweichung = []
  for a in range(0, len(matrix)):
    zeile = matrix[a]
    summe = 0
    for b in range(0, len(teilloesung)):
      summe += zeile[b] * teilloesung[b]
    abweichung.append(summe - rechts[a])
  return abweichung


# bestimmt die Anzahl positiver und negativer Koeffizienten in einer Variable
def holeVariablentyp(matrix, nr):
  positiv = 0
  negativ = 0
  for a in range(0,len(matrix)):
    if matrix[a][nr] > 0:
      positiv += 1
    if matrix[a][nr] < 0:
      negativ += 1
  return (positiv, negativ)


# bestimmt Mittelwert der Koeff. einer Spalte
def holeMittelwert(matrix, nr):
  summe = 0.0
  for a in range(0,len(matrix)):
    summe += matrix[a][nr]
  return summe / len(matrix)


# bestimme Selbstkorrelation in einer Spalte
def sx(matrix, nr):
  summe = 0.0
  mittel = holeMittelwert(matrix, nr)
  for a in range(0,len(matrix)):
    summe += ((matrix[a][nr] - mittel) ** 2) 
  return math.sqrt(summe / (len(matrix) - 1))  # Div durch 0


# bestimme Korrelation in einer Spalte mit einer anderen
def sxy(matrix, nr, nr2):
  summe = 0.0
  mittel = holeMittelwert(matrix, nr)
  mittel2 = holeMittelwert(matrix, nr2)
  for a in range(0,len(matrix)):
    summe += ((matrix[a][nr] - mittel) * (matrix[a][nr2] - mittel2)) 
  return summe / (len(matrix) - 1)  # Div durch 0


# bestimme Korrelationskoeffizient
def vxy(matrix,nr,nr2):
  oben = sxy(matrix,nr,nr2)
  unten1 = sx(matrix,nr)
  unten2 = sx(matrix,nr2)
  if oben == 0.0:
    return oben # TODO
  if unten1 == 0.0 or unten2 == 0.0:
    raise RuntimeError("nicht bestimmbar")
  return oben / (unten1 * unten2)


# Hauptprogramm
lin01transform.anzZeilen = anzZeilen
lin01transform.anzSpalten = anzSpalten
meldung = None
lauf = 0
while lauf < 10:
  if nurEindeutig and meldung == None:
    print ("Suche nach eindeutigem System, kann dauern ...")
    meldung = "x"
  #matrix,rechts, loesungReell = holeBeispiel2()
  matrix,rechts, loesungReell = holeZufallsbeispiel()  # 2 mit spez. Matrix
  # nur weiter, wenn es genau eine Loesung gibt; bei festem Beispiel kann hier endlosschleife sein !
  if nurEindeutig and existiert01(matrix,rechts) != 1:
    continue
  if nurEindeutig:
    warten = input("... Beispiel gefunden, nun Suche ...")
    meldung = None    
  print("Loesung reell: {0}".format(loesungReell))
  print("Matrix: {0}".format(matrix))
  print("rechts: {0}".format(rechts))
  # mit dem gierigen Algorithmus loesen
  loesung01 = loese01(matrix, rechts, loesungReell, None, None, None)
  if nurEindeutig and loesung01 == None:
     raise Exception("Fehler: kein Loesung gefunden, es gibt aber eine")
  #print("Loesungsvektor ", loesung01)
  #
  if sucheTransformiert:
    # schaue nach, was Transformation liefert
    matrix2,rechts2,vorverarbeitet = lin01transform.transformiere(matrix,rechts)
    print("Ergebnis transformiert ", end="")
    print(matrix2,end="")
    print(rechts2,end="")
    print(vorverarbeitet)
    # loest die gefundene Loesung das transformierte System ?
    diff = holeAbweichung(matrix2, rechts2, loesung01)
    if not testeFertig(diff):
      raise Exception("Originalloesung loest nicht das transformierte System")
    print("  Starte transformierte Suche")
    # TODO loesungReell passt nicht mehr gut, moeglichst neu bestimmen, dann kaeme aber Abh. zu numpy rein
    loesung012 = loese01(matrix2,rechts2,loesungReell, vorverarbeitet,matrix,rechts)
    if nurEindeutig and loesung012 == None:
      print ("Hinweis: Kein Loesungsvektor beim transformierten System mit reellen Startloesung")  # raise Exception
    #
  warten = input("weiter ...")
  print("")
  lauf += 1

