Feature Extraction
MLX
ONNX
Safetensors
English
gemma3_text
embeddings
code-retrieval
code-search
distillation
quantization
gemma3
Instructions to use sensiarion/embeddinggemma-300m-code-8L-distill-int8 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use sensiarion/embeddinggemma-300m-code-8L-distill-int8 with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] huggingface-cli download --local-dir embeddinggemma-300m-code-8L-distill-int8 sensiarion/embeddinggemma-300m-code-8L-distill-int8
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
Add thin EmbeddingGemma code model (curvature-pruned 8L + int8 + distill/contrastive)
38bb3d8 verified | """Bidirectional MLX encoder for a thin EmbeddingGemma (true Metal inference). | |
| mlx_lm's Gemma3 model runs CAUSALLY and reassigns sliding/full roles by layer | |
| INDEX (ignoring our sliced layer_types) -> wrong embeddings (parity ~0.75-0.92). | |
| EmbeddingGemma is BIDIRECTIONAL. This module reuses mlx_lm's Gemma3 layers but: | |
| * builds BIDIRECTIONAL masks (padding-only for full layers; symmetric sliding | |
| window for sliding layers), | |
| * sets the correct rope base per kept layer (sliding=10000, full=1e6) from | |
| config.layer_types, | |
| then mean-pools. Result matches torch (parity ~1.0). | |
| """ | |
| import os, json | |
| import numpy as np | |
| import mlx.core as mx | |
| from mlx_lm import load as mlx_load | |
| from mlx_lm.models.rope_utils import initialize_rope | |
| NEG = -1e9 | |
| def load_encoder(mlx_dir): | |
| """Load the converted mlx model and fix per-layer rope to match kept roles.""" | |
| model, _ = mlx_load(mlx_dir) | |
| cfg = json.load(open(os.path.join(mlx_dir, "config.json"))) | |
| layer_types = cfg["layer_types"] | |
| sw = int(cfg.get("sliding_window", 512)) | |
| head_dim = int(cfg.get("head_dim", 256)) | |
| hidden = int(cfg.get("hidden_size", 768)) | |
| assert len(layer_types) == len(model.model.layers), "layer_types/layers mismatch" | |
| for role, blk in zip(layer_types, model.model.layers): | |
| base = 10000.0 if role == "sliding_attention" else 1_000_000.0 | |
| blk.self_attn.rope = initialize_rope(dims=head_dim, base=base, traditional=False) | |
| model.eval() | |
| return model, layer_types, sw, hidden | |
| def _masks(attn_mask, L, sw): | |
| """Additive bidirectional masks: (full=[B,1,1,L] padding, sliding=+window).""" | |
| pad = (1.0 - attn_mask)[:, None, None, :] * NEG # [B,1,1,L] | |
| idx = mx.arange(L) | |
| within = mx.abs(idx[:, None] - idx[None, :]) < sw # [L,L] bool, symmetric | |
| win = mx.where(within, 0.0, NEG)[None, None] # [1,1,L,L] | |
| return pad, pad + win | |
| def encode(model, layer_types, sw, hidden, input_ids, attn_mask): | |
| """input_ids/attn_mask: mx int arrays [B,L]. Returns L2-normalized mean-pooled embs [B,d].""" | |
| h = model.model.embed_tokens(input_ids) | |
| h = h * mx.array(hidden ** 0.5, mx.bfloat16).astype(h.dtype) | |
| B, L = input_ids.shape | |
| full_mask, slide_mask = _masks(attn_mask.astype(h.dtype), L, sw) | |
| for role, blk in zip(layer_types, model.model.layers): | |
| m = slide_mask if role == "sliding_attention" else full_mask | |
| h = blk(h, m.astype(h.dtype), None) | |
| h = model.model.norm(h) # [B,L,d] | |
| am = attn_mask.astype(h.dtype)[:, :, None] | |
| v = (h * am).sum(1) / mx.maximum(am.sum(1), 1e-9) # mean pool | |
| v = v / mx.linalg.norm(v, axis=-1, keepdims=True) | |
| return v | |
| def embed_texts(model, layer_types, sw, hidden, tok, texts, max_len=256, batch=32): | |
| out = [] | |
| for i in range(0, len(texts), batch): | |
| chunk = texts[i:i + batch] | |
| enc = tok(chunk, padding=True, truncation=True, max_length=max_len, return_tensors="np") | |
| ids = mx.array(enc["input_ids"].astype(np.int32)) | |
| am = mx.array(enc["attention_mask"].astype(np.int32)) | |
| v = encode(model, layer_types, sw, hidden, ids, am) | |
| mx.eval(v) | |
| out.append(np.array(v)) | |
| return np.concatenate(out, 0) | |