# your code goes here
import numpy as np

class XNORNeuralNetwork:
    def __init__(self, lr=0.5):
        np.random.seed(42)
        # Два скрытых слоя для стабильного обучения XNOR
        self.W1 = np.random.randn(2, 3) * 0.8
        self.b1 = np.zeros((1, 3))
        self.W2 = np.random.randn(3, 3) * 0.8
        self.b2 = np.zeros((1, 3))
        self.W3 = np.random.randn(3, 1) * 0.8
        self.b3 = 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)
        
        self.z3 = np.dot(self.a2, self.W3) + self.b3
        self.a3 = self.sigmoid(self.z3)
        
        return self.a3
    
    def backward(self, X, y, output):
        m = X.shape[0]
        
        # Градиенты выходного слоя
        delta3 = (output - y) * self.sigmoid_derivative(output)
        dW3 = np.dot(self.a2.T, delta3) / m
        db3 = np.sum(delta3, axis=0, keepdims=True) / m
        
        # Градиенты второго скрытого слоя
        delta2 = np.dot(delta3, self.W3.T) * self.sigmoid_derivative(self.a2)
        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.W3 -= self.lr * dW3
        self.b3 -= self.lr * db3
        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=5000, verbose=False):
        X_inverted = 1 - X
        y_inverted = 1 - y
        
        # Объединяем прямые и инвертированные данные
        X_combined = np.vstack([X, X_inverted])
        y_combined = np.vstack([y, y_inverted])
        
        losses = []
        
        for epoch in range(epochs):
            # Перемешиваем данные
            indices = np.random.permutation(len(X_combined))
            X_shuffled = X_combined[indices]
            y_shuffled = y_combined[indices]
            
            # Прямой проход на всех данных
            output = self.forward(X_shuffled)
            
            # Ошибка
            loss = np.mean((output - y_shuffled) ** 2)
            losses.append(loss)
            
            # Обратный проход
            self.backward(X_shuffled, y_shuffled, output)
            
            if verbose and epoch % 1000 == 0:
                # Проверка на исходных данных
                pred = self.forward(X)
                pred_inv = self.forward(X_inverted)
                acc = np.mean((pred > 0.5) == y)
                acc_inv = np.mean((pred_inv > 0.5) == y_inverted)
                print(f"Epoch {epoch}, Loss: {loss:.6f}, "
                      f"Acc: {acc:.2f}, Acc_inv: {acc_inv:.2f}")
        
        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)

best_model = None
best_accuracy = 0

for attempt in range(5):
    print(f"\nПопытка {attempt + 1}:")
    model = XNORNeuralNetwork(lr=0.3)
    model.train(X, y, epochs=3000, verbose=False)
    
    # Проверка точности
    predictions = model.forward(X)
    accuracy = np.mean((predictions > 0.5) == y)
    print(f"  Точность: {accuracy*100:.0f}%")
    
    if accuracy > best_accuracy:
        best_accuracy = accuracy
        best_model = model
    
    if accuracy == 1.0:
        print("  ✓ Достигнута 100% точность!")
        break

model = best_model
print(f"\nЛучшая точность: {best_accuracy*100:.0f}%")

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

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

correct_direct = 0
correct_inv = 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 = 1 - X[i:i+1]
    f_not_x = model.forward(X_inv)[0][0]
    
    # Предсказание через инверсию
    pred_from_inv = 1 if (1 - f_not_x) > 0.5 else 0
    
    # Ансамбль
    ensemble = (f_x + (1 - f_not_x)) / 2
    pred_ensemble = 1 if ensemble > 0.5 else 0
    
    # Сумма для проверки свойства
    sum_check = f_x + f_not_x
    
    # Подсчёт точности
    target = int(y[i][0])
    if pred_direct == target: correct_direct += 1
    if pred_from_inv == target: correct_inv += 1
    if pred_ensemble == target: correct_ensemble += 1
    
    print(f"{str(X[i]):<12} {target:<8} {f_x:<12.4f} {pred_direct:<10} "
          f"{f_not_x:<12.4f} {1-f_not_x:<12.4f} {ensemble:<12.4f} {sum_check:<10.4f}")

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

# ===== ПРОВЕРКА СВОЙСТВА =====
print("\n" + "=" * 70)
print("ПРОВЕРКА СВОЙСТВА ИНВЕРСИИ")
print("=" * 70)

v1 = model.forward(np.array([[0, 0]]))[0][0]
v2 = model.forward(np.array([[1, 1]]))[0][0]
print(f"f(0,0) + f(1,1) = {v1:.4f} + {v2:.4f} = {v1+v2:.4f}")

v1 = model.forward(np.array([[0, 1]]))[0][0]
v2 = model.forward(np.array([[1, 0]]))[0][0]
print(f"f(0,1) + f(1,0) = {v1:.4f} + {v2:.4f} = {v1+v2:.4f}")

# ===== БОНУС: ВИЗУАЛИЗАЦИЯ УВЕРЕННОСТИ =====
print("\n" + "=" * 70)
print("АНАЛИЗ УВЕРЕННОСТИ ПРЕДСКАЗАНИЙ")
print("=" * 70)

for i in range(len(X)):
    f_x = model.forward(X[i:i+1])[0][0]
    X_inv = 1 - X[i:i+1]
    f_not_x = model.forward(X_inv)[0][0]
    
    direct_conf = abs(f_x - 0.5) * 2  # 0-1, где 1 = максимальная уверенность
    inv_conf = abs((1 - f_not_x) - 0.5) * 2
    ensemble_score = (f_x + (1 - f_not_x)) / 2
    ensemble_conf = abs(ensemble_score - 0.5) * 2
    
    # Проверка согласованности
    if abs(f_x + f_not_x - 1) < 0.1:
        consistency_bonus = "✓ согласованы"
    else:
        consistency_bonus = "✗ расхождение"
    
    print(f"Вход {X[i]}:")
    print(f"  Прямое: {f_x:.4f} (уверенность: {direct_conf:.2f})")
    print(f"  Инверсия: {1-f_not_x:.4f} (уверенность: {inv_conf:.2f})")
    print(f"  Ансамбль: {ensemble_score:.4f} (уверенность: {ensemble_conf:.2f}) - {consistency_bonus}")