#!/usr/bin/env python3
import torch
import torch.nn as nn
import torch.nn.functional as F

class GeGLU(nn.Module):
    def __init__(self, dim, d_large):
        super().__init__()
        # We project to 2 * d_large so we can chunk it in half later
        self.w12 = nn.Linear(dim, 2 * d_large, bias=False)
        self.w3 = nn.Linear(d_large, dim, bias=False)

    def forward(self, x):
# projecting to double size
        projected = self.w12(x)   
# chunking in half
        x_linear, x_gate = projected.chunk(2, dim=-1)
        
# gating and down-projecting
        return self.w3(x_linear * F.gelu(x_gate))

# Example usage:
# Batch size=2, Sequence length=4, Embedding dim=128
x = torch.randn(2, 4, 128)
layer = GeGLU(dim=128, d_large=512)
output = layer(x)
print(output.shape)  # torch.Size([2, 4, 128])
