(self, pruned_steps)
| 207 | return target_sparsity |
| 208 | |
| 209 | def lagrangian_regularization(self, pruned_steps): |
| 210 | target_sparsity = self.get_target_sparsity( |
| 211 | pruned_steps) if self.lagrangian_warmup > 0 else self.target_sparsity |
| 212 | expect_sparsity = 1 - self.get_num_parameters_and_constraint( |
| 213 | "hidden" in self.types) / self.prunable_model_size |
| 214 | |
| 215 | # lagrangian_loss = ( |
| 216 | # self.lambda_1 * (expect_sparsity - target_sparsity).abs() + |
| 217 | # self.lambda_2 * (expect_sparsity - target_sparsity).square() |
| 218 | # ) |
| 219 | |
| 220 | zero = torch.tensor(0.0, device=expect_sparsity.device) |
| 221 | lagrangian_loss = ( |
| 222 | self.lambda_1 * torch.maximum(target_sparsity - expect_sparsity, zero) + |
| 223 | self.lambda_2 * |
| 224 | torch.maximum(target_sparsity - expect_sparsity, zero).square() |
| 225 | ) |
| 226 | |
| 227 | return lagrangian_loss, expect_sparsity.detach().item(), target_sparsity |
| 228 | |
| 229 | # during training |
| 230 | def _sample_z(self, loga): |
no test coverage detected