Forward pass for training. Args: x: TODO.
(self, x: torch.Tensor)
| 215 | self.merged = True |
| 216 | |
| 217 | def forward(self, x: torch.Tensor): |
| 218 | """Forward pass for training. |
| 219 | |
| 220 | Args: |
| 221 | x: TODO. |
| 222 | """ |
| 223 | def T(w): |
| 224 | """T. |
| 225 | |
| 226 | Args: |
| 227 | w: TODO. |
| 228 | """ |
| 229 | return w.T if self.fan_in_fan_out else w |
| 230 | |
| 231 | if self.r > 0 and not self.merged: |
| 232 | result = F.linear(x, T(self.weight), bias=self.bias) |
| 233 | if self.r > 0: |
| 234 | result += (self.lora_dropout(x) @ self.lora_A.T @ self.lora_B.T) * self.scaling |
| 235 | return result |
| 236 | else: |
| 237 | return F.linear(x, T(self.weight), bias=self.bias) |
| 238 | |
| 239 | |
| 240 | class MergedLinear(nn.Linear, LoRALayer): |