Nemotron-3-Embed-8B-Community-MLX-4bit / nemotron3_embed_mlx.py
mattbusi's picture
CI sync from nemotron-embed-quant @ 8ce22703
47270e3 verified
Raw History Blame Contribute Delete
3.19 kB
"""Nemotron-3-Embed-1B MLX implementation (self-contained single file).
Original: nvidia/Nemotron-3-Embed-1B-BF16 (OpenMDW-1.1)
Changes: converted the causal-decoder mlx-lm ministral3 implementation into a
bidirectional encoder (causal mask removed, key-padding mask added),
then applied mean pooling + L2 normalization. Distributed weights may
carry MLX affine quantization.
Dependencies: mlx, mlx-lm, transformers, numpy
"""
import json
from pathlib import Path
import mlx.core as mx
import mlx.nn as nn
import numpy as np
from mlx_lm.models.ministral3 import (
ModelArgs,
TransformerBlock,
_get_llama_4_attn_scale,
)
from transformers import AutoTokenizer
NEG_INF = -1e9
PREFIXES = {"query": "query: ", "passage": "passage: "}
class NemotronEmbedModel(nn.Module):
def __init__(self, args: ModelArgs):
super().__init__()
self.args = args
self.embed_tokens = nn.Embedding(args.vocab_size, args.hidden_size)
self.layers = [TransformerBlock(args) for _ in range(args.num_hidden_layers)]
self.norm = nn.RMSNorm(args.hidden_size, eps=args.rms_norm_eps)
def __call__(self, input_ids: mx.array, attention_mask: mx.array) -> mx.array:
h = self.embed_tokens(input_ids)
attn_scale = _get_llama_4_attn_scale(
input_ids.shape[1],
0,
self.args.rope_parameters["llama_4_scaling_beta"],
self.args.rope_parameters["original_max_position_embeddings"],
).astype(h.dtype)
pad = (1 - attention_mask[:, None, None, :]).astype(h.dtype) * NEG_INF
for layer in self.layers:
h = layer(h, attn_scale, mask=pad)
h = self.norm(h).astype(mx.float32)
m = attention_mask[:, :, None].astype(mx.float32)
emb = (h * m).sum(axis=1) / m.sum(axis=1)
return emb / mx.linalg.norm(emb, axis=-1, keepdims=True)
def load(path: str):
p = Path(path)
if not p.is_dir():
from huggingface_hub import snapshot_download
p = Path(snapshot_download(path))
cfg = json.loads((p / "config.json").read_text())
model = NemotronEmbedModel(ModelArgs.from_dict(cfg))
q = cfg.get("quantization")
if q:
nn.quantize(
model,
group_size=q["group_size"],
bits=q["bits"],
mode=q.get("mode", "affine"),
)
model.load_weights(str(p / "model.safetensors"))
model.eval()
mx.eval(model.parameters())
tok = AutoTokenizer.from_pretrained(str(p))
tok.padding_side = "right"
return model, tok
def encode(model, tokenizer, texts, input_type=None, batch_size=8, max_length=4096):
if input_type is not None:
texts = [PREFIXES[input_type] + t for t in texts]
out = []
for i in range(0, len(texts), batch_size):
b = tokenizer(
texts[i : i + batch_size],
padding=True,
truncation=True,
max_length=max_length,
return_tensors="np",
)
e = model(mx.array(b["input_ids"]), mx.array(b["attention_mask"]))
out.append(np.asarray(e.astype(mx.float32)))
mx.clear_cache()
return np.vstack(out)