| 137 | return x |
| 138 | |
| 139 | def prune(self): |
| 140 | device = self.c_fc.weight.device |
| 141 | if self.hidden_z is None: |
| 142 | self.hidden_z = torch.ones( |
| 143 | (self.d_model,), dtype=torch.bool, device=device) |
| 144 | if self.intermediate_z is None: |
| 145 | self.intermediate_z = torch.ones( |
| 146 | (self.mlp_width,), dtype=torch.bool, device=device) |
| 147 | hidden_r = torch.where(self.hidden_z != 0)[0] |
| 148 | intermediate_r = torch.where(self.intermediate_z != 0)[0] |
| 149 | d_model = len(hidden_r) |
| 150 | mlp_width = len(intermediate_r) |
| 151 | # m = self |
| 152 | m = copy.deepcopy(self) |
| 153 | m.c_fc = nn.Linear(hidden_r.shape[0], intermediate_r.shape[0]) |
| 154 | m.c_proj = nn.Linear(intermediate_r.shape[0], hidden_r.shape[0]) |
| 155 | m.d_model = d_model |
| 156 | m.mlp_width = mlp_width |
| 157 | m.c_fc.weight = nn.Parameter( |
| 158 | (self.c_fc.weight[intermediate_r][:, hidden_r]).contiguous()) |
| 159 | m.c_fc.bias = nn.Parameter( |
| 160 | (self.c_fc.bias[intermediate_r]).contiguous()) |
| 161 | |
| 162 | m.c_proj.weight = nn.Parameter(((self.c_proj.weight * |
| 163 | self.intermediate_z.view(1, -1) * self.hidden_z.view(-1, 1))[hidden_r][:, intermediate_r]).contiguous()) |
| 164 | m.c_proj.bias = nn.Parameter( |
| 165 | ((self.c_proj.bias * self.hidden_z)[hidden_r]).contiguous()) |
| 166 | return m |
| 167 | |
| 168 | |
| 169 | class MultiheadAttention(nn.MultiheadAttention): |