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

Method prune

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

Source from the content-addressed store, hash-verified

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
169class MultiheadAttention(nn.MultiheadAttention):

Callers 8

pruneMethod · 0.45
pruneMethod · 0.45
pruneMethod · 0.45
pruneMethod · 0.45
pruneMethod · 0.45
train_one_epochFunction · 0.45
mainFunction · 0.45
mainFunction · 0.45

Calls

no outgoing calls

Tested by 1

mainFunction · 0.36