kaamd/ll
011
1from .zero_neuron import ZeroNeuron2from sentence_transformers.models import Module3import torch4 5class SparseEmbedding(Module):6 """ This module should be applied last (after Pooling, Normalize, etc.) """7 config_keys = ["n_in", "init_mean", "init_std", "temperature", "stretch", "eps"]8 def __init__(self,9 n_in: int,10 init_mean: float = 0.5,11 init_std: float = 0.01,12 temperature: float = 1.0,13 stretch: float = 0.1,14 eps: float = 1e-6):15 super(SparseEmbedding, self).__init__()16 self.n_in = n_in17 self.init_mean = init_mean18 self.init_std = init_std19 self.temperature = temperature20 self.stretch = stretch21 self.eps = eps22 self.sparsifyer = ZeroNeuron(23 in_features=n_in,24 out_features=n_in,25 init_mean=init_mean,26 init_std=init_std,27 temperature=temperature,28 stretch=stretch,29 eps=eps30 )31 32 def forward(self, features, *args, **kwargs):33 mask = self.sparsifyer(features["sentence_embedding"], dim=kwargs.get("dim", None))34 features["mask"] = mask35 features["sparsity_loss"] = self.sparsifyer.l0_norm(features["sentence_embedding"])36 37 return features38 39 def save(self, output_path: str):40 self.save_config(output_path)41 torch.save(self.sparsifyer.state_dict(), output_path + "/pytorch_model.bin")42 