#!/usr/bin/env python3
"""
Simple demonstration script for the talk of M. Koniorczyk at 
Lectures on Modern Scientific Programming, Wigner Institute, Budapest, 2019.
Based on the example at 
https://docs.ocean.dwavesys.com/en/latest/examples/min_vertex.html
"""
import networkx as nx
import matplotlib.pyplot as plt
import pprint
import numpy as np
#Real Dwave
#from dwave.system.samplers import DWaveSampler
#Exact solution now
from dimod.reference.samplers import ExactSolver


pp = pprint.PrettyPrinter(indent=4)

G = nx.Graph()

problem_size=5

G.add_nodes_from(range(1, problem_size))

G.add_edges_from([(1, 2),
                  (1, 3),
                  (2, 4),
                  (3, 4),
                  (3, 5),
                  (4, 5)])
nx.draw(G, with_labels=True)
plt.show(block=False)

#https://docs.ocean.dwavesys.com/projects/system/en/stable/reference/generated/dwave.system.samplers.DWaveSampler.sample_qubo.html#dwave.system.samplers.DWaveSampler.sample_qubo

Q = {}

def add_to_element(i, j, c):
    """Add c to the element i,j of global Q"""
    try:
        Q[(i, j)] += c
    except KeyError:
        Q[(i, j)] = c

for vertex in G.edges():
    i = vertex[0]
    j = vertex[1]
    add_to_element(i, i, -1)
    add_to_element(j, j, -1)
    add_to_element(i, j, 1)
    add_to_element(j, i, 1)

Qnp=np.zeros((problem_size, problem_size))
for (i,j) in Q.keys():
    Qnp[i-1,j-1] = Q[(i,j)]

print("Q matrix:")
pp.pprint(Qnp)

#Real DWave:
#sampler = EmbeddingComposite(DWaveSampler(endpoint='https://URL_to_my_D-Wave_system/', token='ABC-123456789012345678901234567890', solver='My_D-Wave_Solver'))
# or just simply
#sampler = DWaveSampler()

sampler = ExactSolver()
sampleset = sampler.sample_qubo(Q)
print("-----------------------")
print("Samples:")

for sample, energy in sampleset.data(['sample', 'energy']):
    print(sample, energy)

min_energy = next(sampleset.data(['energy']))[0]
print("Minimum energy: %.1f\n"%min_energy)

for sample, energy in sampleset.data(['sample', 'energy']):
    if energy == min_energy:
        x = np.zeros((problem_size,1))
        for i in sample.keys():
            x[i-1] = sample[i]
        obj = np.matmul(np.transpose(x),np.matmul(Qnp,x))[0][0]
        print("----------")
        print("Solution:")
        pp.pprint(x)
        print("Max cut value:", obj)
