MCPcopy Create free account
hub / github.com/SooLab/CGFormer / prune_linear_layer

Function prune_linear_layer

bert/modeling_utils.py:1146–1168  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

1144
1145
1146def 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
1171def prune_conv1d_layer(layer, index, dim=1):

Callers 2

prune_layerFunction · 0.85
prune_headsMethod · 0.85

Calls 1

toMethod · 0.80

Tested by

no test coverage detected