# 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 train(self, X, y, epochs=10000, verbose=False):
        losses = []
        X_inverted = 1 - X
        y_inverted = 1 - y
        
        for epoch in range(epochs):
            # ===== ПРЯМОЙ ПРОХОД (сохраняем все промежуточные значения) =====
            output_direct = self.forward(X)
            a1_direct = self.a1.copy()
            
            # ===== ИНВЕРТИРОВАННЫЙ ПРОХОД =====
            output_inverted = self.forward(X_inverted)
            a1_inverted = self.a1.copy()
            
            # ===== ВЫЧИСЛЕНИЕ ГРАДИЕНТОВ ДЛЯ ПРЯМОГО ПУТИ =====
            m = X.shape[0]
            
            # Ошибка выходного слоя (прямой путь)
            delta2_direct = (output_direct - y) * self.sigmoid_derivative(output_direct)
            dW2_direct = np.dot(a1_direct.T, delta2_direct) / m
            db2_direct = np.sum(delta2_direct, axis=0, keepdims=True) / m
            
            # Ошибка скрытого слоя (прямой путь)
            delta1_direct = np.dot(delta2_direct, self.W2.T) * self.sigmoid_derivative(a1_direct)
            dW1_direct = np.dot(X.T, delta1_direct) / m
            db1_direct = np.sum(delta1_direct, axis=0, keepdims=True) / m
            
            # ===== ВЫЧИСЛЕНИЕ ГРАДИЕНТОВ ДЛЯ ИНВЕРТИРОВАННОГО ПУТИ =====
            delta2_inverted = (output_inverted - y_inverted) * self.sigmoid_derivative(output_inverted)
            dW2_inverted = np.dot(a1_inverted.T, delta2_inverted) / m
            db2_inverted = np.sum(delta2_inverted, axis=0, keepdims=True) / m
            
            delta1_inverted = np.dot(delta2_inverted, self.W2.T) * self.sigmoid_derivative(a1_inverted)
            dW1_inverted = np.dot(X_inverted.T, delta1_inverted) / m
            db1_inverted = np.sum(delta1_inverted, axis=0, keepdims=True) / m
            
            # ===== КОМБИНИРОВАНИЕ И ПРИМЕНЕНИЕ ГРАДИЕНТОВ =====
            # Усредняем градиенты от прямого и инвертированного обучения
            self.W2 -= self.lr * (dW2_direct + dW2_inverted) / 2
            self.b2 -= self.lr * (db2_direct + db2_inverted) / 2
            self.W1 -= self.lr * (dW1_direct + dW1_inverted) / 2
            self.b1 -= self.lr * (db1_direct + db1_inverted) / 2
            
            # ===== ВЫЧИСЛЕНИЕ ОШИБОК ДЛЯ ЛОГИРОВАНИЯ =====
            loss_direct = np.mean((output_direct - y) ** 2)
            loss_inverted = np.mean((output_inverted - y_inverted) ** 2)
            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 % 2000 == 0:
                print(f"Epoch {epoch}, Loss: {total_loss:.6f} "
                      f"[D:{loss_direct:.4f} I:{loss_inverted:.4f} C:{consistency_loss:.4f}]")
        
        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=4, lr=0.5)
losses = model.train(X, y, epochs=10000, verbose=True)

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

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

correct_direct = 0
correct_inv = 0
correct_ensemble = 0

for i in range(len(X)):
    # Прямой проход
    output_direct = model.forward(X[i:i+1])
    f_x = output_direct[0][0]
    pred_direct = 1 if f_x > 0.5 else 0
    
    # Инвертированный проход
    X_inv = 1 - X[i:i+1]
    output_inverted = model.forward(X_inv)
    f_not_x = output_inverted[0][0]
    
    # Предсказание через инверсию: y = 1 - f(~x)
    pred_from_inv = 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_check = f_x + f_not_x
    
    # Подсчёт точности
    if pred_direct == int(y[i][0]):
        correct_direct += 1
    if pred_from_inv == int(y[i][0]):
        correct_inv += 1
    if pred_ensemble == int(y[i][0]):
        correct_ensemble += 1
    
    status = "✓" if abs(sum_check - 1.0) < 0.1 else "✗"
    
    print(f"{str(X[i]):<12} {int(y[i][0]):<8} {f_x:<15.6f} {pred_direct:<10} "
          f"{f_not_x:<15.6f} {1-f_not_x:<15.6f} {ensemble_score:<15.6f} {sum_check:<12.4f} {status}")

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

# ===== ДЕМОНСТРАЦИЯ СВОЙСТВА =====
print(f"\n{'='*70}")
print("ПРОВЕРКА СВОЙСТВА ИНВЕРСИИ")
print(f"{'='*70}")

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 pairs:
    v1 = model.forward(x1)[0][0]
    v2 = model.forward(x2)[0][0]
    print(f"{name} = {v1:.6f} + {v2:.6f} = {v1+v2:.6f} ≈ 1")

# ===== БОНУС: УСТОЙЧИВОСТЬ К ШУМУ =====
print(f"\n{'='*70}")
print("ТЕСТ НА УСТОЙЧИВОСТЬ К ШУМУ")
print(f"{'='*70}")

print("\nДобавляем шум ±0.2 к входам:")
for i in range(len(X)):
    noise = np.random.uniform(-0.2, 0.2, 2)
    noisy_x = np.clip(X[i] + noise, 0, 1).reshape(1, -1)
    
    f_noisy = model.forward(noisy_x)[0][0]
    pred_noisy = 1 if f_noisy > 0.5 else 0
    
    f_noisy_inv = model.forward(1 - noisy_x)[0][0]
    ensemble_noisy = (f_noisy + (1 - f_noisy_inv)) / 2
    pred_ens_noisy = 1 if ensemble_noisy > 0.5 else 0
    
    print(f"  Вход {X[i]} + шум = [{noisy_x[0][0]:.2f}, {noisy_x[0][1]:.2f}]")
    print(f"    Прямое: {f_noisy:.4f} → {pred_noisy}, "
          f"Ансамбль: {ensemble_noisy:.4f} → {pred_ens_noisy}")