# your code goes here
import numpy as np

class XNORNeuralNetwork:
    def __init__(self, hidden_size=4, lr=0.8):
        np.random.seed(42)
        # Инициализация He для лучшей сходимости
        self.W1 = np.random.randn(2, hidden_size) * np.sqrt(2.0/2)
        self.b1 = np.zeros((1, hidden_size))
        self.W2 = np.random.randn(hidden_size, 1) * np.sqrt(2.0/hidden_size)
        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.dW2 = dW2
        self.db2 = db2
        self.dW1 = dW1
        self.db1 = db1
    
    def apply_gradients(self):
        self.W2 -= self.lr * self.dW2
        self.b2 -= self.lr * self.db2
        self.W1 -= self.lr * self.dW1
        self.b1 -= self.lr * self.db1
    
    def train(self, X, y, epochs=15000, verbose=False):
        losses = []
        X_inverted = 1 - X
        y_inverted = 1 - y
        
        for epoch in range(epochs):
            # Прямой проход для обычных данных
            output_direct = self.forward(X)
            loss_direct = np.mean((output_direct - y) ** 2)
            
            # Вычисляем градиенты для прямого прохода
            self.backward(X, y, output_direct)
            # Сохраняем градиенты
            dW1_direct = self.dW1.copy()
            db1_direct = self.db1.copy()
            dW2_direct = self.dW2.copy()
            db2_direct = self.db2.copy()
            
            # Прямой проход для инвертированных данных
            output_inverted = self.forward(X_inverted)
            loss_inverted = np.mean((output_inverted - y_inverted) ** 2)
            
            # Вычисляем градиенты для инвертированного прохода
            self.backward(X_inverted, y_inverted, output_inverted)
            
            # Объединяем градиенты с весами
            alpha = 0.5  # баланс между прямым и инвертированным обучением
            self.dW1 = alpha * dW1_direct + (1 - alpha) * self.dW1
            self.db1 = alpha * db1_direct + (1 - alpha) * self.db1
            self.dW2 = alpha * dW2_direct + (1 - alpha) * self.dW2
            self.db2 = alpha * db2_direct + (1 - alpha) * self.db2
            
            # Применяем комбинированные градиенты
            self.apply_gradients()
            
            # Дополнительный loss для свойства инверсии
            consistency_loss = np.mean((output_direct + output_inverted - 1) ** 2)
            
            total_loss = loss_direct + loss_inverted + 0.3 * consistency_loss
            losses.append(total_loss)
            
            if verbose and epoch % 1000 == 0:
                print(f"Epoch {epoch}, Total Loss: {total_loss:.6f}")
                print(f"  Direct: {loss_direct:.6f}, Inverted: {loss_inverted:.6f}, "
                      f"Consistency: {consistency_loss:.6f}")
                print(f"  Predictions: {output_direct.flatten()}")
        
        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)

print("=" * 70)
print("XNOR С КОМБИНИРОВАННЫМИ ГРАДИЕНТАМИ")
print("=" * 70)

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

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

print(f"\n{'Вход':<12} {'Цель':<8} {'f(x)':<15} {'Предск.':<10} "
      f"{'f(~x)':<15} {'Сумма':<15} {'Ансамбль':<15} {'Статус'}")
print("-" * 95)

correct_direct = 0
correct_inverted = 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]
    pred_inverted_from_not = 1 if (1 - f_not_x) > 0.5 else 0
    
    # Ансамблевое предсказание
    ensemble_score = (f_x + (1 - f_not_x)) / 2
    pred_ensemble = 1 if ensemble_score > 0.5 else 0
    
    # Проверка свойства
    sum_outputs = f_x + f_not_x
    
    # Подсчёт точности
    if pred_direct == int(y[i][0]):
        correct_direct += 1
    
    if pred_inverted_from_not == int(y[i][0]):
        correct_inverted += 1
    
    if pred_ensemble == int(y[i][0]):
        correct_ensemble += 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} {int(y[i][0]):<8} {f_x:<15.6f} {pred_direct:<10} "
          f"{f_not_x:<15.6f} {sum_outputs:<15.6f} {ensemble_score:<15.6f} {consistency}")

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

# ===== ПРОВЕРКА СВОЙСТВА ИНВЕРСИИ =====
print("\n" + "=" * 70)
print("ПРОВЕРКА СВОЙСТВА: f(x) + f(NOT x) ≈ 1")
print("=" * 70)

test_pairs = [
    ("f(0,0) + f(1,1)", np.array([[0, 0]]), np.array([[1, 1]])),
    ("f(0,1) + f(1,0)", np.array([[0, 1]]), np.array([[1, 0]])),
]

for name, x1, x2 in test_pairs:
    f1 = model.forward(x1)[0][0]
    f2 = model.forward(x2)[0][0]
    print(f"{name} = {f1:.6f} + {f2:.6f} = {f1 + f2:.6f} ≈ 1")

# ===== ДЕМОНСТРАЦИЯ УЛУЧШЕНИЯ =====
print("\n" + "=" * 70)
print("ПРЕИМУЩЕСТВО АНСАМБЛЯ")
print("=" * 70)

# Тест с добавлением шума
print("\nТест с шумом (добавляем случайный шум к входам):")
noise_level = 0.15
for i in range(len(X)):
    noisy_x = X[i] + np.random.randn(2) * noise_level
    noisy_x = np.clip(noisy_x, 0, 1).reshape(1, -1)
    
    f_x_noisy = model.forward(noisy_x)[0][0]
    pred_direct_noisy = 1 if f_x_noisy > 0.5 else 0
    
    # Ансамбль с инверсией для зашумлённых данных
    noisy_x_inv = 1 - noisy_x
    f_not_x_noisy = model.forward(noisy_x_inv)[0][0]
    ensemble_score_noisy = (f_x_noisy + (1 - f_not_x_noisy)) / 2
    pred_ensemble_noisy = 1 if ensemble_score_noisy > 0.5 else 0
    
    print(f"Вход: {X[i]} -> Зашумлённый: [{noisy_x[0][0]:.2f}, {noisy_x[0][1]:.2f}]")
    print(f"  Прямое: {f_x_noisy:.4f} -> {pred_direct_noisy}, "
          f"Ансамбль: {ensemble_score_noisy:.4f} -> {pred_ensemble_noisy}")