File size: 2,518 Bytes
c01cda9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
import torch as pt
import torch.nn.functional as F
import numbers
def resize_down_up(x, scale=0.5, mode="bilinear"):
    if isinstance(scale, pt.Tensor):
        scale = scale.item()
    _val_invars(x, scale=scale, mode=mode, type="resize")
    B, C, H, W = x.shape
    h_new, w_new = max(1, int(H * scale)), max(1, int(W * scale))

    kwargs = {}
    if mode in {"bilinear", "bicubic"}:
        kwargs["align_corners"] = False

    x_down = F.interpolate(x, size=(h_new, w_new), mode=mode, **kwargs)
    return F.interpolate(x_down, size=(H,W), mode=mode, **kwargs)

def _val_invars(x, **kwargs):
    if x.ndim != 4:
        raise ValueError(f"x needs shape (B,C,H,W), got {x.shape}")
    op_type = kwargs.get("type")

    if op_type == "resize":
        scale = kwargs.get("scale")
        if not isinstance(scale, numbers.Real):
            raise ValueError(f"scale must be a real number, got {scale}")
        if scale <= 0.0 or scale >= 1.0:
            raise ValueError(f"scale must be in (0,1), got {scale}")
        mode = kwargs.get("mode")
        valid_modes = {"bilinear", "bicubic", "nearest", "area"}
        if mode not in valid_modes:
            raise ValueError(f"mode must be one of {valid_modes}, got {mode}")
    elif op_type == "decimate":
        factor = kwargs.get("factor")
        if factor < 2:
            raise ValueError(f"decimation factor must be >=2, got {factor}")
    elif op_type == "attack":
        epsilon = kwargs.get("epsilon")
        if not (0.0 < epsilon < 1.0):
            raise ValueError(f"epsilon must be (0,1), got {epsilon}")

def decimate(x, factor=2):
    # decimation: keep every factor-th pixel, then upsample back up.
    _val_invars(x, factor=factor, type="decimate")
    B, C, H, W = x.shape
    x_dec = x[:, :, ::factor, ::factor]
    return F.interpolate(x_dec, size=(H,W), mode="nearest")

def blur_decimate(x, blur):
    B, C, H, W = x.shape
    x_down = blur(x)
    return F.interpolate(x_down, size=(H,W), mode="bilinear", align_corners=False)

def checkerboard_alias_attack(
    x: pt.Tensor,
    epsilon: float=0.5,
):
    _val_invars(x, epsilon=epsilon, type="attack")
    B, C, H, W = x.shape
    device = x.device
    dtype = x.dtype

    rows = pt.arange(H, device=device).view(H,1)
    cols = pt.arange(W, device=device).view(1,W)
    checker = ((rows + cols) % 2).float() * 2.0 - 1.0 # values in {-1, +1}
    checker = checker.view( 1, 1, H, W).expand(B, C, H, W).to(dtype)

    x_adv = (x + epsilon * checker).clamp(-1.0, 1.0)
    return x_adv