Prune a linear layer (a model parameters) to keep only entries in index. Return the pruned layer as a new layer with requires_grad=True. Used to remove heads.
(layer, index, dim=0)
| 1144 | |
| 1145 | |
| 1146 | def prune_linear_layer(layer, index, dim=0): |
| 1147 | """ Prune a linear layer (a model parameters) to keep only entries in index. |
| 1148 | Return the pruned layer as a new layer with requires_grad=True. |
| 1149 | Used to remove heads. |
| 1150 | """ |
| 1151 | index = index.to(layer.weight.device) |
| 1152 | W = layer.weight.index_select(dim, index).clone().detach() |
| 1153 | if layer.bias is not None: |
| 1154 | if dim == 1: |
| 1155 | b = layer.bias.clone().detach() |
| 1156 | else: |
| 1157 | b = layer.bias[index].clone().detach() |
| 1158 | new_size = list(layer.weight.size()) |
| 1159 | new_size[dim] = len(index) |
| 1160 | new_layer = nn.Linear(new_size[1], new_size[0], bias=layer.bias is not None).to(layer.weight.device) |
| 1161 | new_layer.weight.requires_grad = False |
| 1162 | new_layer.weight.copy_(W.contiguous()) |
| 1163 | new_layer.weight.requires_grad = True |
| 1164 | if layer.bias is not None: |
| 1165 | new_layer.bias.requires_grad = False |
| 1166 | new_layer.bias.copy_(b.contiguous()) |
| 1167 | new_layer.bias.requires_grad = True |
| 1168 | return new_layer |
| 1169 | |
| 1170 | |
| 1171 | def prune_conv1d_layer(layer, index, dim=1): |
no test coverage detected