File size: 3,191 Bytes
3658b11
 
 
 
 
 
 
 
 
 
47270e3
3658b11
 
 
 
 
 
47270e3
 
 
 
 
3658b11
 
 
 
 
47270e3
3658b11
 
 
 
 
 
 
 
 
 
 
47270e3
 
3658b11
 
 
 
 
 
 
 
 
 
 
47270e3
3658b11
 
 
 
47270e3
3658b11
 
 
 
 
47270e3
 
 
 
 
 
3658b11
 
 
 
 
 
 
47270e3
 
3658b11
 
 
 
47270e3
 
 
 
 
 
 
3658b11
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
"""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)