(self, shape)
| 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 |
no outgoing calls
no test coverage detected