(self, input_tensor)
| 204 | return output |
| 205 | |
| 206 | def forward(self, input_tensor): |
| 207 | if self.disabled: |
| 208 | return F.linear(input_tensor, self.weight) |
| 209 | assert self._initialized, "init_lora was never called, please initialize before proceeding" |
| 210 | |
| 211 | # Gather the sharded base weight |
| 212 | if self.zero_shards > 1: |
| 213 | with torch.no_grad(): |
| 214 | base_weight = self.full_weight() |
| 215 | elif self.quantization_config: |
| 216 | base_weight = self.weight.dequantized() |
| 217 | else: |
| 218 | base_weight = self.weight |
| 219 | |
| 220 | base_weight_output = F.linear(input_tensor, base_weight) |
| 221 | lora_output = self.lora_weight_2(self.lora_weight_1(input_tensor)) |
| 222 | return base_weight_output + self.lora_scaling_factor * lora_output |
nothing calls this directly
no test coverage detected