"""Diagrams for "Object Detection - A Quick Read"."""
from kit import Drawing


def pipeline():
    """What any object detector has to do: pick regions, describe them, then classify and regress each one."""
    d = Drawing(740, 150)
    img = d.box(60, 54, 90, 44, "Image", fill="rust", r=0)
    sel = d.box(200, 50, 130, 52, "Target Region\nSelection", fill="blue", r=8)
    feat = d.box(380, 50, 130, 52, "Feature\nExtraction of\ntargets", fill="blue", r=8, size=12)
    cls = d.ellipse(630, 30, 70, 22, "Classification", fill="rust")
    reg = d.ellipse(630, 122, 70, 22, "Regression", fill="rust")
    d.link(img, sel)
    d.link(sel, feat)
    d.arrow(feat.r, (cls.x, 34))
    d.arrow(feat.r, (reg.x, 118))
    return d


def _chain(d, y, start_x, items, r=26):
    """A row of round nodes joined by arrows; each item is (name, label written over the arrow that leads to it)."""
    x, prev, nodes = start_x, None, []
    for name, how in items:
        node = d.ellipse(x, y, r + 6, r - 6, name, fill="rust", size=9.5)
        if prev is not None:
            d.arrow(prev.r, node.l, label=how or None, size=9, dy=-26)
        prev = node
        nodes.append(node)
        x += 128
    d.line(prev.r, (prev.x2 + 14, y), dash=True)
    return nodes


def families():
    """The three families of deep-learning detectors and how each grew, method by method."""
    d = Drawing(740, 430)
    root = d.box(4, 198, 96, 44, '"Deep" Learning\nObject Detection', fill="grey", r=6, size=10)
    d.text(160, 54, "Two Stage Detectors:\nRegion Proposals based,\nAnchor Based", size=10)
    rcnn = d.ellipse(160, 108, 34, 20, "Region-\nCNN", fill="rust", size=10)
    fast = d.ellipse(330, 108, 36, 20, "Fast-RCNN", fill="rust", size=10)
    d.arrow(rcnn.r, fast.l, label="Fast Region\nExtractor method", size=9.5, dy=-18)
    for y, name, how in ((36, "R-FCN", "Full Convolutional\nNetwork"), (108, "FPN", "Feature\nPyramid"), (176, "Mask-\nRCNN", "Segmentation")):
        node = d.ellipse(540, y, 34, 20, name, fill="rust", size=10)
        if y == 108:
            d.arrow(fast.r, node.l, label=how, size=9.5, dy=-16)
        else:
            d.arrow((330, 88 if y < 108 else 128), (330, y), node.l, label=how, size=9.5, label_at=0.75, dy=-16 if y < 108 else -10)
        d.line(node.r, (node.x2 + 60, y), dash=True)
    d.text(150, 204, "Anchor Free", size=10)
    free = _chain(d, 244, 180, (("DenseBox", ""), ("CornerNet", ""), ("ExtremeNet", "Localization\nImprovement"),
                               ("Feature\nSelective\nAnchor Free", "FPN"), ("CenterNet", "Centerpoint\napproach")), r=30)
    d.text(180, 326, "Single Shot Detectors\nAnchor Based", size=10)
    shot = _chain(d, 384, 180, (("MultiBox", ""), ("YOLO", "Grid\nRegression"), ("SSD", "RPN"),
                               ("YOLOv2", "Batch Norm.\nMultiscale"), ("YOLOv3", "Multiclass\nClassification")))
    d.text(shot[4].cx - 62, 408, "Darknet-53", size=9.5, color="mute")
    d.arrow(root.r, (116, 220), (116, 108), rcnn.l)
    d.arrow((116, 220), (116, 244), free[0].l)
    d.arrow((116, 244), (116, 384), shot[0].l)
    return d
