import torch
import torch.nn as nn
import torch.optim as optim
import matplotlib.pyplot as plt
from IPython import display
display.set_matplotlib_formats('svg')
import numpy as np

print("For regression, you can use e.g. MAE, MSE")
print("For classification, you can use e.g. binary cross entropy (BCELoss, BCEWithLogitsLoss), categorical cross entropy (CrossEntropyLoss) ")

###
# MSELoss
# regression problems
###
print("------ MSELoss ------")

y_predict = torch.tensor(2.5)
y_real = torch.tensor(0.2)

loss_mse = nn.MSELoss()
loss_mse_out = loss_mse(y_predict, y_real)

print("y_predict", y_predict)
print("y_real", y_real)
print("loss_mse_out", loss_mse_out)

###
# BCELoss
# binary classification
# https://pytorch.org/docs/stable/generated/torch.nn.BCELoss.html
#
# Vzorec: L = - [ y * log(y_predict) + (1 - y) * log(1 - y_predict) ]
#
# Funguje jako přepínač podle skutečné třídy y (0 nebo 1):
# - Pokud je skutečnost y = 1:
#     druhá část (1 - y) se vynuluje, zbývá jen -log(y_predict).
#     Pokud síť předpoví 1, chyba je 0. Pokud předpoví 0, chyba je obrovská.
# - Pokud je skutečnost y = 0:
#     první část y se vynuluje, zbývá jen -log(1 - y_predict).
#     Pokud síť předpoví 0, chyba je 0. Pokud předpoví 1, chyba je obrovská.
###
print("------ BCELoss ------")

y_predict = torch.tensor(0.1) # torch.tensor(2.5)
y_real = torch.tensor(1.0)

loss_bce = nn.BCELoss()
loss_bce_out = loss_bce(y_predict, y_real)

print("y_predict", y_predict)
print("y_real", y_real)
print("loss_bce_out", loss_bce_out)

###
# the case that we have raw output, we can use sigmoid
# Sigmoid + BCELoss
#    def forward(self,x):
#        out = self.fc1...
#        out = self.relu...
#        out = F.sigmoid(self.fc2(out)) #sigmoid + BCELoss
#        return out
###

loss_bce = nn.BCELoss()
sigm = nn.Sigmoid()
y_predict = sigm(torch.tensor(2.5))
y_real = torch.tensor(1.0)
loss_sigm_bce_out = loss_bce(y_predict, y_real)

print("loss_sigm_bce_out", loss_sigm_bce_out)

###
# PyTorch library provides BCEWithLogitsLoss
# "This loss combines a Sigmoid layer and the BCELoss in one single class"
# "Logits are the outputs of a neural network before the activation function is applied"
#    def forward(self,x):
#        out = self.fc1...
#        out = self.relu...
#        out = self.fc2(out) # raw output as input for BCEWithLogitsLoss
#        return out
###

print("------ BCEWithLogitsLoss ------")

loss_bcewl = nn.BCEWithLogitsLoss()
y_predict = torch.tensor(2.5)
y_real = torch.tensor(1.0)
loss_bcewl_out = loss_bcewl(y_predict, y_real)

print("y_predict", y_predict)
print("y_real", y_real)
print("loss_bcewl_out", loss_bcewl_out)

###
# CrossEntropyLoss
# in the case that we have multiclass problem
# default loss function in Pytorch
# https://discuss.pytorch.org/t/can-we-use-cross-entropy-loss-for-binary-classification/159501
# https://pytorch.org/docs/stable/generated/torch.nn.CrossEntropyLoss.html
###

print("------ CrossEntropyLoss (LogSoftmax + NLLLoss) ------")

loss_cel = nn.CrossEntropyLoss()

# assume that the output of the network has three nodes
y_predict = torch.tensor([1.5, 3.0, 4.0])
print("y_predict", y_predict)
# the real class is zero
# but for class zero the model returned the smallest value (1.5) - so the error is large
y_real = torch.tensor(0)
loss_cel_out = loss_cel(y_predict, y_real)
print("loss_cel_out for class 0", loss_cel_out)

y_real = torch.tensor(1)
loss_cel_out = loss_cel(y_predict, y_real)
print("loss_cel_out for class 1", loss_cel_out)

y_real = torch.tensor(2)
loss_cel_out = loss_cel(y_predict, y_real)
print("loss_cel_out for class 2", loss_cel_out)

###
# NLLLoss
# https://pytorch.org/docs/stable/generated/torch.nn.NLLLoss.html#torch.nn.NLLLoss
# https://ljvmiranda921.github.io/notebook/2017/08/13/softmax-and-the-negative-log-likelihood/
###

print("------ NLLLoss ------")

loss_NLLLoss = nn.NLLLoss()
log_soft_max = nn.LogSoftmax(dim=0)
in_val = torch.FloatTensor([1.5, 3.0, 4.0])
print("in_val", in_val)
y_predict = log_soft_max(in_val)
y_real = torch.tensor(2)
loss_NLLLoss_out = loss_NLLLoss(y_predict, y_real)
print("loss_NLLLoss_out for class 2", loss_NLLLoss_out)
#hand = y_predict[0] * 0 + y_predict[1] * 0 + y_predict[2] * 1
#print(-hand)

# https://pytorch.org/docs/stable/generated/torch.nn.Softmax.html

# https://discuss.pytorch.org/t/what-classification-loss-should-i-choose-when-i-have-used-a-softmax-function/65121/10

print("------ Softmax, LogSoftmax ------")

l = [1, 2, 3]
num = np.exp(l)
den = np.sum(num)
soft_m = num/den 
print(num)
print(den)
print(soft_m, np.sum(soft_m))

soft_max = nn.Softmax(dim=0)
log_soft_max = nn.LogSoftmax(dim=0)

in_val = torch.FloatTensor([1, 4, 3])
output_sotfmax = soft_max(in_val)
output_logsotfmax = log_soft_max(in_val)
print("in_val", in_val)
print("output_sotfmax", output_sotfmax, np.log(output_sotfmax))
print("output_logsotfmax", output_logsotfmax)


