@jaxtyped(typechecker=beartype) def pool(x: Float[Tensor, "B T d"], mask: Bool[Tensor, "B T"]) -> Float[Tensor, "B d"]: m = mask[:, :, None].float() return (x * m).sum(1) / m.sum(1)