import math
from typing import Optional
import mlx.core as mx
from mlx.nn.layers.base import Module
from mlx.nn.layers.quantized import QuantizedEmbedding
class Embedding(Module):
def __init__(self, num_embeddings: int, dims: int):
super().__init__()
scale = math.sqrt(1 / dims)
self.weight = mx.random.normal(shape=(num_embeddings, dims), scale=scale)
def _extra_repr(self):
return f"{self.weight.shape[0]}, {self.weight.shape[1]}"
def __call__(self, x):
return self.weight[x]
def as_linear(self, x):
return x @ self.weight.T
def to_quantized(
self,
group_size: Optional[int] = None,
bits: Optional[int] = None,
mode: str = "affine",
quantize_input: bool = False,
):
if quantize_input:
raise ValueError("Quantized input is not supported.")
return QuantizedEmbedding.from_embedding(self, group_size, bits, mode)