File size: 6,454 Bytes
320e2b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
import torch
import torch.nn as nn

def _make_pruned_conv(old_conv: nn.Conv2d, keep_out_idx: torch.Tensor) -> nn.Conv2d:
    """
    Create a new Conv2d with fewer output channels (keep_out_idx),
    copying weights/bias from old_conv.
    """
    device = old_conv.weight.device
    dtype = old_conv.weight.dtype

    new_out = keep_out_idx.numel()
    new_conv = nn.Conv2d(
        in_channels=old_conv.in_channels,
        out_channels=new_out,
        kernel_size=old_conv.kernel_size,
        stride=old_conv.stride,
        padding=old_conv.padding,
        dilation=old_conv.dilation,
        groups=old_conv.groups,
        bias=(old_conv.bias is not None),
        padding_mode=old_conv.padding_mode,
    ).to(device=device, dtype=dtype)

    with torch.no_grad():
        new_conv.weight.copy_(old_conv.weight.data[keep_out_idx].contiguous())
        if old_conv.bias is not None:
            new_conv.bias.copy_(old_conv.bias.data[keep_out_idx].contiguous())

    return new_conv


def _make_pruned_bn(old_bn: nn.BatchNorm2d, keep_idx: torch.Tensor) -> nn.BatchNorm2d:
    """
    Create a new BatchNorm2d with fewer channels, copying params + running stats.
    """
    device = old_bn.weight.device
    dtype = old_bn.weight.dtype

    new_nf = keep_idx.numel()
    new_bn = nn.BatchNorm2d(
        num_features=new_nf,
        eps=old_bn.eps,
        momentum=old_bn.momentum,
        affine=old_bn.affine,
        track_running_stats=old_bn.track_running_stats,
    ).to(device=device, dtype=dtype)

    with torch.no_grad():
        if old_bn.affine:
            new_bn.weight.copy_(old_bn.weight.data[keep_idx].contiguous())
            new_bn.bias.copy_(old_bn.bias.data[keep_idx].contiguous())

        if old_bn.track_running_stats:
            new_bn.running_mean.copy_(old_bn.running_mean.data[keep_idx].contiguous())
            new_bn.running_var.copy_(old_bn.running_var.data[keep_idx].contiguous())
            new_bn.num_batches_tracked.copy_(old_bn.num_batches_tracked)

    return new_bn


def prune_conv_bn_pair(conv: nn.Conv2d, bn: nn.BatchNorm2d, amount: float = 0.3):
    """
    Structurally prune Conv2d output channels using L1 norm.
    Returns: (new_conv, new_bn, keep_out_idx)
    """
    if not (0.0 <= amount < 1.0):
        raise ValueError("amount must be in [0, 1).")

    W = conv.weight.data  # (out, in, kH, kW)
    out_ch = W.shape[0]
    num_prune = int(round(amount * out_ch))

    if num_prune <= 0:
        keep_idx = torch.arange(out_ch, device=W.device)
        return conv, bn, keep_idx

    # L1 norm per output channel
    channel_l1 = W.abs().sum(dim=(1, 2, 3))  # (out,)

    # Keep the highest-L1 channels
    sorted_idx = torch.argsort(channel_l1, descending=True)
    keep_idx = sorted_idx[num_prune:]  # (kept,)

    # Keep indices sorted for nicer determinism
    keep_idx, _ = torch.sort(keep_idx)

    new_conv = _make_pruned_conv(conv, keep_idx)
    new_bn = _make_pruned_bn(bn, keep_idx)

    return new_conv, new_bn, keep_idx


def prune_conv_input_channels(conv: nn.Conv2d, keep_in_idx: torch.Tensor) -> nn.Conv2d:
    """
    Prune Conv2d input channels by selecting keep_in_idx on dim=1 of weight.
    Returns a new conv with in_channels = len(keep_in_idx).
    """
    device = conv.weight.device
    dtype = conv.weight.dtype

    new_in = keep_in_idx.numel()
    new_conv = nn.Conv2d(
        in_channels=new_in,
        out_channels=conv.out_channels,
        kernel_size=conv.kernel_size,
        stride=conv.stride,
        padding=conv.padding,
        dilation=conv.dilation,
        groups=conv.groups,  # assumes groups-compatible; your model uses groups=1
        bias=(conv.bias is not None),
        padding_mode=conv.padding_mode,
    ).to(device=device, dtype=dtype)

    with torch.no_grad():
        # weight shape: (out, in, kH, kW)
        new_conv.weight.copy_(conv.weight.data[:, keep_in_idx].contiguous())
        if conv.bias is not None:
            new_conv.bias.copy_(conv.bias.data.contiguous())

    return new_conv


def rebuild_fc_after_pruning(model: nn.Module, example_input: torch.Tensor) -> None:
    """
    Rebuild model.fc input dim based on the current conv/ssrp path.
    Assumes model has attributes: conv1, conv2, conv3, ssrp_ms, flatten, fc.
    """
    device = next(model.parameters()).device
    model.eval()
    with torch.no_grad():
        x = example_input.to(device)
        feats = model.flatten(model.ssrp_ms(model.conv3(model.conv2(model.conv1(x)))))
        in_dim = feats.shape[1]

    old_out = model.fc.out_features
    model.fc = nn.Linear(in_dim, old_out).to(device)
    model.train()

def apply_structural_pruning(model: nn.Module, amount: float = 0.8, example_input: torch.Tensor = torch.randn(1, 1, 40, 862)) -> nn.Module:
    """
    Structural channel pruning for your CNN_PCAw_SSRPMS_KAN conv blocks.

    - Prunes conv1 out channels + bn1, updates conv2 input channels accordingly
    - Prunes conv2 out channels + bn2, updates conv3 input channels accordingly
    - Prunes conv3 out channels + bn3
    - Optionally rebuilds fc using example_input

    example_input should be shaped like your model input, e.g. (1, 1, F, T)
    """
    device = next(model.parameters()).device

    # ---- conv1 prune (conv1[1] is Conv2d, conv1[2] is BN) ----
    conv1_old = model.conv1[1]
    bn1_old = model.conv1[2]
    conv1_new, bn1_new, keep1 = prune_conv_bn_pair(conv1_old, bn1_old, amount)

    model.conv1[1] = conv1_new
    model.conv1[2] = bn1_new

    # ---- conv2 input prune to match conv1 kept outputs ----
    model.conv2[1] = prune_conv_input_channels(model.conv2[1], keep1)

    # ---- conv2 prune ----
    conv2_old = model.conv2[1]
    bn2_old = model.conv2[2]
    conv2_new, bn2_new, keep2 = prune_conv_bn_pair(conv2_old, bn2_old, amount)

    model.conv2[1] = conv2_new
    model.conv2[2] = bn2_new

    # ---- conv3 input prune to match conv2 kept outputs ----
    model.conv3[0] = prune_conv_input_channels(model.conv3[0], keep2)

    # ---- conv3 prune ----
    conv3_old = model.conv3[0]
    bn3_old = model.conv3[1]
    conv3_new, bn3_new, keep3 = prune_conv_bn_pair(conv3_old, bn3_old, amount)

    model.conv3[0] = conv3_new
    model.conv3[1] = bn3_new

    # Ensure the whole model stays on the same device
    model.to(device)

    # ---- fc rebuild (needed because flatten dim changes) ----
    if example_input is not None:
        rebuild_fc_after_pruning(model, example_input)

    return model