"""Diagrams for the three YOLOv4 reads: Introduction, Bag of Specials and Final Verdict."""
from kit import Drawing, tree, stack_trees

# ------------------------------------------------------------------ what goes where (the trees)
BOF_BACKBONE = "1. Class Label Smoothing\n2. Mosaic and CutMix Data\n   Augmentations.\n3. DropBlock Regularization"
BOF_DETECTOR = "1. CIoU-loss\n2. Cross Minibatch\n   Normalization.\n3. DropBlock regularization\n4. Mosaic data augmentation\n5. Self-Adversarial Training"
BOS_BACKBONE = "1. Mish Activation\n2. Cross Stage Partial\n   Connections\n3. Multi input Weighted\n   Residual Connections"
BOS_DETECTOR = "1. Mish Activation\n2. Spatial Pyramid Pooling\n3. Spatial Attention Module\n4. Path Aggregation Networks\n5. DIoU-NMS"


def _bag(title, backbone, detector):
    return (title, "blue", [("Backbone", "blue", [(backbone, "grey", [])]), ("Detector", "blue", [(detector, "grey", [])])])


def bag_of_freebies():
    """Bag of Freebies: what YOLOv4 uses at training time, for the backbone and for the detector."""
    return tree(_bag("Bag of Freebies\nUsed at training time", BOF_BACKBONE, BOF_DETECTOR), width=560)


def bag_of_specials():
    """Bag of Specials: what YOLOv4 uses at inference time, for the backbone and for the detector."""
    return tree(_bag("Bag of Specials\nUsed at inference time", BOS_BACKBONE, BOS_DETECTOR), width=560)


def optimization_approaches():
    """The two kinds of optimisation in YOLOv4 and what each contains."""
    return tree(("Optimization\nApproaches/Methods", "rust", [
        _bag("Bag of Freebies\nUsed at training time", BOF_BACKBONE, BOF_DETECTOR),
        _bag("Bag of Specials\nUsed at inference time", BOS_BACKBONE, BOS_DETECTOR)]), size=9.5, gap=8, pad=8, width=740)


def baseline_architecture():
    """The final YOLOv4 recipe: backbone, neck and head."""
    return tree(("YOLOv4 Main\nRecipies", "rust", [
        ("Backbone", "blue", [("CSPDarknet53", "grey", [])]),
        ("Neck", "blue", [("Concatenated\nPath Aggregation\nNetworks\nwith SPP\nModules", "grey", [])]),
        ("Head", "blue", [("YOLOv3 Head\n\nNXNX(5 + C)\nFeature Map", "grey", [])])]), size=12, gap=60, width=640)


def candidate_modules():
    """Everything the YOLOv4 authors considered during the ablation study, grouped by where it goes."""
    leaf = lambda s: [(s, "grey", [])]
    architectures = ("Baseline\nArchitectures", "blue", [
        ("Backbone", "rust", leaf("1. VGG16\n2. ResNet-50\n3. SpineNet\n4. EfficientNet-B0/B7\n5. CSPResNeXt50\n6. CSPDarknet53")),
        ("Neck", "rust", [
            ("Additional Blocks", "tint", leaf("1. SPP\n2. ASPP\n3. RFB\n4. SAM")),
            ("Path Aggregation\nBlocks", "tint", leaf("1. FPN\n2. PAN\n3. NAS-FPN\n4. Fully-connected FPN\n5. BiFPN\n6. ASFF\n7. SFAM"))]),
        ("Head", "rust", [
            ("Dense Predictions", "tint", leaf("1. RPN\n2. SSD\n3. YOLO\n4. RetinaNet\n5. CornerNet\n6. CenterNet\n7. MatrixNet\n8. FCOS")),
            ("Sparse Predictions", "tint", leaf("1. Faster R-CNN\n2. R-FCN\n3. Mask RCNN\n4. RepPoints"))])])
    methods = ("Tunable and Performance\nOptimization Methodologies", "blue", [
        ("Data\nAugmentation", "rust", leaf("1. CutMix\n2. MixUp\n3. CutOut")),
        ("Activation\nFunctions", "rust", leaf("1. ReLU\n2. leaky-ReLU\n3. parametric-ReLU\n4. ReLU6\n5. SELU\n6. Mish")),
        ("Regularization\nMethods", "rust", leaf("1. DropOut\n2. DropPath\n3. Spatial DropOut\n4. DropBlock")),
        ("Normalization\nMethods", "rust", leaf("1. Batch Normalization\n2. Cross-GPU Batch\n   Normalization\n3. Cross-Iteration Batch\n   Normalization (CBN)")),
        ("Regression Loss\nFunctions", "rust", leaf("1. MSE\n2. IoU\n3. GIoU\n4. CIoU\n5. DIoU"))])
    top = Drawing(740, 40)
    top.box(300, 4, 140, 32, "YOLOv4 Considerable\nCandidates", fill="rustmid", size=10.5, r=5)
    return stack_trees(top, tree(architectures, size=9.5, gap=8, pad=8, width=740), tree(methods, size=9.5, gap=8, pad=8, width=740), gap=10)


# ------------------------------------------------------------------ introduction
def _cloud(d, x, y, w, h, label, items, label_side="right"):
    d.group(x, y, w, h, fill="rust", r=h / 2)
    for (ix, iy), name in items:
        d.ellipse(x + ix, y + iy, 26, 15, name, fill="blue", size=9.5)
    if label_side == "right":
        d.text(x + w + 10, y + h / 2, label, size=11, anchor="start")
    else:
        d.text(x + w / 2, y + h + 12, label, size=11)
    return (x, y, w, h)


def detection_flow():
    """How an object detector is put together: input, backbone, an optional neck, and a prediction head."""
    d = Drawing(740, 520)
    d.text(370, 14, "OBJECT DETECTION FLOW", bold=True)
    _cloud(d, 60, 34, 190, 86, "Sparse Prediction", (((52, 28), "Mask\nRCNN"), ((36, 62), "R-FCN"), ((130, 56), "Faster\nRCNN")))
    _cloud(d, 420, 46, 190, 90, "Dense Prediction", (((56, 26), "RPN"), ((134, 30), "YOLO"), ((52, 64), "Corner\nNet"), ((120, 66), "FCN")))
    head = d.box(300, 178, 96, 40, "Head", fill="rustmid", r=6)
    d.text(410, 198, "PREDICTION STAGE", size=11, anchor="start")
    d.arrow((190, 120), (310, 176), sw=2.4, color="grey")
    d.arrow((470, 136), (386, 176), sw=2.4, color="grey")
    inp = d.box(150, 280, 86, 40, "Image", fill="tint", r=0)
    d.text(193, 268, "INPUT", size=11)
    back = d.box(300, 280, 96, 40, "Backbone", fill="rustmid", r=6)
    d.text(430, 262, "FEATURE EXTRACTION\nSTAGE", size=11)
    d.link(inp, back)
    d.link(back, head)
    _cloud(d, 516, 238, 196, 110, "Cloud of Backbone", (((36, 40), "FCN"), ((84, 22), "ResNet"), ((152, 36), "DenseNet"), ((86, 68), "VGG"), ((150, 84), "MobileNet")), "below")
    d.arrow((516, 300), (398, 300), sw=2.4, color="grey")
    neck = d.box(296, 378, 104, 44, "Neck (Optional,\nbut effective)", fill="rustmid", r=6, size=11)
    d.text(412, 400, "RICH SEMANTIC EXTRACTION\nMETHODS", size=11, anchor="start")
    d.link(neck, back)
    _cloud(d, 248, 444, 200, 68, "Cloud of Neck", (((34, 30), "PAN"), ((84, 18), "RFB"), ((138, 26), "FPN"), ((70, 50), "ASFF"), ((130, 52), "BiFPN")))
    d.arrow((348, 444), (348, 424), sw=2.4, color="grey")
    return d


def pan_modification():
    """YOLOv4's change to PAN: the two feature maps are concatenated instead of added."""
    d = Drawing(740, 190)
    for x, sym, name, title in ((90, "+", "addition", "(a) PAN"), (430, "c", "concatenation", "(b) Our modified PAN")):
        d.cuboid(x + 70, 20, 70, 8, d=12, fill="mid")
        d.cuboid(x - 40, 70, 60, 8, d=12, fill="mid")
        op = d.op(x + 110, 74, sym, r=10, fill="blue")
        d.text(x + 130, 74, name, size=11, anchor="start")
        d.cuboid(x + 60, 116, 90, 8, d=12, fill="mid")
        d.arrow((x + 34, 74), (op.x, 74))
        d.arrow((x + 110, 128 - 24), (x + 110, op.y2), sw=1.2) if False else d.arrow((x + 110, 104), (x + 110, op.y2))
        d.arrow((x + 110, op.y), (x + 110, 30))
        d.text(x + 100, 160, title, bold=True)
    return d


# ------------------------------------------------------------------ bag of specials
def _classifier(d, cx, cy):
    d.ellipse(cx, cy, 62, 42, fill="blue")
    d.text(cx, cy, "Classifier\nArchitecture\nAvgPool + Linear\n+ Softmax", size=9.5)


def densenet():
    """A DenseNet-like backbone: a stem convolution, then M chains of dense block + transition block, then a classifier."""
    d = Drawing(740, 330)
    img = d.box(8, 146, 96, 44, "Input Image\nH X W X 3", fill="rust", r=0, size=11)
    d.group(124, 10, 470, 310, fill="tint", r=30)
    d.text(360, 30, "Backbone: DenseNet Alike Structure", bold=True)
    stem = d.cuboid(144, 146, 100, 46, d=14, fill="white")
    d.text(194, 169, "Conv+BN+ReLU\nOut Filters: F\nStride: 2", size=10)
    d.text(200, 104, "Output of ConvLayer\n(H/2, W/2, F)", size=11)
    d.group(282, 50, 300, 260, fill="grey", r=2)
    d.text(432, 68, "Chains of Dense and Transition Blocks X M", size=10.5, bold=True)
    d.text(432, 100, "Growth Rate: k\nTotal Concatenations in each Dense Block: N", size=10)
    db = d.cuboid(300, 154, 108, 44, d=14, fill="white", label="Dense Block", size=11)
    tb = d.cuboid(444, 154, 112, 44, d=14, fill="white", label="Transition Block", size=10.5)
    d.text(432, 236, "Output of DB+TB Chains", size=11)
    d.text(432, 268, "( H / (2·M),  W / (2·M),\n(F + (2^M − 1)·N·k) / 2^M )", size=10.5, italic=True)
    d.arrow(img.r, (stem.x, 168))
    d.arrow((stem.x2, 168), (282, 168))
    d.arrow((db.x2, 176), (tb.x, 176))
    _classifier(d, 668, 168)
    d.arrow((582, 168), (606, 168))
    return d


def _res_dense(dense):
    """A residual block (addition keeps 16 channels) or a dense block (concatenation grows by k each layer)."""
    d = Drawing(740, 356)
    k = "k" if dense else "16"
    d.cuboid(6, 150, 78, 40, d=12, fill="white", label="16X100X100\nCXWXW", size=9.5)
    d.text(50, 116, "Input\nFeature Map", size=10, color="mute")
    d.box(108, 8, 524, 340, fill="white", r=0)
    title = ("DENSE BLOCK\nGrowth Rate: k  |  Total Concatenations/Dense Layers: N  |  Total Output Channels: N * k" if dense
             else "ResNet Block\nTotal Concatenations: N  |  Total Output Channels: 16")
    d.text(370, 26, title, size=10, bold=True)
    d.arrow((96, 164), (108, 164))
    d.arrow((632, 164), (644, 164))
    d.cuboid(646, 150, 78, 40, d=12, fill="white", label=("(16+N*k)\nX100X100" if dense else "16X100X100\nCXWXH"), size=9.5)
    d.text(690, 116, "Output of\n%s Block" % ("Dense" if dense else "ResNet"), size=10, color="mute")
    # first layer
    a = d.stack(140, 94, 66, 40, 3, fill="rust")
    d.brace(a.x, a.x2, a.y - 6, "16 Channels", up=True)
    c1 = d.cuboid(246, 88, 78, 40, d=12, fill="grey", label="Conv+BN+ReLU\nOut Filters: %s" % k, size=9)
    b = d.stack(368, 94, 66, 40, 3, fill="blue")
    d.brace(b.x, b.x2, b.y - 6, "%s Channels" % k, up=True)
    op1 = d.op(478, 108, "c" if dense else "+", r=11, fill="blue")
    d.text(478, 150, "Channel\nConcatenation" if dense else "Channel Element\nWise Addition", size=9.5)
    if dense:
        c = d.stack(520, 100, 60, 40, 3, fill="rust")
        d.stack(541, 100, 60, 40, 2, fill="blue")
    else:
        c = d.stack(520, 94, 66, 40, 3, fill="mid")
    d.brace(520, 612 if dense else c.x2, 74, ("16 + k Channels" if dense else "16 Channels"), up=True)
    d.arrow((a.x2, 110), (c1.x, 110))
    d.arrow((c1.x2, 110), (b.x - 2, 110))
    d.arrow((b.x2 - 12, 108), (op1.x, 108))
    d.arrow((op1.x2, 108), (518, 108))
    d.arrow((226, 110), (226, 52), (478, 52), (478, op1.y), dash=True)
    # second layer
    c2 = d.cuboid(520, 240, 78, 40, d=12, fill="grey", label="Conv+BN+ReLU\nOut Filters: %s" % k, size=9)
    d.arrow((560, 142), (560, 226))
    e = d.stack(396, 246, 66, 40, 3, fill="dark")
    d.brace(e.x, e.x2, e.y - 6, "%s Channels" % k, up=True)
    op2 = d.op(352, 260, "c" if dense else "+", r=11, fill="blue")
    d.text(352, 216, "Channel\nConcatenation" if dense else "Channel Element\nWise Addition", size=9.5)
    if dense:
        f = d.stack(226, 252, 60, 40, 2, fill="rust")
        d.stack(240, 252, 60, 40, 2, fill="blue")
        d.stack(254, 252, 60, 40, 2, fill="dark")
        d.brace(226, 322, 302, "16 + 2k Channels")
    else:
        f = d.stack(226, 246, 66, 40, 3, fill="rustmid")
        d.brace(f.x, f.x2, 296, "16 Channels")
    d.arrow((c2.x, 262), (e.x2 + 2, 262))
    d.arrow((e.x - 2, 260), (op2.x2, 260))
    d.arrow((op2.x, 260), (324 if dense else f.x2 + 2, 260))
    d.arrow((612, 126), (622, 126), (622, 318), (352, 318), (352, op2.y2), dash=True)
    d.arrow((216, 264), (150, 264), label="N Times", dy=-10)
    if dense:
        d.text(300, 176, "1st Dense Layer", size=10, color="mute")
        d.text(490, 334, "2nd Dense Layer", size=10, color="mute")
        d.text(170, 290, "N Dense Layers", size=9.5, color="mute")
    return d


def resnet_block():
    """A ResNet block: each layer's output is added to its input, so the channel count stays at 16."""
    return _res_dense(False)


def dense_block():
    """A dense block: each layer's k new channels are concatenated to everything before, so channels keep growing."""
    return _res_dense(True)


def csp():
    """Baseline DenseNet, and the CSP version that sends part of the base layer around the dense block."""
    d = Drawing(740, 500)
    # baseline
    d.group(196, 8, 320, 170, "Baseline DenseNet", fill="tint", r=24)
    img = d.box(60, 76, 100, 40, "Input Image\nH X W X 3", fill="rust", r=0, size=11)
    base = d.cuboid(212, 76, 70, 40, d=12, fill="white", label="Base Layer", size=10.5, above="(H/2, W/2, F)")
    db = d.box(312, 62, 80, 68, "Dense Block", fill="grey", r=0, size=10.5)
    tb = d.box(410, 62, 90, 68, "Transition\nBlock", fill="grey", r=0, size=10.5)
    d.text(352, 50, "Output Xk", size=9.5, color="mute")
    d.text(455, 50, "Output XT", size=9.5, color="mute")
    d.arrow(img.r, (base.x, 96))
    d.arrow((base.x2, 96), db.l)
    d.link(db, tb)
    d.brace(214, 500, 140, "Chain of DB+TB")
    _classifier(d, 600, 96)
    d.arrow(tb.r, (538, 96))
    # CSP
    d.group(96, 196, 520, 296, fill="tint", r=24)
    d.text(520, 214, "CSP + DenseNet", bold=True)
    img = d.box(2, 330, 84, 40, "Input Image\nH X W X 3", fill="rust", r=0, size=10)
    base = d.cuboid(110, 330, 70, 40, d=12, fill="white", label="Base Layer", size=10.5, above="(H/2, W/2, F)")
    b1 = d.cuboid(222, 268, 76, 40, d=12, fill="white", label="Base Layer 1", size=10.5, above="New Base Layer\nfor Dense Block\n(H/2, W/2, F1)")
    b2 = d.cuboid(222, 392, 76, 40, d=12, fill="white", label="Base Layer 2", size=10.5, above="(H/2, W/2, F2)")
    db = d.box(330, 252, 78, 68, "Dense Block", fill="grey", r=0, size=10.5)
    tb = d.box(424, 252, 78, 68, "Transition\nBlock", fill="grey", r=0, size=10.5)
    add = d.op(524, 286, "+", r=11)
    tb2 = d.box(546, 252, 62, 68, "Transition\nBlock", fill="grey", r=0, size=9.5)
    for box, name in ((db, "Xk"), (tb, "XT"), (tb2, "XU")):
        d.text(box.cx, 240, "Output %s" % name, size=9.5, color="mute")
    d.arrow(img.r, (base.x, 350))
    d.arrow((base.x2, 350), (204, 350), (204, 288), (b1.x, 288))
    d.arrow((204, 350), (204, 412), (b2.x, 412))
    d.arrow((b1.x2, 286), db.l)
    d.link(db, tb)
    d.arrow(tb.r, add.l)
    d.arrow((b2.x2, 412), (524, 412), add.b)
    d.arrow(add.r, tb2.l)
    d.brace(330, 508, 330, "Partial Dense Block")
    d.brace(516, 608, 330, "Partial Transition\nBlock")
    d.brace(112, 606, 452, "Chain of CSP Connections")
    _classifier(d, 676, 286)
    d.arrow(tb2.r, (615, 286))
    return d


def spp():
    """Spatial pyramid pooling in YOLOv3-SPP: four parallel branches with growing kernels, concatenated to 4x the channels."""
    d = Drawing(740, 360)
    d.box(4, 4, 732, 352, fill="white", r=0)
    d.text(150, 22, "YOLOv3-SPP Backbone Architecture", bold=True)
    d.brace(300, 480, 52, "FCN-SPP Version", up=True)
    fcn = d.cuboid(20, 158, 76, 50, d=12, fill="grey", label="YOLO-FCN\nNetwork", size=10)
    fmap = d.stack(138, 170, 54, 30, 3, step=5, fill="mid", label="CXHXW", size=9.5)
    d.text(170, 140, "Output Feature\nMap", size=10)
    d.arrow((fcn.x2, 184), (fmap.x, 184))
    add = d.op(470, 190, "c", r=10, fill="blue")
    for i, (kern, fill) in enumerate((("(1, 1)", "rustmid"), ("(3, 3)", "blue"), ("(9, 9)", "dark"), ("(13, 13)", "rust"))):
        y = 84 + i * 72
        conv = d.cuboid(262, y, 40, 24, d=8, fill="white", label="Conv", size=9.5)
        d.text(316, y - 18, "K: %s  S: 1, P: 1" % kern, size=9.5, italic=True, anchor="start")
        out = d.stack(350, y + 2, 50, 22, 3, step=4, fill=fill, label="CXHXW", size=9)
        d.arrow((fmap.x2, 180), (conv.x - 2, y + 12), sw=1.1)
        d.arrow((conv.x2, y + 10), (out.x - 2, y + 10), sw=1.1)
        d.arrow((out.x2, y + 10), (470 if i in (0, 3) else 440, y + 10), (470 if i in (0, 3) else 440, 190) if i not in (0, 3) else (470, add.y if i == 0 else add.y2), sw=1.1,
                head=i in (0, 3))
    d.arrow((440, 190), (add.x, 190), sw=1.1)
    cat = d.stack(506, 196, 50, 26, 6, step=5, fill="blue", label="4CXHXW", size=9)
    d.text(540, 138, "4XC Channels\nafter concatenation", size=10)
    d.arrow((add.x2, 190), (504, 200), sw=1.1)
    conv = d.cuboid(622, 196, 40, 26, d=8, fill="white", label="Conv", size=9.5)
    d.text(646, 166, "K: (1, 1)\nS: 1, P: 1", size=9.5, italic=True)
    d.arrow((cat.x2 - 22, 210), (conv.x - 2, 210), sw=1.1)
    head = d.box(646, 250, 90, 46, "YOLOv3 Head\nHXWX(KX(5+C))", fill="rust", r=0, size=9.5)
    d.arrow((648, 224), (680, 250), sw=1.1)
    return d


def sam():
    """Spatial attention: the original SAM pools across channels first; YOLOv4's version applies the convolution point-wise."""
    d = Drawing(740, 400)
    d.box(120, 4, 500, 150, fill="white", r=0, sw=2)
    d.text(370, 20, "Modified SAM -- Point Wise", bold=True)
    src = d.cuboid(140, 74, 56, 30, d=10, fill="blue", label="CXHXW", size=10, above="Output\nFeature Map")
    conv = d.cuboid(290, 76, 44, 26, d=8, fill="white", label="Conv", size=10, above="Kernel (7, 7)")
    sig = d.cuboid(374, 72, 5, 32, d=8, fill="dark")
    d.text(380, 120, "Sigmoid\nLayer", size=10, color="mute")
    att = d.cuboid(424, 74, 56, 30, d=10, fill="rustmid")
    d.text(456, 46, "Spatial\nAttention Map", size=10, color="mute")
    mul = d.op(528, 90, "x", r=10)
    out = d.cuboid(560, 74, 50, 30, d=10, fill="solid", label="CXHXW", size=10)
    d.arrow((src.x2, 90), (conv.x, 90))
    d.arrow((conv.x2, 90), (sig.x, 90))
    d.arrow((sig.x2, 90), (att.x, 90))
    d.arrow((att.x2, 90), (mul.x, 90))
    d.arrow((mul.x2, 90), (out.x, 90))
    d.arrow((224, 90), (224, 136), (528, 136), (528, mul.y2))

    d.box(60, 166, 660, 228, fill="white", r=0, sw=2)
    d.text(390, 182, "Original SAM -- Spatial Wise", bold=True)
    src = d.cuboid(76, 278, 56, 30, d=10, fill="blue", label="CXHXW", size=10, above="Output\nFeature Map")
    avg = d.cuboid(238, 216, 8, 30, d=10, fill="rustmid")
    d.text(196, 222, "Average Pooling\nalong channel axis", size=9.5, anchor="end") if False else d.text(244, 198, "Average Pooling along channel axis", size=9.5)
    mx = d.cuboid(238, 336, 8, 30, d=10, fill="mid")
    d.text(244, 380, "MaxPooling along channel axis", size=9.5)
    cat = d.op(312, 294, "c", r=10, fill="blue")
    d.text(300, 322, "Concatenation", size=9.5)
    both = d.cuboid(344, 280, 8, 30, d=10, fill="rustmid")
    d.cuboid(354, 280, 8, 30, d=10, fill="mid")
    conv = d.cuboid(404, 280, 44, 26, d=8, fill="white", label="Conv", size=10, above="Kernel (7, 7)")
    sig = d.cuboid(486, 276, 5, 32, d=8, fill="dark")
    d.text(492, 324, "Sigmoid\nLayer", size=10, color="mute")
    att = d.cuboid(528, 280, 8, 30, d=10, fill="rustmid")
    d.text(540, 248, "Spatial\nAttention Map", size=10, color="mute")
    mul = d.op(590, 294, "x", r=10)
    out = d.cuboid(634, 278, 50, 30, d=10, fill="solid", label="CXHXW", size=10)
    d.arrow((src.x2, 294), (170, 294), (170, 232), (avg.x, 232))
    d.arrow((170, 294), (170, 352), (mx.x, 352))
    d.arrow((avg.x2, 230), (312, 230), (312, cat.y))
    d.arrow((mx.x2, 350), (312, 350), (312, cat.y2))
    d.arrow((cat.x2, 294), (both.x, 294))
    d.arrow((372, 294), (conv.x, 294))
    d.arrow((conv.x2, 294), (sig.x, 294))
    d.arrow((sig.x2, 294), (att.x, 294))
    d.arrow((att.x2, 294), (mul.x, 294))
    d.arrow((mul.x2, 294), (out.x, 294))
    d.arrow((150, 294), (150, 386), (590, 386), (590, mul.y2))
    return d


def _pan(d, y0, modified):
    widths = (150, 110, 80, 56)                   # bottom to top
    step = 58
    d.text(370, y0 + 14, "Modified PAN" if modified else "Original PAN", bold=True)
    d.arrow((30, y0 + 292), (30, y0 + 80), color="grey", sw=3)
    d.text(78, y0 + 56, "Many Layers (100+)", size=10.5)
    d.text(72, y0 + 190, "Downsampling\nLayers", size=10)
    d.cuboid(112, y0 + 296, 190, 4, d=7, fill="rust")
    back, ps, ns = [], [], []
    for i, w in enumerate(widths):
        y = y0 + 250 - i * step
        back.append(d.cuboid(112, y, w, 14, d=8, fill="blue"))
        ps.append(d.cuboid(330, y + 10, w, 14, d=8, fill="grey", label="P%d" % (i + 2), size=9.5))
        n = d.cuboid(520, y - 6 if modified else y + 10, w, 14, d=8, fill="mid", label="N%d" % (i + 2), size=9.5)
        if modified:
            d.cuboid(520, y + 14, w, 12, d=8, fill="grey")
        ns.append(n)
    for i in range(4):
        d.arrow((back[i].x2, back[i].cy + 4), (330, ps[i].cy + 4), sw=1)
        if i:
            d.arrow((back[i].x + 24, back[i - 1].y), (back[i].x + 24, back[i].y2), sw=1)
            d.arrow((ps[i].x + 24, ps[i].y2), (ps[i].x + 24, ps[i - 1].y), sw=1)
            d.arrow((ns[i - 1].x + 24, ns[i - 1].y), (ns[i].x + 24, ns[i].y2 + (20 if modified else 0)), sw=1)
            d.arrow((ps[i].x2, ps[i].cy + 4), (520, ns[i].cy + (24 if modified else 4)), sw=1)
        d.arrow((ns[i].x2 + 4, ns[i].cy + 4), (690, ns[i].cy + 4), sw=1)
    d.arrow((142, back[3].y), (142, y0 + 36), (360, y0 + 36), (360, ps[3].y), sw=1)
    d.arrow((400, ps[0].y2), (400, y0 + 312), (560, y0 + 312), (560, ns[0].y2 + (20 if modified else 0)), sw=1)
    d.text(440, y0 + 60, ("Channel concatenation\nat each layer" if modified else "Element wise addition\nat each layer"), size=10)
    d.text(590, y0 + 40, "Only ~10 Layers, with a\nConvolution between levels", size=10)
    d.cuboid(690, y0 + 76, 4, 220, d=8, fill="grey")
    d.text(676, y0 + 16, "ROIAlign + fusion\nof all layers", size=10)
    d.arrow((704, y0 + 180), (736, y0 + 180), sw=1)


def pan():
    """Path Aggregation Network: a short bottom-up path (N2..N5) on top of FPN. YOLOv4 concatenates where PAN added."""
    d = Drawing(740, 668)
    d.box(2, 2, 736, 326, fill="white", r=0, sw=2)
    _pan(d, 4, True)
    d.box(2, 338, 736, 326, fill="white", r=0, sw=2)
    _pan(d, 340, False)
    return d
