"""Drawings that more than one read uses."""
from kit import Drawing


def dilation():
    """Three 3x3 kernels with dilation 1, 2 and 4, stacked: the area each output pixel sees grows 3 -> 7 -> 15."""
    d = Drawing(740, 284)
    size, cells, gap = 13, 17, 22
    seen = 1                                    # side of the area one tap already covers
    for i, rate in enumerate((1, 2, 4)):
        x, mid = 12 + i * (cells * size + gap), cells // 2
        taps = [(mid + a * rate, mid + b * rate) for a in (-1, 0, 1) for b in (-1, 0, 1)]
        count = {}
        for r, c in taps:
            for a in range(-(seen // 2), seen // 2 + 1):
                for b in range(-(seen // 2), seen // 2 + 1):
                    count[(r + a, c + b)] = count.get((r + a, c + b), 0) + 1
        top = max(count.values())
        shade = {k: 0.22 + 0.5 * v / top for k, v in count.items()}
        d.text(x + cells * size / 2, 14, "%d-dilated convolution" % rate, bold=True)
        d.cells(x, 30, cells, cells, size, shade=shade, dots=taps)
        seen += 2 * rate
        d.text(x + cells * size / 2, 30 + cells * size + 5 + 8, "sees %d x %d" % (seen, seen), size=12, color="mute")
    return d


def _residual(d, x, y, bottleneck=True, downsample=False, se=False, w=360):
    """One residual block inside a rounded panel at (x, y). Returns nothing; the panel is w wide and 176 tall."""
    kind = "Bottleneck" if bottleneck else "Simple Block"
    d.group(x + 64, y, w - 150, 176, fill="tint")
    d.text(x + 64 + (w - 150) / 2, y + 16, "%s %s Downsample Layer" % (kind, "with" if downsample else "w/o"), size=11, bold=True)
    my = y + 72
    d.cuboid(x + 24, my - 14, 8, 30, d=10, fill="rustsolid")
    d.text(x + 34, my + 34, "Input\n(W1, R, R)", size=9.5)
    names = ["1x1 Conv\n+BN+ReLU", "3x3 Conv\n+BN+ReLU", "1x1 Conv\n+BN+ReLU"] if bottleneck else ["3x3 Conv\n+BN+ReLU", "3x3 Conv\n+BN"]
    notes = ["bottleneck\nratio b", "group\nsize g", ""] if bottleneck else ["", ""]
    slots = len(names) + (1 if se else 0)
    span = w - 150 - 56
    step = span / slots
    bx = x + 64 + 22
    prev_x2 = x + 42
    order = list(zip(names, notes))
    if se:
        order.insert(len(order) - 1, ("SE", ""))
    for name, note in order:
        if name == "SE":
            b = d.cuboid(bx + 6, my - 13, 20, 28, d=9, fill="mid")
            d.text(bx + 20, my - 34, "SE Attention\nModule", size=9)
        else:
            b = d.cuboid(bx, my - 13, 34, 28, d=9, fill="rust")
            d.text(bx + 22, my + 32, name, size=9)
            if note:
                d.text(bx + 22, my - 36, note, size=9, italic=True, color="mute")
        d.arrow((prev_x2, my), (b.x - 1, my), sw=1.1)
        prev_x2 = b.x2
        bx += step
    add = d.op(x + w - 106, my, "+", r=8)
    d.text(add.cx, my - 20, "Add+ReLU", size=9)
    d.arrow((prev_x2, my), (add.x, my), sw=1.1)
    out = d.cuboid(x + w - 54, my - 14, 8, 30, d=10, fill="grey")
    d.text(x + w - 44, my + 34, "Output\n(W1, R/2, R/2)" if downsample else "Output\n(W1, R, R)", size=9.5)
    d.arrow((add.x2, my), (out.x - 1, my), sw=1.1)
    sx, sy = x + 64 + 10, y + 142
    if downsample:
        c = d.cuboid(x + w / 2 - 30, sy - 13, 34, 26, d=8, fill="rust")
        d.text(x + w / 2 - 8, sy + 24, "1x1 Conv+BN, Stride = 2", size=9)
        d.arrow((sx, my), (sx, sy), (c.x - 1, sy), sw=1.1)
        d.arrow((c.x2, sy), (add.cx, sy), (add.cx, add.y2), sw=1.1)
    else:
        d.arrow((sx, my), (sx, sy), (add.cx, sy), (add.cx, add.y2), sw=1.1)


def residual_blocks():
    """The four residual blocks: simple or bottleneck, with or without a downsampling shortcut."""
    d = Drawing(740, 380)
    _residual(d, 2, 6, bottleneck=False, downsample=False)
    _residual(d, 374, 6, bottleneck=True, downsample=False)
    _residual(d, 2, 198, bottleneck=False, downsample=True)
    _residual(d, 374, 198, bottleneck=True, downsample=True)
    d.line((370, 4), (370, 376), color="grey")
    d.line((4, 190), (736, 190), color="grey")
    return d


def bottleneck_blocks():
    """ResNet residual block design: the bottleneck block, without and with a downsampling shortcut."""
    d = Drawing(740, 380)
    d.text(370, 14, "ResNet Residual Block Design", bold=True, size=14)
    _residual(d, 130, 28, bottleneck=True, downsample=False, w=480)
    _residual(d, 130, 204, bottleneck=True, downsample=True, w=480)
    return d


def se_block():
    """A bottleneck residual block with a Squeeze-and-Excitation module before its last convolution."""
    d = Drawing(740, 190)
    _residual(d, 110, 6, bottleneck=True, downsample=False, se=True, w=520)
    return d
