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 

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…