"""Diagrams for "MixConv - Mixed Depthwise Convolutional Kernels". The multiplication counts are computed here."""
from kit import Drawing

C, K, H, W, N = 5, 3, 10, 10, 64        # input channels, kernel side, feature height and width, output channels


def _filters(d, x, y, fill, total, size=22, depth=8):
    for i in (0, 1, 3):
        d.cuboid(x + i * (size + depth + 8), y, size, size, d=depth, fill=fill)
    d.dots(x + 2 * (size + depth + 8) + size / 2, y + size / 2, gap=7)
    w = 3 * (size + depth + 8) + size + depth
    d.brace(x, x + w, y + size + 8, "Total %d Filters" % total)
    return x + w / 2


def vanilla_conv():
    """A vanilla convolution: every one of the 64 filters looks at all 5 input channels."""
    d = Drawing(740, 452)
    per_filter = C * K * K * H * W
    d.curly(40, 330, 110, 110)
    d.cuboid(70, 78, 76, 76, d=18, fill="blue")
    d.text(118, 176, "(%d, %d, %d)" % (C, H, W), size=12)
    d.text(196, 112, "*", size=30, bold=True)
    d.cuboid(236, 96, 34, 34, d=14, fill="rust")
    d.text(262, 150, "(%d, %d, %d)" % (C, K, K), size=12)
    fx = _filters(d, 440, 30, "rust", N)
    d.arrow((440, 46), (350, 46), (350, 70))
    d.text(350, 30, "Padding = 1, Stride = 1\nOutput Channels = %d" % N, size=11, anchor="end")
    d.flow(118, 196, 118, 226)
    d.cuboid(70, 250, 76, 76, d=18, fill="blue")
    d.raw('<rect x="70" y="250" width="34" height="34" fill="#F6E3D6" fill-opacity=".8" stroke="#A34A12" stroke-width="2"/>')
    d.text(190, 282, "Total Multiplications / Filter\n%d · %d · %d · %d · %d = %d" % (C, K, K, H, W, per_filter), size=12, anchor="start")
    d.flow(118, 336, 118, 362)
    for i in (0, 1, 2, 4):
        d.cuboid(64 + i * 26, 390, 5, 34, d=14, fill="mid")
    d.dots(64 + 3 * 26 + 8, 402, gap=6)
    d.text(124, 438, "Stacking of each output", size=11, color="mute")
    d.flow(216, 400, 270, 400)
    d.cuboid(296, 384, 150, 36, d=18, fill="blue")
    d.text(380, 354, "(%d, %d, %d)" % (N, H, W), size=12)
    d.text(486, 392, "Total Multiplications /\nWhole Convolution Operation\n%d · %d · %d · %d · %d · %d = %d" % (N, C, K, K, H, W, N * per_filter), size=12, anchor="start")
    return d


def depthwise_separable():
    """Depthwise then pointwise: filter each channel on its own, then mix channels with 1x1 filters."""
    d = Drawing(740, 470)
    depthwise = C * 1 * K * K * H * W
    per_point = 1 * 1 * C * H * W
    pointwise = N * per_point
    # depthwise half
    d.text(170, 14, "Depthwise convolution", bold=True)
    d.curly(24, 250, 96, 96)
    d.cuboid(52, 66, 64, 64, d=16, fill="blue")
    d.text(92, 150, "(%d, %d, %d)" % (C, H, W), size=12)
    d.text(152, 96, "*", size=28, bold=True)
    d.box(184, 82, 30, 30, fill="rust", r=0)
    d.text(199, 128, "(1, %d, %d)" % (K, K), size=12)
    _filters(d, 268, 38, "rust", C, size=20, depth=0)
    d.flow(92, 166, 92, 190)
    d.cuboid(52, 216, 64, 64, d=16, fill="blue")
    d.raw('<rect x="52" y="216" width="26" height="26" fill="#F6E3D6" fill-opacity=".8" stroke="#A34A12" stroke-width="2"/>')
    d.text(140, 240, "Total Multiplications / Filter\n1 · %d · %d · %d · %d = %d" % (K, K, H, W, K * K * H * W), size=11.5, anchor="start")
    d.flow(92, 290, 92, 314)
    for i in (0, 1, 2, 4):
        d.cuboid(46 + i * 24, 338, 5, 34, d=14, fill="rustmid")
    d.dots(46 + 3 * 24 + 8, 350, gap=6)
    d.text(20, 392, "Stacking of each output", size=11, color="mute", anchor="start")
    d.text(20, 424, "Total Multiplications /\nDepthwise Convolution Operation\n%d · 1 · %d · %d · %d · %d = %d" % (C, K, K, H, W, depthwise), size=11.5, anchor="start")
    d.arrow((180, 356), (384, 356), (384, 96), (398, 96), sw=2.6, color="grey")
    # pointwise half
    d.text(560, 14, "Pointwise convolution", bold=True)
    d.curly(400, 620, 96, 96)
    d.cuboid(428, 66, 64, 64, d=16, fill="rustmid")
    d.text(468, 150, "(%d, %d, %d)" % (C, H, W), size=12)
    d.text(528, 96, "*", size=28, bold=True)
    d.cuboid(556, 86, 18, 18, d=12, fill="grey")
    d.text(574, 128, "(%d, 1, 1)" % C, size=12)
    _filters(d, 636, 44, "grey", N, size=11, depth=5)
    d.flow(468, 166, 468, 190)
    d.cuboid(428, 216, 64, 64, d=16, fill="rustmid")
    d.cuboid(428, 216, 14, 14, d=16, fill="grey")
    d.text(520, 240, "Total Multiplications / Filter\n1 · 1 · %d · %d · %d = %d" % (C, H, W, per_point), size=11.5, anchor="start")
    d.flow(468, 290, 468, 314)
    for i in (0, 1, 2, 4):
        d.cuboid(424 + i * 24, 338, 5, 34, d=14, fill="mid")
    d.dots(424 + 3 * 24 + 8, 350, gap=6)
    d.flow(550, 352, 590, 352)
    d.cuboid(606, 338, 100, 30, d=16, fill="blue")
    d.text(664, 310, "(%d, %d, %d)" % (N, H, W), size=12)
    d.text(730, 404, "Total Multiplications / Pointwise Convolution Operation\n%d · 1 · 1 · %d · %d · %d = %d" % (N, C, H, W, pointwise), size=11.5, anchor="end")
    d.text(370, 452, "Total Multiplication / Whole Convolution Operation:  %d + %d = %d" % (pointwise, depthwise, pointwise + depthwise), size=13, bold=True)
    return d


def vanilla_vs_mixconv():
    """Vanilla depthwise convolution uses one kernel size for every channel; MixConv gives each group of channels its own."""
    d = Drawing(740, 250)
    for x, w, title in ((30, 280, "(a) Vanilla Depthwise Convolution"), (400, 310, "(b) Our proposed MixConv")):
        d.box(x, 20, w, 30, "Input Tensor", fill="blue", r=0)
        d.box(x, 170, w, 30, "Output Tensor", fill="blue", r=0)
        d.text(x + w / 2, 226, title)
    for i in range(9):
        d.arrow((48 + i * 30, 50), (48 + i * 30, 170), sw=1, color="grey")
    d.ellipse(170, 110, 130, 18, "kxk", fill="rust")
    d.text(340, 148, "channels", size=11, color="mute", anchor="middle")
    groups = ((440, "3x3", "blue", "blue"), (530, "5x5", "mid", "blue"), (650, "kxk", "rust", "rust"))
    for cx, label, fill, color in groups:
        for dx in (-22, 0, 22):
            d.arrow((cx + dx, 50), (cx + dx, 170), sw=1.2, color=color)
        d.ellipse(cx, 110, 34, 16, label, fill=fill)
    d.dots(590, 110, gap=7)
    return d


def flow():
    """MixConv end to end: split the channels into groups, give each group its own kernel size, concatenate, then mix with 1x1."""
    d = Drawing(740, 560)
    src = d.cuboid(292, 30, 96, 50, d=22, fill="rust")
    d.text(340, 54, "Input Feature Map\n(C, H, W)", size=11)
    d.text(350, 104, "C1 + C2 + C3 + ... + Cg = C", size=12, italic=True)
    xs = (70, 200, 330, 520)
    names = ("1", "2", "3", "g")
    kernels = ("3X3", "5X5", "7X7", "kXk")
    fills = ("rust", "grey", "blue", "mid")
    for x, name, kern, fill in zip(xs, names, kernels, fills):
        d.text(x + 34, 126, "Group %s" % name, size=11)
        d.cuboid(x, 152, 56, 44, d=14, fill=fill)
        d.text(x + 34, 210, "(C%s, H, W)" % name, size=11)
        d.text(x + 34, 232, "*", size=22, bold=True)
        d.box(x + 4, 244, 60, 34, "Kernel Size\n%s" % kern, fill="tint", size=10.5, r=0)
        d.flow(x + 34, 282, x + 34, 302)
        out = d.cuboid(x, 322, 56, 44, d=14, fill=fill)
        d.text(x + 34, 380, "(C%s, H, W)" % name, size=11)
        d.line((x + 34, 392), (340, 440))
    d.dots(458, 178, gap=8)
    d.dots(458, 262, gap=8)
    d.dots(458, 346, gap=8)
    d.text(630, 232, "*  Depthwise\n   Convolution", size=12, bold=True, anchor="start")
    d.text(340, 452, "Feature Map Concatenation", size=11)
    mid = d.cuboid(80, 496, 110, 44, d=18, fill="grey")
    d.text(135, 518, "Intermediate\nFeature Map (C, H, W)", size=10)
    d.arrow((340, 462), (340, 478), (150, 478), (150, 492), sw=1)
    d.text(240, 518, "x", size=22, bold=True)
    d.cuboid(274, 506, 22, 22, d=10, fill="rustmid")
    d.text(292, 548, "Kernel Size (Co, 1, 1)", size=10.5)
    d.flow(330, 518, 392, 518)
    out = d.cuboid(410, 496, 120, 44, d=18, fill="rustmid")
    d.text(470, 518, "Output Feature Map\n(Co, H, W)", size=10.5)
    d.text(566, 518, "x  Pointwise\n   Convolution", size=12, bold=True, anchor="start")
    return d


S_NET = (("stem", "112x112x16"), ("3", "112x112x16"), ("3", "56x56x24"), ("3", "56x56x24"), ("357", "28x28x40"), ("35", "28x28x40"),
         ("35", "28x28x40"), ("35", "28x28x40"), ("357", "28x28x80"), ("35", "28x28x80"), ("35", "28x28x80"), ("357", "14x14x120"),
         ("357", "14x14x120"), ("357", "14x14x120"), ("3579B", "7x7x200"), ("3579", "7x7x200"), ("3579", "7x7x200"))
M_NET = (("stem", "112x112x24"), ("3", "112x112x24"), ("357", "56x56x32"), ("3", "56x56x32"), ("3579", "28x28x40"), ("35", "28x28x40"),
         ("35", "28x28x40"), ("35", "28x28x40"), ("357", "28x28x80"), ("3579", "28x28x80"), ("3579", "28x28x80"), ("3579", "28x28x80"),
         ("3", "14x14x120"), ("3579", "14x14x120"), ("3579", "14x14x120"), ("3579", "14x14x120"), ("3579", "7x7x200"), ("3579", "7x7x200"),
         ("3579", "7x7x200"), ("3579", "7x7x200"))
STYLE = {"stem": ("white", "stem"), "3": ("grey", "3x3"), "35": ("blue", "3x3, 5x5"), "357": ("mid", "3x3, 5x5, 7x7"),
         "3579": ("rust", "3x3, 5x5, 7x7, 9x9"), "3579B": ("rustmid", "3x3, 5x5, 7x7, 9x9, 11x11")}


def _mixnet_row(d, y, name, net):
    step = 708 / len(net)
    d.text(16, y - 70, "224x224x3", size=9, rotate=-90, color="mute")
    for i, (kind, size) in enumerate(net):
        fill, label = STYLE[kind]
        h = 14 + 5.9 * len(label)
        x = 24 + i * step
        d.box(x, y - h / 2, step * 0.48, h, fill=fill, r=1)
        d.text(x + step * 0.24, y, label, size=9.5, rotate=-90, bold=kind != "stem")
        d.text(x + step * 0.78, y - 70, size, size=9, rotate=-90, color="mute")
        d.arrow((x + step * 0.48, y), (x + step, y), sw=1)
    d.arrow((10, y), (24, y), sw=1)
    d.text(370, y + 96, name, size=14, bold=True)


def mixnet():
    """MixNet-S and MixNet-M: taller blocks mix more kernel sizes. The size after each block is written above the arrow."""
    d = Drawing(740, 480)
    _mixnet_row(d, 124, "MixNet-S", S_NET)
    _mixnet_row(d, 360, "MixNet-M", M_NET)
    return d
