import marimo

__generated_with = "0.21.1"
app = marimo.App(width="medium")


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    # CNN for Square vs Circle (from scratch)

    Train a tiny CNN to classify 28×28 squares vs circles using
    only NumPy — no frameworks. The same model as the article, with vectorized convolution and pooling for faster training.
    """)
    return


@app.cell
def _():
    import marimo as mo
    import numpy as np
    import matplotlib.pyplot as plt

    np.random.seed(42)
    return mo, np, plt


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ## Layer classes
    """)
    return


@app.cell
def _(np):
    class Conv2D:
        def __init__(self, num_filters, kernel_size, padding='same'):
            if padding not in ('same', 'valid') or kernel_size % 2 == 0:
                raise ValueError("Use an odd kernel size and 'same' or 'valid' padding")
            self.kernels = np.random.randn(num_filters, kernel_size, kernel_size) * 0.1
            self.biases = np.zeros(num_filters)
            self.padding = (kernel_size - 1) // 2 if padding == 'same' else 0

        def forward(self, input):
            if self.padding > 0:
                p = self.padding
                input = np.pad(input, ((p, p), (p, p)), mode='constant')
            self.input = input
            k = self.kernels.shape[1]
            self.patches = np.lib.stride_tricks.sliding_window_view(input, (k, k))
            self.z = np.einsum('rcij,fij->frc', self.patches, self.kernels)
            self.z += self.biases[:, None, None]
            self.out = np.maximum(0, self.z)
            return self.out

        def backward(self, upstream_gradient):
            grad = upstream_gradient * (self.z > 0)
            self.grad_kernels = np.einsum('frc,rcij->fij', grad, self.patches)
            self.grad_biases = grad.sum(axis=(1, 2))

        def update(self, lr):
            self.kernels -= lr * self.grad_kernels
            self.biases -= lr * self.grad_biases

    class MaxPool2D:
        def __init__(self, size=2):
            self.size = size

        def forward(self, x):
            self.input = x
            s = self.size
            channels, h, w = x.shape
            oh, ow = h // s, w // s
            windows = x[:, :oh*s, :ow*s].reshape(channels, oh, s, ow, s)
            windows = windows.transpose(0, 1, 3, 2, 4).reshape(channels, oh, ow, s*s)
            self.max_indices = windows.argmax(axis=-1)
            out = windows.max(axis=-1)
            self.output_shape = out.shape
            return out

        def backward(self, gradient):
            s = self.size
            out = np.zeros_like(self.input)
            ch, row, col = np.indices(self.output_shape)
            out[ch, row*s + self.max_indices//s, col*s + self.max_indices%s] = gradient
            return out

    class DenseLayer:
        def __init__(self, n_inputs, n_outputs):
            self.W = np.random.randn(n_outputs, n_inputs) * np.sqrt(2.0 / n_inputs)
            self.b = np.zeros(n_outputs)

        def forward(self, x):
            self.x = x
            return self.W @ x + self.b

        def backward(self, grad_output):
            self.grad_W = np.outer(grad_output, self.x)
            self.grad_b = grad_output
            return self.W.T @ grad_output

        def update(self, lr):
            self.W -= lr * self.grad_W
            self.b -= lr * self.grad_b

    def sigmoid(x):
        z = np.exp(-np.abs(x))
        return np.where(np.asarray(x) >= 0, 1 / (1 + z), z / (1 + z))

    def binary_cross_entropy(predicted, label):
        p = np.clip(predicted, 1e-10, 1 - 1e-10)
        return -label * np.log(p) - (1 - label) * np.log(1 - p)

    return Conv2D, MaxPool2D, DenseLayer, sigmoid, binary_cross_entropy


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ## Generate dataset
    """)
    return


@app.cell
def _(np):
    SIZE = 28
    LINE_WIDTH = 2.0
    rr, cc = np.mgrid[0:SIZE, 0:SIZE] + 0.5

    def make_shape(kind):
        radius = np.random.uniform(4.0, 10.0)
        margin = radius + 2.0
        cy, cx = np.random.uniform(margin, SIZE - margin, size=2)
        dy, dx = np.abs(rr - cy), np.abs(cc - cx)
        distance = np.maximum(dy, dx) if kind == 'square' else np.hypot(dy, dx)
        edge_distance = np.abs(distance - radius)
        coverage = np.clip((LINE_WIDTH / 2 + 0.5 - edge_distance) / 0.5, 0, 1)
        return (255 * coverage).astype(np.float32)

    N = 2000
    squares = np.array([make_shape('square') for _ in range(N)])
    circles = np.array([make_shape('circle') for _ in range(N)])

    X = np.concatenate([squares, circles], axis=0)
    y = np.concatenate([np.zeros(N), np.ones(N)])
    X = X / 255.0

    # Split each class separately, then shuffle each subset.
    splits = [[], [], []]
    for label in (0, 1):
        class_indices = np.random.permutation(np.flatnonzero(y == label))
        for target, part in zip(splits, np.split(class_indices, [1440, 1600])):
            target.extend(part)
    train_idx, val_idx, test_idx = [np.random.permutation(part) for part in splits]
    X_train, y_train = X[train_idx], y[train_idx]
    X_val, y_val = X[val_idx], y[val_idx]
    X_test, y_test = X[test_idx], y[test_idx]
    print(f"Train: {len(X_train)}, validation: {len(X_val)}, test: {len(X_test)}")
    return SIZE, X_test, X_train, X_val, y_test, y_train, y_val


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ## Visualize samples
    """)
    return


@app.cell
def _(X_train, plt, y_train):
    def _show_samples():
        fig, axes = plt.subplots(2, 8, figsize=(12, 3))
        fig.suptitle("Sample squares (top) and circles (bottom)", fontsize=13)
        sq_idx = [i for i, l in enumerate(y_train) if l == 0][:8]
        ci_idx = [i for i, l in enumerate(y_train) if l == 1][:8]
        for i, j in enumerate(sq_idx):
            axes[0, i].imshow(X_train[j], cmap="gray")
            axes[0, i].axis("off")
        for i, j in enumerate(ci_idx):
            axes[1, i].imshow(X_train[j], cmap="gray")
            axes[1, i].axis("off")
        plt.tight_layout()
        plt.show()
    _show_samples()
    return


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ## Build and train
    """)
    return


@app.cell
def _(Conv2D, DenseLayer, MaxPool2D, X_test, X_train, X_val,
      binary_cross_entropy, np, sigmoid, y_test, y_train, y_val):
    # Build model
    conv = Conv2D(num_filters=2, kernel_size=3, padding='same')
    pool = MaxPool2D(size=2)
    dense = DenseLayer(n_inputs=2*14*14, n_outputs=1)

    def forward(image):
        activated = conv.forward(image)
        pooled = pool.forward(activated)
        flat = pooled.reshape(-1)
        output = sigmoid(dense.forward(flat))
        return output[0]

    def backward(output, label):
        grad = np.array([output - label])
        grad = dense.backward(grad)
        grad = grad.reshape(pool.output_shape)
        grad = pool.backward(grad)
        conv.backward(grad)

    # Train
    learning_rate = 0.25
    num_epochs = 20
    batch_size = 32
    train_history = {"train_loss": [], "train_acc": [], "val_loss": [], "val_acc": []}

    for epoch in range(num_epochs):
        indices = np.random.permutation(len(X_train))
        total_loss = 0
        correct = 0

        for batch_start in range(0, len(X_train), batch_size):
            batch_idx = indices[batch_start:batch_start+batch_size]
            bs = len(batch_idx)

            acc_conv_k = np.zeros_like(conv.kernels)
            acc_conv_b = np.zeros_like(conv.biases)
            acc_dense_W = np.zeros_like(dense.W)
            acc_dense_b = np.zeros_like(dense.b)

            for ix in batch_idx:
                pred = forward(X_train[ix])
                sample_loss = binary_cross_entropy(pred, y_train[ix])
                total_loss += sample_loss
                correct += (1 if (pred > 0.5) == y_train[ix] else 0)

                backward(pred, y_train[ix])
                acc_conv_k += conv.grad_kernels
                acc_conv_b += conv.grad_biases
                acc_dense_W += dense.grad_W
                acc_dense_b += dense.grad_b

            conv.grad_kernels = acc_conv_k / bs
            conv.grad_biases = acc_conv_b / bs
            dense.grad_W = acc_dense_W / bs
            dense.grad_b = acc_dense_b / bs
            conv.update(learning_rate)
            dense.update(learning_rate)

        train_acc = correct / len(X_train)
        train_loss = total_loss / len(X_train)

        val_loss = 0
        val_correct = 0
        for vi in range(len(X_val)):
            v_pred = forward(X_val[vi])
            val_loss += binary_cross_entropy(v_pred, y_val[vi])
            val_correct += (1 if (v_pred > 0.5) == y_val[vi] else 0)
        val_acc = val_correct / len(X_val)
        val_loss = val_loss / len(X_val)

        train_history["train_loss"].append(float(train_loss))
        train_history["train_acc"].append(train_acc)
        train_history["val_loss"].append(float(val_loss))
        train_history["val_acc"].append(val_acc)

        print(f"Epoch {epoch+1}/{num_epochs} — loss: {train_loss:.4f}, "
              f"train: {train_acc:.2%}, val: {val_acc:.2%}")

    test_correct = sum(1 for ti in range(len(X_test))
                       if (forward(X_test[ti]) > 0.5) == y_test[ti])
    final_test_acc = test_correct / len(X_test)
    print(f"\nTest accuracy: {final_test_acc:.4f}")
    test_loss = float(np.mean([
        binary_cross_entropy(forward(image), label)
        for image, label in zip(X_test, y_test)
    ]))
    training_data = {
        "config": {"seed": 42, "epochs": num_epochs, "batch_size": batch_size,
                   "learning_rate": learning_rate, "renderer": "shared-edge-coverage-v1"},
        "split_sizes": {"train": len(X_train), "validation": len(X_val), "test": len(X_test)},
        "epochs": {"baseline": {"epoch": list(range(1, num_epochs + 1)), **train_history}},
        "test": {"accuracy": final_test_acc, "loss": test_loss},
        "examples": {"squareA": X_test[y_test == 0][0].tolist(),
                     "circleA": X_test[y_test == 1][0].tolist()},
        "kernels": conv.kernels.tolist(), "biases": conv.biases.tolist(),
    }
    return final_test_acc, train_history, training_data


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ## Results
    """)
    return


@app.cell
def _(final_test_acc, train_history, plt):
    def _show_results():
        print(f"Test accuracy: {final_test_acc:.4f}")

        fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))
        ep = range(1, len(train_history["train_acc"]) + 1)

        ax1.plot(ep, train_history["train_acc"], label="Train", linestyle="--", color="#4a6cf7")
        ax1.plot(ep, train_history["val_acc"], label="Validation", color="#e67e22")
        ax1.set_title("Accuracy")
        ax1.set_xlabel("Epoch")
        ax1.set_ylabel("Accuracy")
        ax1.legend()
        ax1.grid(True, alpha=0.3)

        ax2.plot(ep, train_history["train_loss"], label="Train", linestyle="--", color="#4a6cf7")
        ax2.plot(ep, train_history["val_loss"], label="Validation", color="#e67e22")
        ax2.set_title("Loss")
        ax2.set_xlabel("Epoch")
        ax2.set_ylabel("Loss")
        ax2.legend()
        ax2.grid(True, alpha=0.3)

        plt.tight_layout()
        plt.show()
    _show_results()
    return


@app.cell
def _(mo, training_data):
    import json
    mo.download(json.dumps(training_data, indent=2).encode(),
                filename="training.json", label="Download metrics and learned filters")
    return


if __name__ == "__main__":
    app.run()
