MCPcopy Create free account
hub / github.com/microsoft/Cream / prune

Method prune

TinyCLIP/src/open_clip/model.py:820–844  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

818 return self.encode_text(text, normalized=normalized, **mask)
819
820 def prune(self):
821 device = self.token_embedding.weight.device
822 if self.hidden_z is None:
823 self.hidden_z = torch.ones(
824 self.text_projection.size(0), device=device)
825 if self.embed_dim_z is None:
826 self.embed_dim_z = torch.ones(
827 self.text_projection.size(1), device=device)
828 mod = self
829 self_copy = copy.deepcopy(self)
830 hidden_r = self.hidden_z > 0
831 mod.token_embedding = nn.Embedding(
832 self_copy.token_embedding.weight.shape[0], hidden_r.sum())
833 mod.positional_embedding = nn.Parameter(
834 torch.empty(self_copy.context_length, hidden_r.sum()))
835 mod.token_embedding.weight = nn.Parameter(
836 (self_copy.token_embedding.weight * self_copy.hidden_z.view(1, -1))[:, hidden_r])
837 mod.positional_embedding = nn.Parameter(
838 (self_copy.positional_embedding * self_copy.hidden_z.view(1, -1))[:, hidden_r])
839 mod.transformer = self.transformer.prune()
840 mod.ln_final = self.ln_final.prune()
841 embed_dim_r = self.embed_dim_z > 0
842 mod.text_projection = nn.Parameter(
843 (self.text_projection * self.hidden_z.view(-1, 1) * self.embed_dim_z.view(1, -1))[hidden_r][:, embed_dim_r])
844 return mod
845
846
847class LogitScale(nn.Module):

Callers

nothing calls this directly

Calls 2

sizeMethod · 0.45
pruneMethod · 0.45

Tested by

no test coverage detected