# your code goes here
import numpy as np

class XNORNeuralNetwork:
    def __init__(self, hidden_size=4, lr=0.5):
        np.random.seed(42)
        self.W1 = np.random.randn(2, hidden_size) * 0.5
        self.b1 = np.zeros((1, hidden_size))
        self.W2 = np.random.randn(hidden_size, 1) * 0.5
        self.b2 = np.zeros((1, 1))
        self.lr = lr
    
    def sigmoid(self, x):
        x = np.clip(x, -500, 500)
        return 1 / (1 + np.exp(-x))
    
    def sigmoid_derivative(self, x):
        return x * (1 - x)
    
    def forward(self, X):
        self.z1 = np.dot(X, self.W1) + self.b1
        self.a1 = self.sigmoid(self.z1)
        self.z2 = np.dot(self.a1, self.W2) + self.b2
        self.a2 = self.sigmoid(self.z2)
        return self.a2
    
    def backward(self, X, y, output):
        m = X.shape[0]
        delta2 = (output - y) * self.sigmoid_derivative(output)
        dW2 = np.dot(self.a1.T, delta2) / m
        db2 = np.sum(delta2, axis=0, keepdims=True) / m
        
        delta1 = np.dot(delta2, self.W2.T) * self.sigmoid_derivative(self.a1)
        dW1 = np.dot(X.T, delta1) / m
        db1 = np.sum(delta1, axis=0, keepdims=True) / m
        
        self.W2 -= self.lr * dW2
        self.b2 -= self.lr * db2
        self.W1 -= self.lr * dW1
        self.b1 -= self.lr * db1
    
    def train(self, X, y, epochs=10000, verbose=False):
        losses = []
        for epoch in range(epochs):
            output = self.forward(X)
            
            # Стандартная ошибка
            standard_loss = np.mean((output - y) ** 2)
            
            # Дополнительная ошибка инверсии:
            # f(NOT x) + f(x) должно быть ≈ 1
            X_inverted = 1 - X  # Инвертируем входные значения
            output_inverted = self.forward(X_inverted)
            
            # y_inverted = 1 - y (инвертированная целевая переменная)
            y_inverted = 1 - y
            
            # Комбинированная функция потерь
            consistency_loss = np.mean((output + output_inverted - 1) ** 2)
            inverted_loss = np.mean((output_inverted - y_inverted) ** 2)
            
            # Общая ошибка
            total_loss = standard_loss + 0.3 * consistency_loss + 0.5 * inverted_loss
            
            losses.append(total_loss)
            
            # Обратное распространение с учётом всех компонент
            self.backward(X, y, output)
            
            if verbose and epoch % 2000 == 0:
                print(f"Epoch {epoch}, Loss: {total_loss:.6f}, "
                      f"Std: {standard_loss:.6f}, "
                      f"Inv: {inverted_loss:.6f}, "
                      f"Cons: {consistency_loss:.6f}")
        return losses


# ===== ДАННЫЕ =====
X = np.array([[0, 0], [0, 1], [1, 0], [1, 1]], dtype=np.float64)
y = np.array([[1], [0], [0], [1]], dtype=np.float64)

# Инвертированные данные
X_inv = 1 - X
y_inv = 1 - y

print("=" * 70)
print("XNOR С ИСПОЛЬЗОВАНИЕМ СВОЙСТВА ИНВЕРСИИ")
print("f(x) + f(NOT x) ≈ 1")
print("=" * 70)

# Обучение
model = XNORNeuralNetwork(hidden_size=4, lr=0.5)
losses = model.train(X, y, epochs=12000, verbose=True)

# ===== ТЕСТИРОВАНИЕ =====
print("\n" + "=" * 70)
print("РЕЗУЛЬТАТЫ")
print("=" * 70)

print(f"\n{'Вход':<12} {'f(x)':<15} {'Предск.':<10} "
      f"{'f(~x)':<15} {'Сумма':<15} {'Статус'}")
print("-" * 75)

correct_direct = 0
correct_ensemble = 0

for i in range(len(X)):
    # Прямой проход
    f_x = model.forward(X[i:i+1])[0][0]
    pred_direct = 1 if f_x > 0.5 else 0
    
    # Инвертированный проход
    X_inv_i = 1 - X[i]
    f_not_x = model.forward(X_inv_i.reshape(1, -1))[0][0]
    
    # Проверка свойства
    sum_outputs = f_x + f_not_x
    
    # Ансамблевое предсказание с учётом инверсии
    # f(x) = 1 - f(NOT x), поэтому:
    # ensemble_prediction = (f(x) + (1 - f(NOT x))) / 2
    ensemble_score = (f_x + (1 - f_not_x)) / 2
    pred_ensemble = 1 if ensemble_score > 0.5 else 0
    
    # Проверка корректности
    if pred_direct == y[i][0]:
        correct_direct += 1
    
    if pred_ensemble == y[i][0]:
        correct_ensemble += 1
    
    # Статус суммы (должна быть ≈ 1)
    if abs(sum_outputs - 1.0) < 0.1:
        consistency = "✓"
    elif abs(sum_outputs - 1.0) < 0.2:
        consistency = "~"
    else:
        consistency = "✗"
    
    print(f"{str(X[i]):<12} {f_x:<15.6f} {pred_direct:<10} "
          f"{f_not_x:<15.6f} {sum_outputs:<15.6f} {consistency}")

print(f"\nТочность прямых предсказаний: {correct_direct}/4")
print(f"Точность ансамблевых предсказаний: {correct_ensemble}/4")

# ===== ДЕМОНСТРАЦИЯ СВОЙСТВА =====
print("\n" + "=" * 70)
print("ДЕМОНСТРАЦИЯ СВОЙСТВА ИНВЕРСИИ")
print("=" * 70)

print("\nf(0,0) + f(1,1) = ", end="")
f00 = model.forward(np.array([[0, 0]]))[0][0]
f11 = model.forward(np.array([[1, 1]]))[0][0]
print(f"{f00:.4f} + {f11:.4f} = {f00 + f11:.4f} ≈ 1")

print("f(0,1) + f(1,0) = ", end="")
f01 = model.forward(np.array([[0, 1]]))[0][0]
f10 = model.forward(np.array([[1, 0]]))[0][0]
print(f"{f01:.4f} + {f10:.4f} = {f01 + f10:.4f} ≈ 1")