"""Small U-Net family used for all learned variants (same backbone, different inputs)."""
from __future__ import annotations

import torch
import torch.nn as nn
import torch.nn.functional as F


def _block(cin, cout):
    return nn.Sequential(
        nn.Conv2d(cin, cout, 3, padding=1), nn.BatchNorm2d(cout), nn.ReLU(inplace=True),
        nn.Conv2d(cout, cout, 3, padding=1), nn.BatchNorm2d(cout), nn.ReLU(inplace=True))


class UNet(nn.Module):
    def __init__(self, in_ch: int, base: int = 16, depth: int = 5, residual: bool = False, out_ch: int = 1, cap: float = 0.0, cap_down: float | None = None):
        """residual=True: channel 2 (the physics baseline) is added to the output, so the
        network only learns the correction (the 'hybrid' of the research plan).
        cap > 0 (with residual): the correction is cap * tanh(.), i.e. bounded to +/- cap in the training scale, so
        the output cannot leave the physics map by more than that (train --cap, in dB). cap_down: a different bound
        for corrections downwards (removing signal), e.g. large, while adding signal stays bounded by cap.
        out_ch=2: the second output is a 'covered' logit (signal above the floor), see train.unet_predict."""
        super().__init__()
        self.residual, self.cap = residual, cap
        self.cap_down = cap if cap_down is None else cap_down
        chs = [base * 2 ** i for i in range(depth)]
        self.down = nn.ModuleList()
        c = in_ch
        for ch in chs:
            self.down.append(_block(c, ch))
            c = ch
        self.up = nn.ModuleList()
        self.upconv = nn.ModuleList()
        for ch in reversed(chs[:-1]):
            self.upconv.append(nn.ConvTranspose2d(c, ch, 2, stride=2))
            self.up.append(_block(ch * 2, ch))
            c = ch
        self.head = nn.Conv2d(c, out_ch, 1)

    def forward(self, x):
        skips = []
        h = x
        for i, blk in enumerate(self.down):
            h = blk(h)
            if i < len(self.down) - 1:
                skips.append(h)
                h = F.max_pool2d(h, 2)
        for upc, blk in zip(self.upconv, self.up):
            h = upc(h)
            h = blk(torch.cat([h, skips.pop()], 1))
        out = self.head(h)
        if self.residual:
            t = torch.tanh(out[:, :1])
            corr = torch.where(t > 0, self.cap * t, self.cap_down * t) if self.cap > 0 else out[:, :1]
            out = torch.cat([corr + x[:, 2:3], out[:, 1:]], 1)
        return out


def _convrelu(cin, cout, k, pad, pool):
    return nn.Sequential(nn.Conv2d(cin, cout, k, padding=pad), nn.ReLU(inplace=True), nn.MaxPool2d(pool, stride=pool))


def _convreluT(cin, cout, k, pad):
    return nn.Sequential(nn.ConvTranspose2d(cin, cout, k, stride=2, padding=pad), nn.ReLU(inplace=True))


class RadioUNet(nn.Module):
    """First U-Net of RadioWNet (Levie et al., github.com/RonLevie/RadioUNet, lib/modules.py, MIT licence),
    ported layer for layer. Few channels at full resolution, 5x5 convs, no BatchNorm, ReLU output.
    residual=True adds channel 2 (the physics map) to the output, as in UNet."""

    def __init__(self, in_ch: int, residual: bool = False):
        super().__init__()
        self.residual = residual
        w = 6 if in_ch <= 3 else 10
        self.layer00 = _convrelu(in_ch, w, 3, 1, 1)
        self.layer0 = _convrelu(w, 40, 5, 2, 2)
        self.layer1 = _convrelu(40, 50, 5, 2, 2)
        self.layer10 = _convrelu(50, 60, 5, 2, 1)
        self.layer2 = _convrelu(60, 100, 5, 2, 2)
        self.layer20 = _convrelu(100, 100, 3, 1, 1)
        self.layer3 = _convrelu(100, 150, 5, 2, 2)
        self.layer4 = _convrelu(150, 300, 5, 2, 2)
        self.layer5 = _convrelu(300, 500, 5, 2, 2)
        self.conv_up5 = _convreluT(500, 300, 4, 1)
        self.conv_up4 = _convreluT(600, 150, 4, 1)
        self.conv_up3 = _convreluT(300, 100, 4, 1)
        self.conv_up20 = _convrelu(200, 100, 3, 1, 1)
        self.conv_up2 = _convreluT(200, 60, 6, 2)
        self.conv_up10 = _convrelu(120, 50, 5, 2, 1)
        self.conv_up1 = _convreluT(100, 40, 6, 2)
        self.conv_up0 = _convreluT(80, 20, 6, 2)
        self.conv_up00 = _convrelu(20 + w + in_ch, 20, 5, 2, 1)
        # the original ends in convrelu (ReLU on the output); kept, and dropped for the residual variant
        self.conv_up000 = _convrelu(20 + in_ch, 1, 5, 2, 1) if not residual else nn.Conv2d(20 + in_ch, 1, 5, padding=2)

    def forward(self, x):
        l00 = self.layer00(x)
        l0 = self.layer0(l00)
        l1 = self.layer1(l0)
        l10 = self.layer10(l1)
        l2 = self.layer2(l10)
        l20 = self.layer20(l2)
        l3 = self.layer3(l20)
        l4 = self.layer4(l3)
        l5 = self.layer5(l4)
        u = torch.cat([self.conv_up5(l5), l4], 1)
        u = torch.cat([self.conv_up4(u), l3], 1)
        u = torch.cat([self.conv_up3(u), l20], 1)
        u = torch.cat([self.conv_up20(u), l2], 1)
        u = torch.cat([self.conv_up2(u), l10], 1)
        u = torch.cat([self.conv_up10(u), l1], 1)
        u = torch.cat([self.conv_up1(u), l0], 1)
        u = torch.cat([self.conv_up0(u), l00, x], 1)
        u = torch.cat([self.conv_up00(u), x], 1)
        out = self.conv_up000(u)
        if self.residual:
            out = out + x[:, 2:3]
        return out


def build(arch: str, in_ch: int, base: int = 32, residual: bool = False, out_ch: int = 1, cap: float = 0.0, cap_down: float | None = None) -> nn.Module:
    if arch == "radiounet":
        assert out_ch == 1 and cap == 0, "the coverage gate and the capped residual are only implemented for the unet arch"
        return RadioUNet(in_ch, residual)
    return UNet(in_ch, base, residual=residual, out_ch=out_ch, cap=cap, cap_down=cap_down)
