"""Diagrams for "DetectoRS - A Comprehensive Review"."""
from kit import Drawing


def _n(d, x, y, label, fill, w=36, h=24):
    return d.box(x - w / 2, y - h / 2, w, h, label, fill=fill, size=12, bold=True)


def _legend(d, x, y, items):
    for i, (fill, label) in enumerate(items):
        d.box(x, y + i * 22, 22, 12, fill=fill, r=1)
        d.text(x + 30, y + i * 22 + 6, label, size=11, anchor="start")


# ------------------------------------------------------------------ cascades
def cascades():
    """Faster R-CNN has one detection head; Cascade R-CNN chains three, each fed the boxes of the one before."""
    d = Drawing(740, 250)
    top, mid, pool, base = 34, 92, 150, 208
    # Faster R-CNN
    d.text(120, 12, "Faster R-CNN", bold=True)
    img = _n(d, 30, base, "I", "white", 34, 34)
    conv = _n(d, 96, base, "conv", "solid", 46)
    h0 = _n(d, 96, pool, "H0", "solid")
    c0, b0 = _n(d, 76, mid, "C0", "grey"), _n(d, 118, mid, "B0", "blue")
    p1 = _n(d, 180, pool, "pool", "rustsolid", 44)
    h1 = _n(d, 180, mid, "H1", "solid")
    c1, b1 = _n(d, 160, top, "C1", "grey"), _n(d, 202, top, "B1", "blue")
    d.link(img, conv)
    d.link(conv, h0)
    d.arrow(h0.t, c0.b)
    d.arrow(h0.t, b0.b)
    d.arrow(conv.r, (180, base), p1.b)
    d.arrow(b0.r, (150, mid), (150, pool), p1.l)
    d.link(p1, h1)
    d.arrow(h1.t, c1.b)
    d.arrow(h1.t, b1.b)
    d.line((250, 8), (250, 242), dash=True, color="grey")
    # Cascade R-CNN
    d.text(500, 12, "Cascade R-CNN", bold=True)
    img = _n(d, 290, base, "I", "white", 34, 34)
    conv = _n(d, 356, base, "conv", "solid", 46)
    d.link(img, conv)
    b_prev = _n(d, 356, top, "B0", "blue")
    for i in range(3):
        x = 450 + i * 110
        p = _n(d, x, pool, "pool", "rustsolid", 44)
        h = _n(d, x, mid, "H%d" % (i + 1), "solid")
        c, b = _n(d, x - 21, top, "C%d" % (i + 1), "grey"), _n(d, x + 21, top, "B%d" % (i + 1), "blue")
        d.arrow((x, base), p.b)
        d.arrow(b_prev.bottom(0.8), (p.x - 16, pool - 14), (p.x, pool))
        d.link(p, h)
        d.arrow(h.t, c.b)
        d.arrow(h.t, b.b)
        b_prev = b
    d.line(conv.r, (670, base))
    return d


def mask_cascades():
    """Cascade Mask R-CNN adds a mask head per stage; Hybrid Task Cascade also passes masks stage to stage and adds a semantic branch."""
    d = Drawing(740, 262)
    top, pool, base = 60, 128, 196
    d.text(170, 12, "Cascade Mask R-CNN", bold=True)
    f = _n(d, 28, base, "F", "white", 34, 34)
    rpn = _n(d, 78, top, "RPN", "rust", 44)
    d.arrow((78, base), rpn.b)
    prev = rpn
    for i in range(3):
        x = 150 + i * 80
        p = _n(d, x, pool, "pool", "solid", 40)
        m, b = _n(d, x - 20, top, "M%d" % (i + 1), "grey", 34), _n(d, x + 20, top, "B%d" % (i + 1), "rust", 34)
        d.arrow((x, base), p.b)
        d.arrow(prev.bottom(0.8), (p.x - 10, pool - 16), (p.x, pool))
        d.arrow(p.t, m.b)
        d.arrow(p.t, b.b)
        prev = b
    d.line(f.r, (310, base))
    d.line((362, 8), (362, 254), dash=True, color="grey")
    # HTC
    d.text(550, 12, "Hybrid Task Cascade", bold=True)
    mrow, brow = 62, 108
    pool, base = 160, 214
    f = _n(d, 396, base, "F", "white", 34, 34)
    s = _n(d, 396, 34, "S", "rustmid", 26)
    d.arrow(f.t, s.b)
    rpn = _n(d, 446, brow, "RPN", "rust", 44)
    d.arrow((446, base), rpn.b)
    xs = [520 + i * 64 for i in range(4)]
    pools = [_n(d, x, pool, "pool", "solid", 38) for x in xs]
    for x, p in zip(xs, pools):
        d.arrow((x, base), p.b)
    d.line(f.r, (xs[-1], base))
    boxes = [_n(d, x, brow, "B%d" % (i + 1), "rust", 32) for i, x in enumerate(xs[:3])]
    masks = [_n(d, x, mrow, "M%d" % (i + 1), "grey", 32) for i, x in enumerate(xs[:3])]
    d.arrow(rpn.bottom(0.8), (pools[0].x - 8, pool - 14), (pools[0].x, pool))
    for i in range(3):
        d.arrow(pools[i].t, boxes[i].b)
        d.arrow(boxes[i].bottom(0.85), (pools[i + 1].x - 8, pool - 14), (pools[i + 1].x, pool))
        d.arrow(pools[i + 1].top(0.3), (masks[i].x2 + 10, mrow + 22), masks[i].bottom(0.85))
        d.arrow((masks[i].cx, 34), masks[i].t, color="rust")
    d.line(s.r, (masks[2].cx, 34), color="rust")
    d.arrow(masks[0].r, masks[1].l, color="blue", sw=1.8)
    d.arrow(masks[1].r, masks[2].l, color="blue", sw=1.8)
    d.text(550, 246, "blue: mask information flow    rust: semantic features", size=11, color="mute")
    return d


# ------------------------------------------------------------------ pyramids
WIDTHS = [150, 138, 96, 72, 54, 38]          # x0 .. x5


def _bottom_up(d, cx, y0, k=1.0, tag=""):
    """The backbone: feature maps x0..x5 with a downsampling stage B1..B5 between each pair. Returns the x shapes."""
    xs, y = [], y0
    for i in range(6):
        w = WIDTHS[i] * k
        fm = d.cuboid(cx - w / 2, y, w, 4, d=7, fill="grey" if i == 0 else "rust")
        d.text(cx - w / 2 - 6, y, "x%d%s" % (i, tag), size=10.5, anchor="end", italic=True)
        xs.append(fm)
        if i == 5:
            break
        bw = (WIDTHS[i] * 0.55 + WIDTHS[i + 1] * 0.45) * k
        b = d.cuboid(cx - bw / 2, y - 26, bw, 11, d=7, fill="blue")
        d.text(cx + bw / 2 + 12, y - 22, "B%d" % (i + 1), size=10.5, anchor="start", bold=True)
        d.arrow((cx, y - 7), (cx, y - 15), sw=1)
        d.arrow((cx, y - 33), (cx, y - 39), sw=1)
        y -= 46
    return xs


def _top_down(d, cx, xs, k=1.0):
    """The FPN side: F5..F2 merge the level above with the backbone map beside it and give f5..f2."""
    fs, prev = {}, None
    for i in range(5, 1, -1):
        y = xs[i].y + 7
        w = WIDTHS[i] * k
        F = d.cuboid(cx - w / 2, y - 6, w, 11, d=7, fill="mid")
        d.text(cx + w / 2 + 12, y - 2, "F%d" % i, size=10.5, anchor="start", bold=True)
        f = d.cuboid(cx - w / 2, y + 19, w, 4, d=7, fill="dark")
        d.arrow((cx, y + 5), (cx, y + 12), sw=1)
        if prev:
            d.arrow((cx, prev.y2), (cx, y - 13), sw=1)
        if i < 5:
            d.arrow((xs[i].x2, xs[i].cy + 3), (cx - w / 2, xs[i].cy + 3), sw=1)
        fs[i], prev = f, f
    top = xs[5]
    d.arrow((top.cx, top.y), (top.cx, top.y - 22), (cx, top.y - 22), (cx, top.y - 7), sw=1)
    return fs


LEGEND = (("rust", "Bottom Up Feature Maps"), ("blue", "Bottom Up Downsample Layers"),
          ("mid", "Top Down Downsample Layers"), ("dark", "Top Down Feature Maps"))


def fpn():
    """The original Feature Pyramid Network: a bottom-up backbone, a top-down path, and one prediction per level."""
    d = Drawing(740, 330)
    d.text(330, 14, "Original FPN", bold=True, size=14)
    d.arrow((22, 290), (22, 70), color="grey", sw=3)
    d.text(64, 180, "Bottom Up\nLayers", size=11)
    xs = _bottom_up(d, 210, 290)
    fs = _top_down(d, 420, xs)
    stage = d.box(560, 50, 26, 200, fill="grey", r=0)
    d.text(573, 36, "Prediction\nStage", size=11)
    for i, f in fs.items():
        d.arrow((f.x2, f.cy + 3), (stage.x, f.cy + 3), label="f%d" % i, sw=1, label_at=0.3, dy=-8)
    d.arrow((716, 70), (716, 250), color="grey", sw=3)
    d.text(668, 160, "Top Down\nFeature\nMaps", size=11)
    _legend(d, 420, 250, LEGEND)
    return d


def rfp():
    """Recursive Feature Pyramid, unrolled twice: the first pass's outputs go through ASPP back into the same backbone, and both passes are fused."""
    d = Drawing(740, 400)
    d.text(370, 12, "Recursive Feature Pyramid Network", bold=True, size=14)
    k = 0.62
    d.text(110, 78, "t = 1", italic=True)
    xs1 = _bottom_up(d, 78, 350, k)
    fs1 = _top_down(d, 196, xs1, k)
    d.group(332, 56, 172, 330, fill="tint", dash=True)
    d.text(418, 372, "Same Backbone", size=11)
    d.text(470, 70, "t = 2", italic=True)
    xs2 = _bottom_up(d, 418, 350, k, "")
    fs2 = _top_down(d, 560, xs2, k)
    for n, i in enumerate(range(5, 1, -1)):
        f1, f2 = fs1[i], fs2[i]
        y = f1.cy + 3
        aspp = d.box(266, y - 9, 44, 18, "ASPP", fill="white", size=10)
        d.arrow((f1.x2, y), aspp.l, sw=1, label="f%d" % i, dy=-7, size=10)
        d.arrow(aspp.r, (xs2[i].cx - WIDTHS[i] * k * 0.4, y), sw=1)          # into the backbone stage that produces x_i
        fuse = d.box(662 - n * 6, f2.cy - 6, 50, 18, "Fusion", fill="rust", size=10)
        d.arrow((f2.x2, f2.cy + 3), fuse.l, sw=1)
        lane = 30 + n * 7
        d.arrow((f1.x2 + 16 + n * 5, y), (f1.x2 + 16 + n * 5, lane), (fuse.cx + 8, lane), (fuse.cx + 8, fuse.y), sw=1, color="grey")
        d.arrow(fuse.r, (fuse.x2 + 14, fuse.cy), sw=1)
    return d


# ------------------------------------------------------------------ modules
def aspp():
    """ASPP: four parallel branches look at the feature at different scales and are concatenated back to C channels."""
    d = Drawing(740, 380)
    d.text(370, 14, "Atrous Spatial Pyramid Pooling (ASPP) Module", bold=True, size=14)
    src = d.cuboid(20, 168, 50, 50, d=18, fill="rust")
    d.text(54, 136, "f(i)", italic=True)
    d.text(54, 234, "(C, H, W)", size=11)
    cat = d.op(560, 186, "c", r=15, fill="blue")
    d.text(560, 150, "Concatenation\nChannelwise", size=11)
    out = d.cuboid(640, 162, 50, 50, d=18, fill="mid")
    d.text(674, 130, "R(f(i))", italic=True)
    d.text(674, 228, "(C, H, W)", size=11)
    d.arrow((cat.x2, 186), (out.x, 186))
    rows = ((56, 1, 1, 0, "grey", 22, 300), (132, 3, 3, 3, "rustmid", 40, 290), (226, 3, 3, 3, "rustmid", 40, 290))
    for n, (y, kern, dil, pad, fill, s, x) in enumerate(rows):
        b = d.cuboid(x, y, s, s, d=14, fill=fill)
        d.text(x - 12, y + s / 2 - 4, "Kernel Size = %d\nDilation Value = %d\nPadding = %d\nOut Channels = C/4" % (kern, dil, pad), size=10, anchor="end")
        d.arrow((88, 184), (x - 150, b.cy), head=False)
        d.arrow((x - 150, b.cy), (x - 142, b.cy), sw=1)
        d.arrow((b.x2 + 2, b.cy), (cat.x - 2 if n == 1 else cat.cx - 9, cat.cy if n == 1 else (cat.y + 3 if n == 0 else cat.y2 - 3)))
    gap = d.cuboid(196, 318, 8, 40, d=14, fill="dark")
    d.text(120, 336, "Global Average\nPooling", size=11)
    last = d.cuboid(300, 322, 22, 22, d=14, fill="grey")
    d.text(350, 340, "Kernel Size = 1, Dilation Value = 1\nPadding = 0, Out Channels = C/4\nExpand Spatially", size=10, anchor="start")
    d.arrow((70, 222), (gap.x - 2, gap.cy + 6))
    d.arrow((gap.x2, gap.cy + 6), (last.x - 2, last.cy + 6))
    d.arrow((last.x2, last.cy - 2), (cat.cx + 4, cat.y2 + 2))
    return d


def rfp_into_resnet():
    """How the fed-back feature enters ResNet: one extra 1x1 convolution added to the first block of each stage."""
    d = Drawing(740, 220)
    d.group(40, 14, 520, 126, fill="grey", r=2)
    d.group(40, 146, 520, 54, fill="blue", r=2)
    d.text(66, 104, "Input", anchor="start")
    a = d.box(150, 34, 84, 34, "Conv\n(1x1)", size=12)
    b = d.box(270, 34, 100, 34, "Conv\n(3x3, s=2)", size=12)
    c = d.box(406, 34, 84, 34, "Conv\n(1x1)", size=12)
    s = d.box(270, 88, 100, 34, "Conv\n(1x1, s=2)", size=12)
    add = d.op(524, 105, "+", r=15)
    d.arrow((110, 105), s.l)
    d.arrow((126, 105), (126, 51), a.l)
    d.link(a, b)
    d.link(b, c)
    d.arrow(c.r, (524, 51), add.t)
    d.arrow(s.r, add.l)
    d.text(54, 173, "RFP Features", anchor="start")
    r = d.box(270, 156, 100, 34, "Conv\n(1x1)", size=12)
    d.arrow((160, 173), r.l)
    d.arrow(r.r, (524, 173), add.b)
    d.arrow(add.r, (590, 105))
    d.text(596, 105, "Output", anchor="start")
    _legend(d, 600, 150, (("grey", "ResNet"), ("blue", "RFP")))
    return d


def fusion():
    """The fusion module: a sigmoid gate decides, per position, how much of the new feature to mix with the old one."""
    d = Drawing(740, 190)
    d.text(160, 40, "f(i) at t+1", italic=True, bold=True)
    d.text(380, 40, "f(i) at t", italic=True, bold=True)
    conv = d.box(110, 108, 100, 40, "Conv\n(1x1)")
    sig = d.box(270, 108, 100, 40, "Sigmoid")
    m1 = d.op(470, 60, "x", r=16)
    m2 = d.op(470, 128, "x", r=16)
    add = d.op(560, 94, "+", r=16)
    d.arrow((160, 52), conv.t)
    d.link(conv, sig)
    d.arrow(sig.r, m2.l, label="σ")
    d.arrow((420, 44), (470, 44), head=False)
    d.arrow((424, 40), (454, 56))
    d.arrow(sig.top(0.8), (m1.x, m1.cy + 8), label="1 − σ", dx=-6)
    d.arrow((208, 46), (330, 78), (m2.x + 2, m2.cy - 10), curve=True)
    d.arrow(m1.r, (add.x + 2, add.cy - 8))
    d.arrow(m2.r, (add.x + 2, add.cy + 8))
    d.arrow(add.r, (620, 94))
    d.text(626, 94, "Output", anchor="start", size=14)
    return d


def _context(d, x, title, fill):
    if fill != "none":
        d.group(x, 30, 118, 200, fill=fill, r=2)
    d.text(x + 59, 244, title, size=11)
    pool = d.box(x + 12, 96, 62, 30, "Global\nAvgPool", size=10.5)
    conv = d.box(x + 12, 46, 62, 30, "Conv\n(1x1)", size=10.5)
    add = d.op(x + 94, 150, "+", r=14)
    d.arrow((x + 43, 150), pool.b)
    d.link(pool, conv)
    d.arrow(conv.r, (x + 94, 61), add.t)
    return add


def sac():
    """Switchable Atrous Convolution: the same 3x3 weights run at atrous rate 1 and 3, and a learned switch S blends the two."""
    d = Drawing(740, 262)
    d.text(24, 150, "Input", size=13)
    pre = _context(d, 52, "Pre-Global Context", "rust")
    d.arrow((48, 150), pre.l)
    d.group(176, 30, 372, 200, fill="tint", r=2)
    d.text(362, 244, "Switchable Atrous Convolution", size=11)
    c1 = d.box(296, 46, 124, 30, "Conv\n(3x3, atrous=1)", size=10.5)
    c3 = d.box(296, 190, 124, 30, "Conv\n(3x3, atrous=3)", size=10.5)
    ap = d.box(222, 135, 70, 30, "AvgPool\n(5x5)", size=10.5)
    cv = d.box(316, 135, 62, 30, "Conv\n(1x1)", size=10.5)
    d.line((408, 76), (408, 190), dash=True, color="rust")
    d.text(426, 118, "shared\nweights", size=9.5, color="rust", anchor="start")
    m1 = d.op(484, 86, "x", r=11)
    m3 = d.op(484, 205, "x", r=11)
    add = d.op(524, 150, "+", r=14)
    d.arrow(pre.r, (196, 150), (196, 61), c1.l)
    d.arrow((196, 150), (196, 205), c3.l)
    d.arrow((196, 150), ap.l)
    d.link(ap, cv)
    d.arrow(c1.r, (484, 61), m1.t)
    d.arrow(c3.r, m3.l)
    d.arrow(cv.r, (398, 150), (398, 100), (m1.x - 12, 100), (m1.x + 1, 92), label="S", label_at=0.9, dy=10)
    d.arrow((398, 150), (398, 176), (m3.x - 12, 176), (m3.x + 1, 198), label="1-S", label_at=0.6, dy=-8)
    d.arrow(m1.r, (524, 86), add.t)
    d.arrow(m3.r, (524, 205), add.b)
    d.raw('<rect x="556" y="30" width="118" height="200" rx="2" fill="#DCE7FA" stroke="#0F52BA" stroke-width="1.2"/>')
    post = _context(d, 556, "Post-Global Context", "none")
    d.arrow(add.r, post.l)
    d.arrow(post.r, (690, 150))
    d.text(712, 150, "Output", size=13)
    return d
