Skip to content

A network untangles spirals

A small network, trained as the scene starts, lifts one spiral over the other. Then a flat plane separates them.

examples/neural_untangle.py
"""How a neural network untangles two spirals.

No straight line separates two interleaved spirals, and no smooth deformation of the plane can pull
them apart: in 2D one arm would have to pass through the other. A small residual network lifts
the plane into 3D instead. Each of its four layers nudges every point, h ↦ h + V tanh(Uh + c), and
training pays for how far points move, so the network learns the cheapest way to untangle. Here
that means raising one arm above the other, until one flat plane cuts the two colors apart (after
Olah, "Neural Networks, Manifolds, and Topology", 2014). Carried back through the layers, that
flat cut becomes the winding boundary between the arms in the input plane.
"""

from dataclasses import dataclass

import numpy as np

import manimgx as m

LAYERS, HIDDEN = 4, 16
TURNS = 1.25  # of each spiral
ENERGY = 0.3  # the price of moving points: small, direct moves win
SEED = 0
SCALE = 2.4  # screen units per unit of the network's space
COLORS = ["#ff9f1c", "#2ec4b6"]


def two_spirals(count: int, rng: np.random.Generator) -> tuple[np.ndarray, np.ndarray]:
    """Two interleaved spirals of TURNS turns in the unit disk, and their labels (0, 1)."""
    turn = np.sqrt(rng.uniform(0.04, 1, count)) * TURNS * m.TAU
    arm = np.stack([np.cos(turn), np.sin(turn)], 1) * (turn / (TURNS * m.TAU))[:, None]
    arm += 0.02 * rng.normal(size=arm.shape)
    return np.concatenate([arm, -arm]), np.repeat([0.0, 1.0], count)


@dataclass
class Network:
    u: list[np.ndarray]  # (HIDDEN, 3) per layer
    c: list[np.ndarray]  # (HIDDEN,)
    v: list[np.ndarray]  # (3, HIDDEN)
    w: np.ndarray  # the classifier's normal
    b: np.ndarray  # and offset, shape (1,)

    def run(self, x: np.ndarray) -> tuple[list[np.ndarray], list[np.ndarray]]:
        """The points at every stage (the input in the plane z = 0, then after each layer) and
        each layer's hidden activations."""
        h = np.column_stack([x, np.zeros(len(x))])
        stages, hidden = [h], []
        for u, c, v in zip(self.u, self.c, self.v):
            a = np.tanh(h @ u.T + c)
            h = h + a @ v.T
            stages.append(h)
            hidden.append(a)
        return stages, hidden

    def stages(self, x: np.ndarray) -> list[np.ndarray]:
        return self.run(x)[0]

    def logit(self, x: np.ndarray) -> np.ndarray:
        return self.stages(x)[-1] @ self.w + self.b[0]


def train(
    x: np.ndarray, y: np.ndarray, rng: np.random.Generator, steps: int = 6000
) -> Network:
    """Logistic loss plus ENERGY × the mean squared move of each layer; full-batch Adam, with
    backpropagation written out."""
    net = Network(
        [0.5 * rng.normal(size=(HIDDEN, 3)) for _ in range(LAYERS)],
        [0.5 * rng.normal(size=HIDDEN) for _ in range(LAYERS)],
        [0.1 * rng.normal(size=(3, HIDDEN)) for _ in range(LAYERS)],
        rng.normal(size=3),
        np.zeros(1),
    )
    params = [*net.u, *net.c, *net.v, net.w, net.b]
    first = [np.zeros_like(p) for p in params]
    second = [np.zeros_like(p) for p in params]
    n = len(x)
    for step in range(1, steps + 1):
        stages, hidden = net.run(x)
        prob = 1 / (1 + np.exp(-(stages[-1] @ net.w + net.b[0])))
        g = (prob - y) / n  # d loss / d logit
        grad_u: list[np.ndarray] = [np.zeros(0)] * LAYERS
        grad_c: list[np.ndarray] = [np.zeros(0)] * LAYERS
        grad_v: list[np.ndarray] = [np.zeros(0)] * LAYERS
        grad_w, grad_b = stages[-1].T @ g, np.array([g.sum()])
        back = np.outer(g, net.w)  # d loss / d h, from the top down
        for k in reversed(range(LAYERS)):
            move = hidden[k] @ net.v[k].T
            through = back + 2 * ENERGY * move / n
            grad_v[k] = through.T @ hidden[k]
            pre = (through @ net.v[k]) * (1 - hidden[k] ** 2)
            grad_u[k], grad_c[k] = pre.T @ stages[k], pre.sum(0)
            back = (
                back + pre @ net.u[k]
            )  # h passes through each layer unchanged, plus its move
        for i, (p, grad) in enumerate(
            zip(params, [*grad_u, *grad_c, *grad_v, grad_w, grad_b])
        ):
            first[i] = 0.9 * first[i] + 0.1 * grad
            second[i] = 0.999 * second[i] + 0.001 * grad**2
            p -= (
                0.02
                * (first[i] / (1 - 0.9**step))
                / (np.sqrt(second[i] / (1 - 0.999**step)) + 1e-8)
            )
    return net


def boundary(net: Network, size: float = 1.15, n: int = 240) -> np.ndarray:
    """The network's decision boundary in the input plane (logit = 0), by marching squares:
    segments as (count, 2, 2)."""
    s = np.linspace(-size, size, n)
    gx, gy = np.meshgrid(s, s, indexing="ij")
    f = net.logit(np.column_stack([gx.ravel(), gy.ravel()])).reshape(n, n)
    segments = []
    corners = [(0, 0), (1, 0), (1, 1), (0, 1)]
    for i in range(n - 1):
        for j in range(n - 1):
            values = [f[i + di, j + dj] for di, dj in corners]
            crossings = []
            for k in range(4):
                (ai, aj), (bi, bj) = corners[k], corners[(k + 1) % 4]
                fa, fb = values[k], values[(k + 1) % 4]
                if (fa < 0) != (fb < 0):
                    t = fa / (fa - fb)
                    crossings.append(
                        [
                            s[i + ai] + t * (s[i + bi] - s[i + ai]),
                            s[j + aj] + t * (s[j + bj] - s[j + aj]),
                        ]
                    )
            if len(crossings) == 2:
                segments.append(crossings)
    return np.array(segments)


class NeuralUntangle(m.ThreeDScene):
    def construct(self) -> None:
        rng = np.random.default_rng(SEED)
        x, y = two_spirals(300, rng)
        net = train(x, y, rng)
        accuracy = float(np.mean((net.logit(x) > 0) == (y == 1)))
        # 0: the input plane; k: after layer k (in between: a blend of the two)
        stage = m.ValueTracker(0.0)

        def at(stages: list[np.ndarray]) -> np.ndarray:
            s = stage.get_value()
            k = min(int(s), LAYERS - 1)
            f = s - k
            return SCALE * ((1 - f) * stages[k] + f * stages[k + 1])

        # the grid of the input plane, and the data, carried through every layer
        # (the grid's lines are cut off at radius 1.2: far from the data the layers are wild)
        grid_inputs = []
        for t in np.linspace(-1.1, 1.1, 13):
            half = np.sqrt(1.2**2 - t**2)
            along = np.linspace(-half, half, 90)
            grid_inputs += [
                np.column_stack([np.full(90, t), along]),
                np.column_stack([along, np.full(90, t)]),
            ]
        grid_stages = [net.stages(g) for g in grid_inputs]
        grid = m.VGroup(
            *[
                m.VMobject(stroke_color="#5c677d", stroke_width=1.4, shade_in_3d=True)
                for _ in grid_inputs
            ]
        )

        def bend_grid(group: m.Mobject) -> None:
            for line, stages in zip(group.submobjects, grid_stages):
                assert isinstance(line, m.VMobject)
                line.set_points_as_corners(at(stages))

        bend_grid(grid)
        grid.add_updater(bend_grid)
        data_stages = net.stages(x)
        rows = np.ones((len(x), 4))
        rows[:, :3] = np.array([m.ManimColor(c).to_rgb() for c in COLORS])[
            y.astype(int)
        ]
        data = m.PMobject(stroke_width=7)
        data.add_points(at(data_stages), rgbas=rows)

        def carry(mob: m.Mobject) -> None:
            mob.points = at(data_stages)

        data.add_updater(carry)

        # the classifier: a plane in the last layer's space, w·h + b = 0
        normal = net.w / np.linalg.norm(net.w)
        foot = -net.b[0] * net.w / float(net.w @ net.w)
        u = np.cross(normal, [0.3, 0.5, 0.8])
        u /= np.linalg.norm(u)
        v = np.cross(normal, u)
        cut = m.Polygon(
            *[
                SCALE * (foot + 1.3 * (a * u + b * v))
                for a, b in [(-1, -1), (1, -1), (1, 1), (-1, 1)]
            ],
            stroke_width=1.5,
            stroke_color=m.WHITE,
            fill_color=m.WHITE,
            fill_opacity=0.18,
            shade_in_3d=True,
        )
        # the boundary it cuts: a curve on the sheet, drawn back to the input plane at the end
        pieces = boundary(net)
        piece_stages = net.stages(pieces.reshape(-1, 2))
        edge = m.VMobject(stroke_color=m.WHITE, stroke_width=3.5, shade_in_3d=True)

        def trace_edge(mob: m.Mobject) -> None:
            assert isinstance(mob, m.VMobject)
            ends = at(piece_stages).reshape(-1, 2, 3)
            mob.reset_points()
            for a, b in zip(ends[:, 0], ends[:, 1]):
                mob.start_new_path(a)
                mob.add_line_to(b)

        trace_edge(edge)
        edge.add_updater(trace_edge)

        # HUD
        title = m.Text("How a network untangles two spirals", font_size=38).to_corner(
            m.UL
        )
        subtitle = m.Text(
            "four residual layers, each nudging space: h ↦ h + V tanh(Uh + c)",
            font_size=22,
        ).set_color(m.GREY_B)
        subtitle.next_to(title, m.DOWN, aligned_edge=m.LEFT, buff=0.12)
        layer_value = m.Integer(0, font_size=30)
        layer_value.add_updater(lambda d: d.set_value(int(round(stage.get_value()))))
        layer_row = m.VGroup(m.Text("layer", font_size=28), layer_value).arrange(
            m.RIGHT, buff=0.15
        )
        layer_row.to_corner(m.UR)
        score = m.Text(
            f"one flat cut separates them: {100 * accuracy:.0f}% of points",
            font_size=24,
        ).to_edge(m.DOWN, buff=0.35)
        back = m.Text(
            "carried back to the input, the flat cut winds between the arms",
            font_size=24,
        ).to_edge(m.DOWN, buff=0.35)
        self.add_fixed_in_frame_mobjects(title, subtitle, layer_row, score, back)
        self.remove(score, back)

        self.set_camera_orientation(phi=35 * m.DEGREES, theta=-90 * m.DEGREES, zoom=1.0)
        self.add(grid, data)
        self.wait(1.5)
        self.begin_ambient_camera_rotation(rate=0.12)
        self.move_camera(phi=64 * m.DEGREES, run_time=2)
        for k in range(1, LAYERS + 1):
            self.play(stage.animate.set_value(k), run_time=3.2)
            self.wait(0.3)
        self.play(m.FadeIn(cut), m.FadeIn(score), m.Create(edge), run_time=1.5)
        self.wait(2.5)
        self.play(m.FadeOut(cut), m.FadeOut(score), run_time=0.6)
        self.play(
            stage.animate.set_value(0),
            m.FadeIn(back, rate_func=m.rush_from),
            run_time=4.4,
        )
        self.stop_ambient_camera_rotation()
        self.move_camera(phi=20 * m.DEGREES, theta=-90 * m.DEGREES, run_time=2)
        self.wait(1.5)


if __name__ == "__main__":
    NeuralUntangle().render("neural_untangle.mp4")