MCPcopy Create free account
hub / github.com/KohakuBlueleaf/LyCORIS / get_weight

Method get_weight

lycoris/modules/loha.py:194–226  ·  view source on GitHub ↗
(self, shape)

Source from the content-addressed store, hash-verified

192 self.register_buffer(
193 "scalar", torch.ones_like(self.scalar), persistent=False
194 )
195
196 def get_weight(self, shape):
197 scale = torch.tensor(
198 self.scale, dtype=self.hada_w1_b.dtype, device=self.hada_w1_b.device
199 )
200 if self.tucker:
201 weight = loha_diff_weight(
202 self.hada_w1_b,
203 self.hada_w1_a,
204 self.hada_w2_b,
205 self.hada_w2_a,
206 self.hada_t1,
207 self.hada_t2,
208 gamma=scale,
209 )
210 else:
211 weight = loha_diff_weight(
212 self.hada_w1_b,
213 self.hada_w1_a,
214 self.hada_w2_b,
215 self.hada_w2_a,
216 None,
217 None,
218 gamma=scale,
219 )
220 if shape is not None:
221 weight = weight.reshape(shape)
222 if self.training and self.rank_dropout:
223 drop = (torch.rand(weight.size(0)) > self.rank_dropout).to(weight.dtype)
224 drop = drop.view(-1, *[1] * len(weight.shape[1:])).to(weight.device)
225 if self.rank_dropout_scale:
226 drop /= drop.mean()
227 weight *= drop
228 return weight
229

Callers 4

get_diff_weightMethod · 0.95
apply_max_normMethod · 0.95
bypass_forward_diffMethod · 0.95
forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected