(self, x: torch.Tensor, *args, **kwargs)
| 299 | gamma=gamma, |
| 300 | extra_args=self.kw_dict, |
| 301 | ) |
| 302 | return self.drop(diff * scalar) |
| 303 | |
| 304 | def bypass_forward(self, x, scale=1): |
| 305 | return self.org_forward(x) + self.bypass_forward_diff(x, scale=scale) |
| 306 | |
| 307 | def forward(self, x: torch.Tensor, *args, **kwargs): |
| 308 | if self.module_dropout and self.training: |
| 309 | if torch.rand(1) < self.module_dropout: |
| 310 | return self.org_forward(x, *args, **kwargs) |
| 311 | |
| 312 | if self.bypass_mode: |
| 313 | return self.bypass_forward(x, scale=self.multiplier) |
| 314 | |
| 315 | base = self.org_forward(x, *args, **kwargs) |
| 316 | base_weight = self._current_weight().to(x.device) |
| 317 | diff_weight = self.get_weight(self.shape).to(base_weight.dtype) * self.scalar |
| 318 | |
| 319 | if self.wd: |
| 320 | new_weight = self.apply_weight_decompose( |
| 321 | base_weight + diff_weight, self.multiplier |
| 322 | ) |
| 323 | else: |
| 324 | new_weight = base_weight + diff_weight * self.multiplier |
| 325 |
nothing calls this directly
no test coverage detected