# pade_resnet_cifar10.py
# ---------------------------------------------------------------
# Reproduction of:
#   Keles, O. & Tekalp, A. M. "PAON: A New Neuron Model using
#   Pade Approximants", IEEE ICIP 2024.
#   "Pade Neurons for Efficient Neural Models", 2026 (expanded).
#
# Specifically: their CIFAR-10 experiment (Table VIII).
#   - Baseline : ResNet(3,3,3)        -> 20 layers, ReLU
#   - Paon-S   : PadeResNet(2,2,2)   -> 14 layers, Paon-S[1/1], no ReLU
#   - Paon-S-II: PadeResNet-II(2,2,2) -> + element-wise (deformable)
#                                       1x1 offset Shifter
#
#   Both use AdamW(lr=1e-3, wd=5e-4), cosine annealing -> 2e-6,
#   600 epochs, batch 250. Augmentations identical to the paper.
# ---------------------------------------------------------------

import argparse, os, time, math, random
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, Subset
import torchvision
from torchvision import datasets, transforms

# ----------------------------------------------------------------
# Reproducibility
# ----------------------------------------------------------------
def seed_all(s=0):
    random.seed(s); np.random.seed(s)
    torch.manual_seed(s); torch.cuda.manual_seed_all(s)

# ----------------------------------------------------------------
# Smoothed Pade Neuron  (Paon-S [K/L])
#   Paon-S_{[K/L]} = (Q_L*P_K + Q_{L-1}*P_{K-1}) /
#                    (Q_L^2 + Q_{L-1}^2)
#   For [K/L] = [1/1] which the paper uses:
#       P_1 = a0 + a1 * x        (a0 = bias, a1 = conv weight)
#       Q_1 = 1 + b1 * x
#       P_0 = a0
#       Q_0 = 1
# ----------------------------------------------------------------
def pades(x, a0, aks, bks, K, L):
    """Evaluate Eq.(4) of the paper for arbitrary (K, L).

    Args
    ----
    x  : input tensor (N,*). Powers are taken elementwise.
    a0 : bias for the numerator (N,*).
    aks: list of K learned kernels (one per numerator power).
    bks: list of L learned kernels (one per denominator power).
    K,L: polynomial orders.
    """
    # P_K = a0 + a1*x + a2*x^2 + ... + a_K*x^K
    # Q_L = 1 + b1*x + ... + b_L*x^L
    xk = x
    PK = a0
    for k in range(1, K + 1):
        PK = PK + aks[k - 1] * xk
        xk = xk * x
    xb = x
    QL = torch.ones_like(x)
    for k in range(1, L + 1):
        QL = QL + bks[k - 1] * xb
        xb = xb * x

    # P_{K-1} and Q_{L-1} (define empty as 1 when K=0 or L=0)
    K1, L1 = max(K - 1, 0), max(L - 1, 0)
    xk = x
    PK1 = a0 if K1 >= 1 else a0 * 1.0
    for k in range(1, K1 + 1):
        PK1 = PK1 + aks[k - 1] * xk if k <= K else PK1
        xk = xk * x
    PK1 = a0 if K == 0 else PK1

    xb = x
    QL1 = torch.ones_like(x) if L1 == 0 else QL  # Q_0 = 1
    for k in range(1, L1 + 1):
        QL1 = QL1 + (bks[k - 1] * xb if k <= L else bks[L - 1] * xb)
        xb = xb * x

    num = QL * PK + QL1 * PK1
    den = QL * QL + QL1 * QL1
    return num / (den + 1e-8)            # tiny eps for safety


# ----------------------------------------------------------------
# Shifter Module  (two variants from the paper)
# ----------------------------------------------------------------
class ShifterKernelWise(nn.Module):
    """First Shifter variant: kernel-wise fixed-bounded shift.

    b<0  -> deactivate,
    b>=0 -> tanh-bounded continuous shift in [-b, b],
            re-sampled per channel via a 1x1 conv (init = 0).
    """
    def __init__(self, channels, b=0):
        super().__init__()
        self.b = b
        if b >= 0:
            self.offset = nn.Conv2d(channels, channels, 1, bias=True)
            nn.init.zeros_(self.offset.weight)
            nn.init.zeros_(self.offset.bias)

    def forward(self, x):
        if self.b < 0 or not hasattr(self, "offset"):
            return x
        B, C, H, W = x.shape
        m = max(H, W) // 4 if self.b == 0 else self.b
        sh = m * torch.tanh(self.offset(x))           # B,C,H,W
        # bilinear-shift via grid_sample
        sh_h = sh / max(W - 1, 1) * 2.0                # to [-2,2]
        sh_w = sh.permute(0, 1, 3, 2) / max(H - 1, 1) * 2.0
        # simpler: do two axis shifts separately
        # here we approximate the 2D shift via a single axis version
        # to keep shape stable; for the canonical CIFAR-10 setup
        # kernel-wise shift is the default Shifter-I choice.
        grid_y, grid_x = torch.meshgrid(
            torch.linspace(-1, 1, H, device=x.device),
            torch.linspace(-1, 1, W, device=x.device),
            indexing="ij",
        )
        grid = torch.stack((grid_x + sh_w.squeeze(1),
                            grid_y + sh_h.squeeze(1)), dim=-1)
        grid = grid.unsqueeze(1).expand(-1, C, -1, -1, -1) \
                 .reshape(B * C, H, W, 2)
        x_rep = x.reshape(B * C, 1, H, W)
        out = F.grid_sample(x_rep, grid,
                            mode="bilinear",
                            padding_mode="border",
                            align_corners=True)
        return out.reshape(B, C, H, W)


class ShifterElementWise(nn.Module):
    """Second Shifter variant (used in PadéResNet-II):
       element-wise 1x1 deformable offsets, bounded by m = max(h,w)/4."""
    def __init__(self, channels, b=0):
        super().__init__()
        self.b = b
        self.offset = nn.Conv2d(channels, 2 * channels, 1, bias=True)
        nn.init.zeros_(self.offset.weight)
        nn.init.zeros_(self.offset.bias)

    def forward(self, x):
        B, C, H, W = x.shape
        m = max(H, W) // 4 if self.b <= 0 else self.b
        sh = self.offset(x) * m                       # B,2C,H,W
        sh_y, sh_x = sh.chunk(2, dim=1)              # each B,C,H,W
        grid_y, grid_x = torch.meshgrid(
            torch.linspace(-1, 1, H, device=x.device),
            torch.linspace(-1, 1, W, device=x.device),
            indexing="ij",
        )
        grid = torch.stack((grid_x + sh_x / max(W - 1, 1) * 2,
                            grid_y + sh_y / max(H - 1, 1) * 2), dim=-1)
        grid = grid.unsqueeze(1).expand(-1, C, -1, -1, -1) \
                 .reshape(B * C, H, W, 2)
        x_rep = x.reshape(B * C, 1, H, W)
        out = F.grid_sample(x_rep, grid,
                            mode="bilinear",
                            padding_mode="border",
                            align_corners=True)
        return out.reshape(B, C, H, W)


# ----------------------------------------------------------------
# Pade Conv2d Layer  (PaLa in the paper)
# ----------------------------------------------------------------
class PaLaConv2d(nn.Module):
    """Convolutional layer built from Paon-S [K/L] neurons.

    Replaces (Conv -> BN -> ReLU).  BN is *kept* but the
    post-addition ReLU is removed (per the paper).
    """
    def __init__(self, in_ch, out_ch, kernel_size=3,
                 K=1, L=1, stride=1, padding=None,
                 shifter=None, shifter_b=0,
                 use_bn=True):
        super().__init__()
        if padding is None:
            padding = kernel_size // 2
        self.K, self.L = K, L
        self.in_ch, self.out_ch, self.stride = in_ch, out_ch, stride
        self.use_bn = use_bn

        # One conv per numerator power k=1..K and per denom power l=1..L
        self.num_kernels = nn.ParameterList([
            nn.Parameter(torch.empty(out_ch, in_ch, kernel_size, kernel_size))
            for _ in range(K)
        ])
        self.den_kernels = nn.ParameterList([
            nn.Parameter(torch.empty(out_ch, in_ch, kernel_size, kernel_size))
            for _ in range(L)
        ])
        self.bias = nn.Parameter(torch.zeros(out_ch))

        # init as the paper: small identity-ish for stability
        for p in self.num_kernels:
            nn.init.kaiming_normal_(p, mode="fan_out", nonlinearity="relu")
        for p in self.den_kernels:
            nn.init.zeros_(p)        # start with b_l = 0  -> Q_L = 1

        self.shift = shifter(in_ch, b=shifter_b) if shifter is not None else None
        if use_bn:
            self.bn = nn.BatchNorm2d(out_ch)

    def forward(self, x):
        if self.shift is not None:
            x = self.shift(x)

        # Compute element-wise powers of x AFTER padding
        N, C, H, W = x.shape
        pad = self.num_kernels[0].shape[-1] // 2
        x_p = F.pad(x, (pad, pad, pad, pad), mode="replicate")

        # numerator and denominator contributions
        xb = x
        num = self.bias.view(1, -1, 1, 1)           # broadcast bias
        xk = x
        for k in range(self.K):
            num = num + F.conv2d(xk, self.num_kernels[k],
                                 stride=self.stride)
            xk = xk * x

        den = torch.ones_like(num)
        xb2 = x
        for k in range(self.L):
            den = den + F.conv2d(xb2, self.den_kernels[k],
                                 stride=self.stride)
            xb2 = xb2 * x

        # one-degree-lower polynomials
        if self.K >= 1:
            xk1 = x
            num1 = self.bias.view(1, -1, 1, 1)
            for k in range(self.K - 1):
                num1 = num1 + F.conv2d(xk1,
                                       self.num_kernels[k],
                                       stride=self.stride)
                xk1 = xk1 * x
        else:
            num1 = self.bias.view(1, -1, 1, 1)

        if self.L >= 1:
            xb1 = x
            den1 = torch.ones_like(den)
            for k in range(self.L - 1):
                den1 = den1 + F.conv2d(xb1,
                                       self.den_kernels[k],
                                       stride=self.stride)
                xb1 = xb1 * x
        else:
            den1 = torch.ones_like(den)

        out = (den * num + den1 * num1) / (den * den + den1 * den1 + 1e-8)
        if self.use_bn:
            out = self.bn(out)
        return out


# ----------------------------------------------------------------
# Pade Linear Layer  (for the classifier head)
# ----------------------------------------------------------------
class PaLaLinear(nn.Module):
    """Fully-connected Paon-S layer, [K/L] = [1/1] (paper default)."""
    def __init__(self, in_f, out_f, K=1, L=1):
        super().__init__()
        self.K, self.L = K, L
        self.w_n = nn.ParameterList([
            nn.Parameter(torch.empty(out_f, in_f)) for _ in range(K)
        ])
        self.w_d = nn.ParameterList([
            nn.Parameter(torch.empty(out_f, in_f)) for _ in range(L)
        ])
        self.b = nn.Parameter(torch.zeros(out_f))
        for p in self.w_n:
            nn.init.kaiming_normal_(p, nonlinearity="relu")
        for p in self.w_d:
            nn.init.zeros_(p)

    def forward(self, x):
        # P_1 = b + w1*x ; Q_1 = 1 + v1*x
        P1 = F.linear(x, self.w_n[0], self.b)
        Q1 = 1.0 + F.linear(x, self.w_d[0])
        P0 = self.b.expand_as(P1)
        Q0 = torch.ones_like(Q1)
        num = Q1 * P1 + Q0 * P0
        den = Q1 * Q1 + Q0 * Q0
        return num / (den + 1e-8)


# ----------------------------------------------------------------
# Residual Blocks
# ----------------------------------------------------------------
class BasicBlock(nn.Module):
    """Vanilla Conv-BN-ReLU block (baseline)."""
    expansion = 1
    def __init__(self, in_ch, out_ch, stride=1):
        super().__init__()
        self.conv1 = nn.Conv2d(in_ch, out_ch, 3, stride, 1, bias=False)
        self.bn1   = nn.BatchNorm2d(out_ch)
        self.conv2 = nn.Conv2d(out_ch, out_ch, 3, 1, 1, bias=False)
        self.bn2   = nn.BatchNorm2d(out_ch)
        self.short = (nn.Sequential(
            nn.Conv2d(in_ch, out_ch, 1, stride, bias=False),
            nn.BatchNorm2d(out_ch)
        ) if stride != 1 or in_ch != out_ch else nn.Identity())

    def forward(self, x):
        out = F.relu(self.bn1(self.conv1(x)))
        out = self.bn2(self.conv2(out))
        out = out + self.short(x)
        return F.relu(out)


class PadeBasicBlock(nn.Module):
    """Padé-S block: two PaLaConv2d + residual + identity (no ReLU).

    shifter_choice: "none", "kw" (kernel-wise), "ew" (element-wise).
    """
    def __init__(self, in_ch, out_ch, stride=1, shifter="none", K=1, L=1):
        super().__init__()
        shifter_cls = None
        if shifter == "kw":
            shifter_cls = ShifterKernelWise
        elif shifter == "ew":
            shifter_cls = ShifterElementWise

        self.conv1 = PaLaConv2d(in_ch, out_ch, 3, K=K, L=L, stride=stride,
                                shifter=shifter_cls,
                                shifter_b=-1 if shifter == "none" else 1)
        self.conv2 = PaLaConv2d(out_ch, out_ch, 3, K=K, L=L, stride=1,
                                shifter=shifter_cls,
                                shifter_b=-1 if shifter == "none" else 1)
        self.short = (PaLaConv2d(in_ch, out_ch, 1, K=K, L=L, stride=stride,
                                 shifter=None, use_bn=True)
                      if stride != 1 or in_ch != out_ch else nn.Identity())

    def forward(self, x):
        out = self.conv1(x)
        out = self.conv2(out)
        return out + self.short(x)            # no ReLU after add


# ----------------------------------------------------------------
# Networks
# ----------------------------------------------------------------
class ResNetCIFAR(nn.Module):
    """Vanilla ResNet20-style network with (n1,n2,n3) blocks per stage.

    Stages: 3, channels: [16, 32, 64], imagenet-style stem removed
    (32x32 kept) per the paper.
    """
    def __init__(self, blocks_per_stage=(3, 3, 3), num_classes=10):
        super().__init__()
        self.in_ch = 16
        self.stem  = nn.Sequential(
            nn.Conv2d(3, 16, 3, 1, 1, bias=False),
            nn.BatchNorm2d(16),
        )
        self.layer1 = self._make_layer(16, blocks_per_stage[0], stride=1,
                                       block_cls=BasicBlock)
        self.layer2 = self._make_layer(32, blocks_per_stage[1], stride=2,
                                       block_cls=BasicBlock)
        self.layer3 = self._make_layer(64, blocks_per_stage[2], stride=2,
                                       block_cls=BasicBlock)
        self.pool   = nn.AdaptiveAvgPool2d(1)
        self.fc     = nn.Linear(64, num_classes)

    def _make_layer(self, out_ch, n_blocks, stride, block_cls):
        layers = [block_cls(self.in_ch, out_ch, stride)]
        self.in_ch = out_ch
        for _ in range(1, n_blocks):
            layers.append(block_cls(self.in_ch, out_ch, stride=1))
        return nn.Sequential(*layers)

    def forward(self, x):
        x = F.relu(self.stem(x))
        x = self.layer1(x); x = self.layer2(x); x = self.layer3(x)
        x = self.pool(x).flatten(1)
        return self.fc(x)


class PadeResNetCIFAR(nn.Module):
    """Padé-S version of ResNetCIFAR: (PaLaConv2d replaces Conv+ReLU)."""
    def __init__(self, blocks_per_stage=(2, 2, 2), num_classes=10,
                 shifter="none", head_pade=True):
        super().__init__()
        self.in_ch = 16
        self.shifter = shifter
        # stem uses a Padé conv too (degree [1/0] = ordinary conv style)
        self.stem = PaLaConv2d(3, 16, 3, K=1, L=0, stride=1, use_bn=True)

        self.layer1 = self._make_layer(16, blocks_per_stage[0], stride=1)
        self.layer2 = self._make_layer(32, blocks_per_stage[1], stride=2)
        self.layer3 = self._make_layer(64, blocks_per_stage[2], stride=2)
        self.pool   = nn.AdaptiveAvgPool2d(1)
        self.fc = PaLaLinear(64, num_classes) if head_pade \
                  else nn.Linear(64, num_classes)

    def _make_layer(self, out_ch, n_blocks, stride):
        layers = [PadeBasicBlock(self.in_ch, out_ch, stride,
                                 shifter=self.shifter)]
        self.in_ch = out_ch
        for _ in range(1, n_blocks):
            layers.append(PadeBasicBlock(self.in_ch, out_ch, stride=1,
                                         shifter=self.shifter))
        return nn.Sequential(*layers)

    def forward(self, x):
        x = self.stem(x)               # no ReLU after stem
        x = self.layer1(x); x = self.layer2(x); x = self.layer3(x)
        x = self.pool(x).flatten(1)
        return self.fc(x)


# ----------------------------------------------------------------
# Data pipeline (matches the paper)
# ----------------------------------------------------------------
class AddGaussianNoise:
    """40 dB SNR Gaussian noise injected during training only."""
    def __init__(self, snr_db=40.0):
        self.snr_db = snr_db

    def __call__(self, x):
        if not torch.is_tensor(x):
            return x
        signal_pwr = x.pow(2).mean()
        noise_pwr  = signal_pwr / (10 ** (self.snr_db / 10))
        noise      = torch.randn_like(x) * noise_pwr.sqrt()
        return x + noise


def make_transforms():
    train_t = transforms.Compose([
        transforms.RandomCrop(32, padding=4),
        transforms.RandomHorizontalFlip(),
        transforms.RandomVerticalFlip(),
        transforms.RandomRotation(90),
        # channel-shuffle: torchvision has no direct op, do it in __getitem__
        transforms.ToTensor(),
        transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
        AddGaussianNoise(40.0),
    ])
    test_t = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
    ])
    return train_t, test_t


class ChannelShuffleWrapper:
    """Wraps a Dataset to apply random channel-permutation after the transform."""
    def __init__(self, ds): self.ds = ds
    def __len__(self): return len(self.ds)
    def __getitem__(self, i):
        x, y = self.ds[i]
        if torch.rand(1).item() < 0.5:
            perm = torch.randperm(3)
            x = x[perm]
        return x, y


# ----------------------------------------------------------------
# Train / evaluate
# ----------------------------------------------------------------
def evaluate(model, loader, device):
    model.eval()
    correct, total = 0, 0
    with torch.no_grad():
        for x, y in loader:
            x, y = x.to(device, non_blocking=True), y.to(device, non_blocking=True)
            correct += (model(x).argmax(1) == y).sum().item()
            total += y.size(0)
    return 100.0 * correct / total


def train(model, train_loader, val_loader, args, device, name=""):
    optim = torch.optim.AdamW(model.parameters(),
                              lr=args.lr, weight_decay=args.wd)
    sched = torch.optim.lr_scheduler.CosineAnnealingLR(
        optim, T_max=args.epochs, eta_min=args.lr_min)
    best_val = 0.0
    step = 0
    log_every = max(1, len(train_loader) // 5)
    t0 = time.time()
    for epoch in range(args.epochs):
        model.train()
        running = 0.0
        for i, (x, y) in enumerate(train_loader):
            x = x.to(device, non_blocking=True)
            y = y.to(device, non_blocking=True)
            logits = model(x)
            loss = F.cross_entropy(logits, y)
            optim.zero_grad(set_to_none=True)
            loss.backward()
            optim.step()
            running += loss.item()
            step += 1
            if (i + 1) % log_every == 0:
                print(f"[{name}] ep {epoch+1}/{args.epochs} "
                      f"step {step} loss {running/log_every:.4f}")
                running = 0.0
        sched.step()
        acc = evaluate(model, val_loader, device)
        best_val = max(best_val, acc)
        print(f"==> [{name}] epoch {epoch+1} finished, val acc "
              f"{acc:.2f}% (best {best_val:.2f}%)")
    test_acc = evaluate(model, val_loader, device)   # reuse split
    print(f"\n{name}: best val acc = {best_val:.2f}%, "
          f"final val acc = {test_acc:.2f}%, "
          f"time = {time.time()-t0:.0f}s\n")
    return best_val


def make_loaders(args):
    train_t, test_t = make_transforms()
    full_train = datasets.CIFAR10(
        args.data_root, train=True,  download=True, transform=train_t)
    test       = datasets.CIFAR10(
        args.data_root, train=False, download=True, transform=test_t)

    # 5000-image val split (paper); used both as val and as test set here
    idx       = torch.randperm(len(full_train), generator=torch.Generator().manual_seed(0))
    val_idx   = idx[:5000]
    train_idx = idx[5000:]
    val_ds  = Subset(datasets.CIFAR10(args.data_root, train=True,
                                       download=True, transform=test_t),
                     val_idx.tolist())
    train_ds = ChannelShuffleWrapper(Subset(full_train, train_idx.tolist()))

    return (DataLoader(train_ds, batch_size=args.bs, shuffle=True,
                       num_workers=args.workers, pin_memory=True,
                       persistent_workers=args.workers > 0),
            DataLoader(val_ds, batch_size=args.bs, shuffle=False,
                       num_workers=args.workers, pin_memory=True,
                       persistent_workers=args.workers > 0))


def build_model(mode, args):
    if mode == "baseline":
        return ResNetCIFAR(blocks_per_stage=(3, 3, 3), num_classes=10)
    elif mode == "pade":
        return PadeResNetCIFAR(blocks_per_stage=tuple(args.blocks),
                               shifter="none")
    elif mode == "pade_ii":
        return PadeResNetCIFAR(blocks_per_stage=tuple(args.blocks),
                               shifter="ew")
    else:
        raise ValueError(mode)


# ----------------------------------------------------------------
# Main
# ----------------------------------------------------------------
def main():
    p = argparse.ArgumentParser()
    p.add_argument("--mode", default="both",
                   choices=["baseline", "pade", "pade_ii", "both"])
    p.add_argument("--data_root", default="../data")
    p.add_argument("--bs",        type=int, default=250)
    p.add_argument("--epochs",    type=int, default=600)
    p.add_argument("--lr",        type=float, default=1e-3)
    p.add_argument("--lr_min",    type=float, default=2e-6)
    p.add_argument("--wd",        type=float, default=5e-4)
    p.add_argument("--workers",   type=int, default=4)
    p.add_argument("--blocks", nargs=3, type=int,
                   default=[2, 2, 2], help="blocks per stage for Paon-S")
    p.add_argument("--device",    default="cuda" if torch.cuda.is_available() else "cpu")
    p.add_argument("--seed",      type=int, default=0)
    args = p.parse_args()

    seed_all(args.seed)
    train_loader, val_loader = make_loaders(args)

    results = {}
    modes = ["baseline", "pade", "pade_ii"] if args.mode == "both" else [args.mode]
    for m in modes:
        model = build_model(m, args).to(args.device)
        nparams = sum(p.numel() for p in model.parameters())
        print(f"\n=== Training {m}: {nparams:,} params on {args.device} ===")
        best = train(model, train_loader, val_loader, args, args.device, name=m)
        results[m] = (nparams, best)

    if len(results) > 1:
        print("\n===== Final comparison =====")
        for m, (np_, acc) in results.items():
            print(f"  {m:10s}  params={np_:>9,}   best val acc = {acc:5.2f}%")
        b_p = results["baseline"][1]
        p_p = results["pade"][1]
        if b_p and p_p:
            print(f"\n  PadeResNet(2,2,2) vs ResNet(3,3,3) on CIFAR-10: "
                  f"{p_p - b_p:+.2f} pp (paper reports +0.37 pp)")


if __name__ == "__main__":
    main()
