INNER CODE UNIT · Python
forward
Leeroo-AI/mergoo · mergoo/compose_layers.py:119
def forward(self, x, *args, **kwargs):
"""
This method is designed to be a drop-in-replacement for the peft LoRA layers' .forward method.
To use it, a bound method must be created (bound to an instance of the LoRALayer class).
"""
previous_dtype = x.dtype
gate_logits = self.gate(x) # b,s,N
weights, selected_experts = torch.topk(
gate_logits, self.num_experts_per_tok
) # b,s,n
weights = F.softmax(weights, dim=2, dtype=torch.float).to(
previous_dtype
) # b,s,n
result = self.base_layer(x, *args, **kwargs)
"""TODO MAYBE
- tensorize this loop add learnable weights here