import torch

def layer_norm(x, gamma, beta, eps=1e-5):
    """x: [B, T, d]; gamma / beta: [d]."""
    z = x.to(torch.float32) if x.dtype in (
        torch.float16, torch.bfloat16) else x
    mu = z.mean(dim=-1, keepdim=True)
    centered = z - mu
    sigma_sq = centered.square().mean(dim=-1, keepdim=True)
    normalized = centered * torch.rsqrt(sigma_sq + eps)
    y = normalized * gamma + beta
    return y.to(x.dtype)