"""Diagrams for "Simple, Powerful, and Fast - RegNet Architecture"."""
from kit import Drawing


def resnet_baseline():
    """The generic ResNet skeleton RegNet searches over: a stem, four layers, and a head."""
    d = Drawing(740, 190)
    d.raw('<path d="M12 62l14 12v44l-14 -12z" fill="#DCE7FA" stroke="#0F52BA" stroke-width="1.2"/>')
    d.text(22, 40, "Image\n3xHxW", size=11)
    d.group(52, 16, 150, 160, fill="grey", r=2)
    d.text(127, 30, "Stem", size=12)
    stem = d.cuboid(82, 68, 70, 38, d=14, fill="rust")
    d.text(127, 140, "3x3 Conv + BN + ReLU\nOutput Filters = 32\nStride = 2, Padding = Same", size=9.5)
    d.arrow((28, 90), (stem.x, 90))
    prev = stem.x2
    for i in range(4):
        b = d.box(218 + i * 92, 70, 76, 40, "Layer %d" % (i + 1), fill="blue", r=6, size=12)
        d.arrow((prev, 90), b.l)
        prev = b.x2
    d.group(590, 16, 146, 160, fill="grey", r=2)
    d.text(663, 30, "Head", size=12)
    pool = d.cuboid(604, 66, 8, 44, d=12, fill="mid")
    d.text(618, 136, "AveragePool2D\nLayer", size=9.5)
    net = d.net(664, 62, 44, 56, layers=(4, 3))
    d.text(690, 48, "Fully Connected", size=9.5)
    d.arrow((prev, 90), (pool.x, 90))
    d.arrow((pool.x2, 90), (net.x, 90))
    return d


def layer_block():
    """A layer is a chain of d residual blocks; the first one halves the resolution."""
    d = Drawing(740, 230)
    src = d.cuboid(30, 78, 50, 50, d=18, fill="dark")
    d.text(64, 154, "Input Feature Map\n(W1, R, R)", size=11)
    d.group(170, 12, 400, 206, fill="tint")
    d.text(370, 30, "Depth of a Layer = d", bold=True)
    d.text(370, 202, "Layer 2", bold=True)
    xs, names = (206, 306, 466), ("Residual Block 1", "Residual Block 2", "Residual Block d")
    prev = src.x2
    for x, name in zip(xs, names):
        b = d.cuboid(x, 84, 40, 40, d=16, fill="mid")
        d.text(x + 28, 54 if name[-1] != "2" else 152, name, size=11)
        if name[-1] == "d":
            d.dots(410, 104, gap=8)
            d.arrow((428, 104), (b.x, 104))
        else:
            d.arrow((prev, 104), (b.x, 104))
        prev = b.x2
    d.arrow((362, 104), (388, 104))
    out = d.cuboid(630, 78, 50, 50, d=18, fill="grey")
    d.arrow((prev, 104), (out.x, 104))
    d.text(664, 154, "Output Feature Map\n(W2, R/2, R/2)", size=11)
    return d


def se_module():
    """Squeeze-and-Excitation: pool each channel to one number, pass it through two 1x1 convolutions and a sigmoid, and rescale the channels."""
    d = Drawing(740, 220)
    src = d.cuboid(26, 66, 10, 46, d=14, fill="rustsolid")
    d.text(40, 134, "Dimensions\n(X, H, W)", size=10.5)
    d.group(76, 10, 420, 164, fill="rust")
    d.text(286, 26, "Squeeze-Excitation Block", bold=True)
    pool = d.cuboid(100, 66, 8, 46, d=12, fill="mid")
    d.text(112, 46, "AveragePool 2D", size=10.5)
    c1 = d.cuboid(186, 68, 60, 42, d=14, fill="white")
    d.text(224, 46, "SE Ratio q", size=10.5)
    d.text(224, 138, "1x1 Conv + ReLU\nIn Filters: X\nOut Filters: X / q", size=10)
    c2 = d.cuboid(322, 68, 60, 42, d=14, fill="white")
    d.text(360, 138, "1x1 Conv + ReLU\nIn Filters: X / q\nOut Filters: X", size=10)
    sig = d.op(452, 90, "s", r=13)
    d.text(452, 64, "Sigmoid", size=10.5)
    scale = d.cuboid(540, 66, 10, 46, d=14, fill="rustsolid")
    d.text(554, 134, "Dimensions\n(X, H, W)", size=10.5)
    mul = d.op(640, 90, "x", r=13)
    d.text(640, 56, "Elementwise\nMultiplication", size=10.5)
    d.arrow((src.x2, 90), (pool.x, 90))
    d.arrow((pool.x2, 90), (c1.x, 90))
    d.arrow((c1.x2, 90), (c2.x, 90))
    d.arrow((c2.x2, 90), (sig.x, 90))
    d.arrow((sig.x2, 90), (scale.x, 90))
    d.arrow((scale.x2, 90), (mul.x, 90))
    d.arrow((mul.x2, 90), (720, 90))
    d.arrow((62, 90), (62, 200), (640, 200), (640, mul.y2))
    return d
