"""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)