"""Diagrams for "Understanding Attention Modules: CBAM and BAM"."""
from kit import Drawing


def base_module():
    """An attention module: squeeze the feature map, score it, and multiply the scores back in."""
    d = Drawing(740, 250)
    src = d.cuboid(8, 92, 78, 46, d=16, label="Input\nFeature Map")
    d.group(128, 18, 300, 176, "Basic Attention Module", fill="tint")
    neck = d.cuboid(152, 98, 46, 34, d=14, fill="rust")
    d.text(182, 62, "Convolutional\nBottleneck", size=12)
    mlp = d.net(270, 88, 34, 54)
    d.text(287, 62, "Multi Layer\nPerceptron", size=12)
    sig = d.op(380, 115, "s", r=14, fill="rust")
    d.text(380, 84, "Sigmoid", size=12)
    amap = d.cuboid(462, 92, 16, 50, d=14, fill="mid")
    d.text(477, 44, "2-D/3-D\nActivation Map", size=12)
    mul = d.op(552, 115, "x", r=13)
    d.text(566, 170, "Channel Wise\nMultiplication", size=12)
    out = d.cuboid(612, 92, 108, 46, d=16, fill="mid", label="Refined Input\nFeature Map")
    d.arrow((src.x2, 115), (neck.x, 115))
    d.arrow((neck.x2, 115), (mlp.x, 115))
    d.arrow((mlp.x2, 115), (sig.x, 115))
    d.arrow((sig.x2, 115), (amap.x, 115))
    d.arrow((amap.x2, 115), (mul.x, 115))
    d.arrow((mul.x2, 115), (out.x, 115))
    d.arrow((47, 138), (47, 226), (552, 226), (552, mul.y2))
    return d


def cbam_overview():
    """CBAM: channel attention first, then spatial attention, each multiplied into the feature."""
    d = Drawing(740, 220)
    d.group(8, 8, 724, 204, "Convolutional Block Attention Module", fill="white", size=15)
    src = d.cuboid(40, 112, 58, 58, d=18)
    d.text(78, 78, "Input Feature")
    ch = d.group(170, 58, 118, 90, fill="tint")
    d.text(229, 82, "Channel\nAttention\nModule", size=12)
    d.cuboid(196, 124, 58, 10, d=6, fill="mid")
    sp = d.group(390, 58, 118, 90, fill="rust")
    d.text(432, 102, "Spatial\nAttention\nModule", size=12)
    d.cuboid(478, 86, 6, 44, d=12, fill="rustmid")
    m1 = d.op(338, 178, "x", r=12)
    m2 = d.op(560, 178, "x", r=12)
    out = d.cuboid(622, 120, 58, 58, d=18, fill="mid")
    d.text(660, 82, "Refined Feature")
    d.arrow((116, 178), (m1.x, 178), sw=2)
    d.arrow((m1.x2, 178), (m2.x, 178), sw=2)
    d.arrow((m2.x2, 178), (622, 178), sw=2)
    d.arrow((130, 178), (130, 103), (170, 103))
    d.arrow((288, 103), (338, 103), (338, m1.y))
    d.arrow((362, 178), (362, 103), (390, 103))
    d.arrow((508, 103), (560, 103), (560, m2.y))
    return d


def bam_structure():
    """BAM: a channel branch and a dilated spatial branch are added, squashed by a sigmoid, and applied to F."""
    d = Drawing(740, 300)
    src = d.cuboid(34, 208, 44, 52, d=16, fill="white")
    d.text(56, 282, "Input tensor F", size=12)
    # channel branch
    d.text(132, 22, "Global avg pool", size=12)
    v1 = d.cuboid(214, 34, 40, 9, d=6, fill="blue")
    v2 = d.cuboid(318, 34, 22, 9, d=6, fill="blue")
    v3 = d.cuboid(404, 34, 40, 9, d=6, fill="mid")
    d.arrow((v1.x2 + 4, 36), (v2.x - 4, 36), label="FC")
    d.arrow((v2.x2 + 4, 36), (v3.x - 4, 36), label="FC")
    d.text(332, 12, "Channel = C/r", size=11, color="mute")
    d.text(560, 28, "Channel attention\nMc(F) ∈ R^(C×1×1)", size=12)
    # spatial branch
    a = d.cuboid(214, 96, 22, 60, d=12, fill="rust")
    b = d.cuboid(326, 96, 22, 60, d=12, fill="rust")
    c = d.cuboid(440, 90, 8, 66, d=14, fill="rustmid")
    d.arrow((a.x2 + 4, 124), (b.x - 4, 124), label="3x3\nconv", dy=22)
    d.text(293, 92, "x 2", bold=True)
    d.arrow((b.x2 + 4, 124), (c.x - 4, 124), label="1x1\nconv", dy=22)
    d.text(296, 176, "with dilation value d", size=12)
    d.brace(214, 360, 190, "Channel = C/r")
    d.text(468, 180, "Spatial attention\nMs(F) ∈ R^(1×H×W)", size=12)
    d.arrow((66, 192), (120, 70), (208, 36), curve=True)
    d.arrow((72, 192), (130, 140), (208, 124), curve=True, label="1x1\nconv", dx=18, dy=30)
    add = d.op(518, 118, "+", r=11)
    sig = d.op(562, 118, "s", r=11)
    att = d.cuboid(610, 96, 44, 48, d=16, fill="mid")
    d.text(640, 168, "BAM attention\nM(F)", size=12, bold=True)
    d.arrow((v3.x2 + 2, 38), (518, 60), (518, add.y), curve=True)
    d.arrow((c.x2 + 4, 118), (add.x, 118))
    d.arrow((add.x2, 118), (sig.x, 118))
    d.arrow((sig.x2, 118), (att.x, 118))
    mul = d.op(676, 240, "x", r=11)
    plus = d.op(712, 240, "+", r=11)
    d.arrow((94, 240), (mul.x, 240))
    d.arrow((mul.x2, 240), (plus.x, 240))
    d.arrow((att.x2, 112), (676, 130), (676, mul.y), curve=True)
    d.arrow((640, 240), (640, 272), (712, 272), (712, plus.y2))
    d.arrow((plus.x2, 240), (738, 240))
    return d


def _residual_stage(d, x, title, with_cbam):
    """One residual stage: a row of convolutional filters, with CBAM after each pair when with_cbam."""
    if with_cbam:
        g = d.group(x, 34, 232, 122, fill="white")
        d.text(x + 116, 54, title, size=12)
        cx = x + 14
        for k, (a, b) in enumerate((("C1", "C2"), ("C3", "C4"))):
            for j, name in enumerate((a, b)):
                d.cuboid(cx + j * 22, 82, 10, 40, d=10, fill="blue")
                d.text(cx + j * 22 + 8, 134, name, size=11)
            d.arrow((cx + 44, 100), (cx + 56, 100))
            d.box(cx + 56, 80, 44, 40, "CBAM", fill="rust", size=11)
            if k == 0:
                d.arrow((cx + 100, 100), (cx + 108, 100))
            cx += 108
        return g
    g = d.group(x, 34, 126, 150, fill="white")
    d.text(x + 63, 56, title, size=12)
    for j, name in enumerate(("C1", "C2", "Ci")):
        px = x + 12 + j * 26 + (20 if j == 2 else 0)
        d.cuboid(px, 90, 12, 46, d=12, fill="blue")
        d.text(px + 10, 150, name, size=11)
    d.dots(x + 76, 108, gap=5)
    d.text(x + 63, 170, "Convolutional\nFilters", size=10)
    return g


def bam_in_resnet():
    """Where BAM goes in ResNet: once after each stage, at the bottleneck before pooling."""
    d = Drawing(740, 200)
    img = d.box(2, 82, 50, 52, "Input\nImage", fill="grey", size=11)
    x = 64
    d.arrow((img.x2, 108), (x, 108))
    for i in range(3):
        g = _residual_stage(d, x, "Residual Block\nLayer %d" % (i + 1), False)
        if i == 2:
            d.line((g.x2 + 4, 108), (g.x2 + 26, 108), dash=True)
            d.text(g.x2 + 22, 70, "Chains of\nResidual\nBlocks", size=10)
            break
        d.line((g.x2 + 6, 16), (g.x2 + 6, 190))
        bam = d.box(g.x2 + 14, 86, 42, 44, "BAM", fill="rust", size=12)
        d.text(bam.cx, 36, "Bottleneck\nLayer of\nResNet", size=10)
        pool = d.cuboid(bam.x2 + 10, 90, 7, 38, d=9, fill="rustmid")
        d.text(pool.cx + 2, 66, "Pooling", size=10)
        d.arrow((g.x2, 108), (bam.x, 108))
        d.arrow((bam.x2, 108), (pool.x, 108))
        d.line((pool.x2 + 6, 16), (pool.x2 + 6, 190))
        d.arrow((pool.x2, 108), (pool.x2 + 18, 108))
        x = pool.x2 + 18
    return d


def cbam_modules():
    """The two halves of CBAM. Channel: pool, shared MLP, add, sigmoid. Spatial: pool across channels, convolve, sigmoid."""
    d = Drawing(740, 400)
    d.group(20, 10, 700, 176, fill="tint")
    d.text(600, 30, "Channel Attention Module", size=14, bold=True)
    src = d.cuboid(52, 78, 62, 62, d=18)
    d.text(92, 168, "Input feature F", size=12, bold=True)
    mx = d.cuboid(210, 56, 56, 12, d=6, fill="solid")
    av = d.cuboid(210, 126, 56, 12, d=6, fill="rustsolid")
    d.text(241, 36, "MaxPool", size=12)
    d.text(241, 108, "AvgPool", size=12)
    d.arrow((132, 92), (170, 92), (176, 62), (204, 62))
    d.arrow((132, 112), (170, 112), (176, 132), (204, 132))
    mlp = d.net(322, 58, 46, 80, layers=(5, 3, 5))
    d.text(345, 158, "Shared MLP", size=12)
    d.arrow((272, 62), (300, 62), (306, 86), (316, 86))
    d.arrow((272, 132), (300, 132), (306, 110), (316, 110))
    o1 = d.cuboid(430, 56, 56, 12, d=6, fill="mid")
    o2 = d.cuboid(430, 126, 56, 12, d=6, fill="rustmid")
    d.arrow((374, 86), (400, 86), (406, 62), (424, 62))
    d.arrow((374, 110), (400, 110), (406, 132), (424, 132))
    add = d.op(530, 96, "+", r=13)
    sig = d.op(566, 96, "s", r=13)
    d.arrow((492, 60), (530, 60), (530, add.y))
    d.arrow((492, 130), (530, 130), (530, add.y2))
    res = d.cuboid(624, 90, 56, 12, d=6, fill="blue")
    d.arrow((sig.x2, 96), (res.x - 2, 96))
    d.text(645, 130, "Channel\nAttention Mc", size=12, bold=True)

    d.group(100, 204, 540, 186, fill="rust")
    d.text(506, 224, "Spatial Attention Module", size=14, bold=True)
    f2 = d.cuboid(140, 262, 58, 58, d=18)
    d.text(172, 350, "Channel-refined\nfeature F'", size=12, bold=True)
    p1 = d.cuboid(276, 262, 8, 58, d=18, fill="solid")
    p2 = d.cuboid(286, 262, 8, 58, d=18, fill="rustsolid")
    d.text(310, 350, "[MaxPool,\nAvgPool]", size=12)
    d.arrow((218, 284), (268, 284), sw=2)
    conv = d.cuboid(378, 258, 6, 62, d=18, fill="white")
    d.text(346, 238, "conv layer", size=12)
    d.line((300, 276), (378, 282))
    d.line((300, 300), (378, 290))
    d.box(288, 274, 14, 28, fill="none")
    d.arrow((406, 284), (428, 284), sw=2)
    s2 = d.op(446, 284, "s", r=13, fill="white")
    out = d.cuboid(508, 258, 6, 62, d=18, fill="rustmid")
    d.arrow((s2.x2, 284), (502, 284), sw=2)
    d.text(524, 350, "Spatial Attention\nMs", size=12, bold=True)
    return d


def cbam_sequence():
    """CBAM inside a ResBlock: the conv output F is refined by channel attention, then by spatial attention."""
    d = Drawing(740, 190)
    d.text(60, 92, "Previous\nconv blocks")
    d.line((112, 20), (112, 150), dash=True)
    a = d.cuboid(146, 82, 22, 26, d=9, fill="dark")
    f = d.cuboid(236, 82, 22, 26, d=9, fill="mid")
    d.text(256, 60, "F", italic=True, bold=True)
    d.arrow((118, 96), (a.x, 96))
    d.arrow((a.x2, 96), (f.x, 96), label="conv")
    ch = d.box(296, 44, 52, 34, fill="tint", edge="blue", sw=2)
    d.cuboid(306, 58, 28, 6, d=5, fill="mid")
    d.text(322, 70, "Mc", size=10, bold=True)
    d.text(322, 30, "Channel attention", size=12)
    m1 = d.op(384, 96, "x", r=10)
    d.text(388, 74, "F'", italic=True, bold=True)
    sp = d.box(418, 24, 34, 52, fill="rust", edge="rust", sw=2)
    d.cuboid(430, 38, 4, 22, d=8, fill="rustmid")
    d.text(435, 68, "Ms", size=10, bold=True)
    d.text(436, 12, "Spatial attention", size=12)
    m2 = d.op(490, 96, "x", r=10)
    d.text(494, 74, 'F"', italic=True, bold=True)
    add = d.op(540, 96, "+", r=10)
    out = d.cuboid(578, 82, 22, 26, d=9, fill="dark")
    d.arrow((f.x2, 96), (m1.x, 96))
    d.arrow((m1.x2, 96), (m2.x, 96))
    d.arrow((m2.x2, 96), (add.x, 96))
    d.arrow((add.x2, 96), (out.x, 96))
    d.arrow((out.x2, 96), (636, 96))
    d.arrow((262, 88), (296, 66))
    d.arrow((348, 66), (378, 88))
    d.arrow((392, 88), (418, 60))
    d.arrow((452, 60), (484, 88))
    d.arrow((160, 110), (350, 142), (540, add.y2), curve=True)
    d.line((648, 20), (648, 150), dash=True)
    d.text(694, 92, "Next\nconv blocks")
    d.text(380, 168, "ResBlock + CBAM", size=14)
    return d


def cbam_in_resnet():
    """Where CBAM goes in ResNet: after every pair of convolutions inside a stage, and again at the bottleneck."""
    d = Drawing(740, 176)
    img = d.box(2, 76, 48, 48, "Input\nImage", fill="grey", size=11)
    x = 62
    d.arrow((img.x2, 100), (x, 100))
    for i in range(2):
        g = _residual_stage(d, x, "Residual Block Layer %d" % (i + 1), True)
        d.line((g.x2 + 6, 14), (g.x2 + 6, 166))
        bam = d.box(g.x2 + 14, 80, 44, 40, "CBAM", fill="rust", size=11)
        d.text(bam.cx + 14, 30, "Bottleneck\nLayer of ResNet", size=10)
        pool = d.cuboid(bam.x2 + 8, 84, 7, 34, d=9, fill="rustmid")
        d.text(pool.cx, 62, "Pooling", size=10)
        d.arrow((g.x2, 100), (bam.x, 100))
        d.arrow((bam.x2, 100), (pool.x, 100))
        d.line((pool.x2 + 6, 14), (pool.x2 + 6, 166))
        x = pool.x2 + 16
        if i == 0:
            d.arrow((pool.x2, 100), (x, 100))
    return d
