"""Diagrams for "CenterNet: Objects as Points"."""
import math

from kit import Drawing, FILL, BLUE, RUST, INK


def three_heads():
    """One forward pass, three outputs: box size, class heatmap and centre offset, all at 128 x 128."""
    d = Drawing(740, 330)
    d.raw('<path d="M10 110l46 -36v150l-46 36z" fill="%s" stroke="%s" stroke-width="1.2"/>' % (FILL["grey"], "#6B6B6B"))
    d.text(33, 284, "Input\nimage", size=12)
    d.cuboid(76, 96, 220, 150, d=60, fill="rust")
    d.raw('<path d="M186 96l60 -60M186 96v150" stroke="%s" stroke-width="1" stroke-dasharray="4 4" fill="none"/>' % RUST)
    d.text(186, 272, "Some Architecture", size=13)
    d.cuboid(376, 104, 16, 118, d=40, fill="rustmid")
    d.text(404, 246, "Downsampling", size=12)
    d.text(384, 92, "256 x 256", size=11, color="mute")
    d.arrow((356, 162), (376, 162))
    heads = (("w-h head", 2, 46), ("hm head", 80, 146), ("offset head", 2, 246))
    for name, channels, y in heads:
        w = 8 if channels == 2 else 30
        d.cuboid(520, y, w, 58, d=22, fill="rust")
        d.cuboid(520 + w, y, 5, 58, d=22, fill="rustsolid")
        d.text(590, y + 12, name, size=13, anchor="start", bold=True)
        d.text(590, y + 32, "%d x 128 x 128" % channels, size=11, anchor="start", color="mute")
        d.arrow((434, 150), (470, 150), (470, y + 29), (514, y + 29))
    return d


def ground_truth_heatmaps():
    """Each ground-truth box becomes one Gaussian peak on the heatmap of its class."""
    d = Drawing(740, 236)
    cells, size = 24, 7
    boxes = (((2, 6), (4, 3), "ink", "class 1"), ((15, 14), (3, 6), "blue", "class 2"), ((18, 21), (5, 2), "rust", "class 3"))
    d.box(8, 26, cells * size, cells * size, fill="grey", r=0)
    d.text(8 + cells * size / 2, 12, "Ground truths", bold=True)
    for (cx, cy), (w, h), color, _ in boxes:
        edge = {"ink": INK, "blue": BLUE, "rust": RUST}[color]
        d.raw('<rect x="%.1f" y="%.1f" width="%.1f" height="%.1f" fill="#fff" stroke="%s" stroke-width="2"/>'
              % (8 + (cx - w / 2 + .5) * size, 26 + (cy - h / 2 + .5) * size, w * size, h * size, edge))
        d.raw('<circle cx="%.1f" cy="%.1f" r="2.4" fill="%s"/>' % (8 + (cx + .5) * size, 26 + (cy + .5) * size, RUST))
    for i, ((cx, cy), (w, h), color, name) in enumerate(boxes):
        x = 8 + (i + 1) * (cells * size + 16)
        sigma = max(w, h) / 3.0                                # larger objects get a wider peak
        shade = {}
        for r in range(cells):
            for c in range(cells):
                v = math.exp(-((c - cx) ** 2 + (r - cy) ** 2) / (2 * sigma ** 2))
                if v > 0.02:
                    shade[(r, c)] = v
        d.text(x + cells * size / 2, 12, "Heatmap, %s" % name, bold=True)
        d.cells(x, 26, cells, cells, size, shade=shade)
    d.text(370, 212 + 12, "The peak is 1 at the box centre and falls off with a spread set by the box size.", size=12, color="mute")
    return d


def _heads(d, x, y, gap):
    names = ("Heatmap Head\n(W/R, H/R, C)", "Dimension Head\n(W/R, H/R, 2)", "Offset Head\n(W/R, H/R, 2)")
    return [d.ellipse(x + (i - 1) * gap, y, 54, 26, name, fill="blue", size=10.5) for i, name in enumerate(names)]


def flowcharts():
    """Training (left): three losses are summed and backpropagated. Inference (right): peaks become boxes."""
    d = Drawing(740, 560)
    # ---------------- training
    d.text(170, 14, "Training", bold=True, size=14)
    a = d.box(120, 30, 100, 36, "Input Image\n(W, H, 3)", fill="grey", r=8, size=11)
    b = d.box(112, 86, 116, 62, "Feature Extractor\nNetwork\n(ResNet, DLA,\nHourglass)", fill="rust", size=11)
    c = d.box(112, 168, 116, 34, "Downsample with a\ngiven stride R.", fill="rust", size=11)
    d.link(a, b)
    d.link(b, c)
    hs = _heads(d, 170, 248, 112)
    d.arrow(c.l, (hs[0].cx, c.cy), hs[0].t)
    d.arrow(c.b, hs[1].t)
    d.arrow(c.r, (hs[2].cx, c.cy), hs[2].t)
    losses = [d.box(h.cx - 50, 300, 100, 32, t, fill="white", size=11) for h, t in
              zip(hs, ("Compute\nHeatmap Loss", "Compute\nDimension Loss", "Compute\nOffset Loss"))]
    for h, l in zip(hs, losses):
        d.link(h, l)
    add = d.op(170, 374, "+", r=15)
    d.arrow(losses[0].b, (losses[0].cx, 374), add.l)
    d.arrow(losses[1].b, add.t)
    d.arrow(losses[2].b, (losses[2].cx, 374), add.r)
    d.raw('<path d="M170 414l58 26l-58 26l-58 -26z" fill="%s" stroke="%s" stroke-width="1.4"/>' % (FILL["blue"], BLUE))
    d.text(170, 440, "Epoch < Total\nEpochs", size=10.5)
    d.arrow(add.b, (170, 414))
    end = d.box(140, 498, 60, 30, "End", fill="grey", r=8, size=12)
    d.arrow((170, 466), end.t, label="No", dx=16, dy=0)
    d.arrow((228, 440), (346, 440), (346, 117), b.r, label="Yes", label_at=0.12)
    d.text(296, 100, "Backpropagate\nthe Loss", size=11, color="mute")
    # ---------------- inference
    d.text(560, 14, "Inference", bold=True, size=14)
    a = d.box(510, 30, 100, 36, "Input Image\n(W, H, 3)", fill="grey", r=8, size=11)
    b = d.box(502, 86, 116, 62, "Feature Extractor\nNetwork\n(ResNet, DLA,\nHourglass)", fill="rust", size=11)
    c = d.box(502, 168, 116, 34, "Downsample with a\ngiven stride R.", fill="rust", size=11)
    d.link(a, b)
    d.link(b, c)
    hs = _heads(d, 560, 248, 116)
    d.arrow(c.l, (hs[0].cx, c.cy), hs[0].t)
    d.arrow(c.b, hs[1].t)
    d.arrow(c.r, (hs[2].cx, c.cy), hs[2].t)
    x = hs[0].cx
    pool = d.box(x - 56, 292, 112, 30, "MaxPool 2D on\nheatmaps", fill="rust", size=10.5)
    peaks = d.box(x - 56, 338, 112, 30, "Consider top 100\npeaks as centers", fill="white", size=10.5)
    size = d.box(x - 56, 384, 112, 44, "Extract width and\nheight values for\nthe top k peaks", fill="white", size=10.5)
    gen = d.box(x - 56, 446, 112, 44, "Generate boxes\nwith these sizes\nand centers", fill="white", size=10.5)
    d.link(hs[0], pool)
    d.link(pool, peaks)
    d.link(peaks, size)
    d.link(size, gen)
    d.arrow(hs[1].b, (hs[1].cx, 406), size.r)
    off = d.box(hs[2].cx - 52, 446, 104, 44, "Extract offset\nvalues for the\ntop k peaks", fill="white", size=10.5)
    d.arrow(hs[2].b, off.t)
    d.line(peaks.r, (hs[2].cx, peaks.cy))
    add = d.op(hs[1].cx, 468, "+", r=13)
    d.arrow(gen.r, add.l)
    d.arrow(off.l, add.r)
    d.text(hs[1].cx, 440, "Add offsets to\ngenerated boxes", size=10, color="mute")
    final = d.box(add.cx - 46, 500, 92, 26, "Get Final Boxes", fill="white", size=10.5)
    d.arrow(add.b, final.t)
    rem = d.box(add.cx + 62, 494, 122, 38, "Remove boxes whose\npeak value is\nless than 0.3", fill="white", size=10)
    d.arrow(final.r, rem.l)
    out = d.box(add.cx - 176, 496, 112, 34, "Final Generated\nBoxes", fill="grey", r=8, size=10.5)
    d.arrow(rem.b, (rem.cx, 548), (out.cx, 548), out.b)
    return d
