Onnx and safetensors models produce different output

#2
by jmzzomg - opened

Hello, lightonai team!

First of all, thanks for the models you've built!

I wanted to add a couple of them to https://github.com/qdrant/fastembed

Since fastembed uses onnx-converted models, I took the official onnx export from this repository and compared it to the model which uses safetensors.
It turned out that the models differ slightly: the weights are different, as are the embeddings.

I verified it with the following script (AI-slop).

Would be nice if you could update the conversion. (I once again asked AI to make me a new conversion script and the results I've gotten agree to about 1e-6 on the embeddings, with cosine 1.000000. I can add it here if you'd like).

import numpy as np
import onnxruntime as ort
from huggingface_hub import snapshot_download
from pylate import models

path = snapshot_download("lightonai/mLateOn")

texts = [
    "Paris is the capital and most populous city of France.",
    "Ein HNSW-Index baut einen hierarchischen Graphen für die schnelle Nachbarsuche auf.",
    "Le Louvre est le musée le plus visité au monde.",
    "El Volga es el río más largo de Europa.",
    "Roma è la capitale d'Italia.",
    "Lisboa é a capital de Portugal.",
    "القاهرة هي عاصمة جمهورية مصر العربية.",
    "def add(a, b):\n    return a + b",
    "東京は日本の首都です。",
    "Late interaction models keep one vector per token and score with MaxSim. " * 30,
]

# 1) ground truth: torch model from model.safetensors (pylate ColBERT is a SentenceTransformer)
model = models.ColBERT(path, device="cpu")
expected = model.encode(texts, is_query=False)  # L2-normalized, padding removed

# 2) ONNX model from the same repo, fed the same tokens (pylate inserts the [D] prefix token)
session = ort.InferenceSession(f"{path}/model.onnx", providers=["CPUExecutionProvider"])
features = model.tokenize(texts, is_query=False)
output = session.run(
    None,
    {
        "input_ids": features["input_ids"].numpy(),
        "attention_mask": features["attention_mask"].numpy(),
    },
)[0]  # already L2-normalized
mask = features["attention_mask"].numpy().astype(bool)
actual = [output[i][mask[i]] for i in range(len(texts))]

# 3) differences
print("first token embedding of the first text (first 8 of 128 dims)")
print("  torch:", np.round(expected[0][0][:8], 4))
print("  onnx: ", np.round(actual[0][0][:8], 4))
print("  diff: ", np.round(expected[0][0][:8] - actual[0][0][:8], 4))

print("\nper text:")
all_cos, all_abs = [], []
for i, (e, a) in enumerate(zip(expected, actual)):
    cos = (e * a).sum(-1)  # both are unit vectors
    all_cos.append(cos)
    all_abs.append(np.abs(e - a).ravel())
    print(
        f"  text {i}: tokens={len(e):4d}  max|diff|={np.abs(e - a).max():.4f}  "
        f"cos min={cos.min():.4f} mean={cos.mean():.4f}"
    )

all_cos, all_abs = np.concatenate(all_cos), np.concatenate(all_abs)
print(f"\nwhole batch ({len(all_cos)} token embeddings):")
print(f"  abs diff: max={all_abs.max():.4f} mean={all_abs.mean():.6f}")
print(
    f"  cosine:   min={all_cos.min():.4f} mean={all_cos.mean():.4f} "
    f"median={np.median(all_cos):.4f} p1={np.percentile(all_cos, 1):.4f}"
)
print(f"  tokens with cosine < 0.99: {(all_cos < 0.99).mean():.1%}")

And got the following output:

first token embedding of the first text (first 8 of 128 dims)
  torch: [ 0.0521 -0.0984 -0.0034 -0.0073 -0.0416 -0.0248  0.0172 -0.1007]
  onnx:  [ 0.0485 -0.1007 -0.0047 -0.0045 -0.0423 -0.0262  0.0179 -0.0981]
  diff:  [ 0.0036  0.0023  0.0013 -0.0028  0.0008  0.0015 -0.0007 -0.0027]

per text:
  text 0: tokens=  14  max|diff|=0.0477  cos min=0.9871 mean=0.9987
  text 1: tokens=  22  max|diff|=0.0427  cos min=0.9869 mean=0.9988
  text 2: tokens=  15  max|diff|=0.0466  cos min=0.9874 mean=0.9989
  text 3: tokens=  13  max|diff|=0.0429  cos min=0.9889 mean=0.9988
  text 4: tokens=  11  max|diff|=0.0463  cos min=0.9874 mean=0.9985
  text 5: tokens=  10  max|diff|=0.0445  cos min=0.9882 mean=0.9985
  text 6: tokens=  12  max|diff|=0.0476  cos min=0.9880 mean=0.9987
  text 7: tokens=  18  max|diff|=0.0912  cos min=0.9510 mean=0.9949
  text 8: tokens=   9  max|diff|=0.0404  cos min=0.9909 mean=0.9986
  text 9: tokens= 423  max|diff|=0.0953  cos min=0.9420 mean=0.9973

whole batch (547 token embeddings):
  abs diff: max=0.0953 mean=0.004412
  cosine:   min=0.9420 mean=0.9975 median=0.9983 p1=0.9868
  tokens with cosine < 0.99: 3.1%
LightOn AI org

Hey George! Thank you so much for checking this, if you can post your AI script I'll check it and merge:)

Here is the slop to convert the model to model.onnx. The script name was mlateon_onnx_export_simple.py.

# /// script
# requires-python = ">=3.12"
# dependencies = [
#   "pylate==1.6.0",
#   "torch==2.11.0",
#   "transformers==5.3.0",
#   "sentence-transformers==5.3.0",
#   "onnx==1.19.1",
#   "onnxscript==0.7.2",
#   "onnxruntime==1.23.2",
#   "numpy==2.5.3",
# ]
# ///
"""Export lightonai/mLateOn to ONNX from model.safetensors, then check it against PyLate.

    uv run --python 3.12 mlateon_onnx_export_simple.py

Writes mlateon-onnx/model.onnx with the same inputs and output as the current model.onnx:
    input_ids, attention_mask: int64 [batch, sequence], [Q]/[D] token already inserted
    output: float32 [batch, sequence, 128], L2-normalized per token (padding included)
"""

from pathlib import Path

import numpy as np
import onnx
import onnxruntime as ort
import torch
from pylate import models

OUTPUT = Path("mlateon-onnx/model.onnx")


class MLateOnForONNX(torch.nn.Module):
    """Transformer -> 3 Dense layers -> L2 normalization, as in PyLate's encode()."""

    def __init__(self, colbert: models.ColBERT):
        super().__init__()
        transformer, *dense_layers = colbert
        self.transformer = transformer.auto_model
        # Call the whole Dense modules, not only their .linear: layers 1 and 2 have
        # use_residual=True, and Dense.forward() adds that residual projection.
        self.dense_layers = torch.nn.ModuleList(dense_layers)

    def forward(self, input_ids, attention_mask):
        output = self.transformer(input_ids=input_ids, attention_mask=attention_mask)
        features = {"token_embeddings": output.last_hidden_state}
        for dense in self.dense_layers:
            features = dense(features)
        return torch.nn.functional.normalize(features["token_embeddings"], p=2, dim=-1)


# Eager attention exports to plain MatMul/Softmax ops.
colbert = models.ColBERT(
    "lightonai/mLateOn", device="cpu", model_kwargs={"attn_implementation": "eager"}
)

# Trace with a padded batch that is longer than the 128-token local attention window,
# so both the padding mask and the sliding-window mask end up in the graph.
example = colbert.tokenize(["Late interaction " * 100, "A short document."], is_query=False)

OUTPUT.parent.mkdir(parents=True, exist_ok=True)
with torch.no_grad():
    torch.onnx.export(
        MLateOnForONNX(colbert).eval(),
        (example["input_ids"], example["attention_mask"]),
        str(OUTPUT),
        input_names=["input_ids", "attention_mask"],
        output_names=["output"],
        dynamic_shapes=(
            {0: "batch_size", 1: "sequence_length"},
            {0: "batch_size", 1: "sequence_length"},
        ),
        opset_version=18,
        dynamo=True,
        external_data=False,  # one 1.25 GB file, like the current model.onnx
        optimize=True,
    )

# Remove the exporter's debug metadata (Python stack traces with local paths). It is
# the only IR 10 feature in the graph, so the file can then be marked as IR 9, which
# onnxruntime >= 1.17 can load.
model = onnx.load(str(OUTPUT))
for node in model.graph.node:
    del node.metadata_props[:]
del model.graph.metadata_props[:]
model.ir_version = 9
onnx.save(model, str(OUTPUT))
onnx.checker.check_model(str(OUTPUT), full_check=True)
del model
print(f"Saved {OUTPUT}")

# Check: ONNX Runtime vs PyLate's encode() with default settings, on the same tokens.
# Same texts and output format as the comparison script earlier in this thread.
del colbert
reference = models.ColBERT("lightonai/mLateOn", device="cpu")
texts = [
    "Paris is the capital and most populous city of France.",
    "Ein HNSW-Index baut einen hierarchischen Graphen für die schnelle Nachbarsuche auf.",
    "Le Louvre est le musée le plus visité au monde.",
    "El Volga es el río más largo de Europa.",
    "Roma è la capitale d'Italia.",
    "Lisboa é a capital de Portugal.",
    "القاهرة هي عاصمة جمهورية مصر العربية.",
    "def add(a, b):\n    return a + b",
    "東京は日本の首都です。",
    # One sentence repeated 30 times is numerically sensitive. On this text, every float32
    # implementation (PyLate's included) is about 1e-5 away from a float64 run, versus
    # about 1e-7 on the other texts.
    "Late interaction models keep one vector per token and score with MaxSim. " * 30,
]
session = ort.InferenceSession(str(OUTPUT), providers=["CPUExecutionProvider"])
for is_query in (False, True):
    expected = reference.encode(texts, is_query=is_query, show_progress_bar=False)
    features = reference.tokenize(texts, is_query=is_query)
    output = session.run(
        None,
        {
            "input_ids": features["input_ids"].numpy(),
            "attention_mask": features["attention_mask"].numpy(),
        },
    )[0]
    # mLateOn's skiplist is empty, so PyLate only removes padding.
    mask = features["attention_mask"].numpy().astype(bool)
    actual = [output[i][mask[i]] for i in range(len(texts))]

    print("\nqueries:" if is_query else "\ndocuments:")
    for i, (e, a) in enumerate(zip(expected, actual)):
        max_diff = np.abs(e - a).max()
        cos = (e * a).sum(-1)  # both are unit vectors
        print(
            f"  text {i}: tokens={len(e):4d}  max|diff|={max_diff:.1e}  "
            f"cos min={cos.min():.7f} mean={cos.mean():.7f}"
        )
        assert max_diff < 1e-4, f"ONNX output differs from PyLate on text {i}"

Yet another slop to create model_int8.onnx

# /// script
# requires-python = ">=3.12"
# dependencies = [
#   "pylate==1.6.0",
#   "torch==2.11.0",
#   "transformers==5.3.0",
#   "sentence-transformers==5.3.0",
#   "onnx==1.19.1",
#   "onnxruntime==1.23.2",
#   "numpy==2.5.3",
# ]
# ///
"""Quantize mlateon-onnx/model.onnx to INT8, then check it against PyLate.

    uv run --python 3.12 mlateon_onnx_quantize_int8.py

Run mlateon_onnx_export_simple.py first. Writes mlateon-onnx/model_int8.onnx (about
313 MB) with the same inputs and output as model.onnx.
"""

from pathlib import Path

import numpy as np
from onnxruntime import InferenceSession
from onnxruntime.quantization import QuantType, quantize_dynamic
from pylate import models

FP32 = Path("mlateon-onnx/model.onnx")
INT8 = Path("mlateon-onnx/model_int8.onnx")

# Dynamic quantization: weights are stored as int8, and activations are quantized on
# the fly. LightOn's colbert-quantize makes the same call with per_channel=False.
# Per-channel scales give the same size and ops, but a closer match to PyLate.
quantize_dynamic(str(FP32), str(INT8), weight_type=QuantType.QInt8, per_channel=True)
print(f"Saved {INT8} ({INT8.stat().st_size / 1e6:.0f} MB)")

# Check: ONNX Runtime vs PyLate's encode() with default settings, on the same tokens.
# INT8 can't match float32 exactly. Activation scales are picked per batch, so results
# also shift slightly with what else is in the batch.
reference = models.ColBERT("lightonai/mLateOn", device="cpu")
texts = [
    "Paris is the capital and most populous city of France.",
    "Ein HNSW-Index baut einen hierarchischen Graphen für die schnelle Nachbarsuche auf.",
    "Le Louvre est le musée le plus visité au monde.",
    "El Volga es el río más largo de Europa.",
    "Roma è la capitale d'Italia.",
    "Lisboa é a capital de Portugal.",
    "القاهرة هي عاصمة جمهورية مصر العربية.",
    "def add(a, b):\n    return a + b",
    "東京は日本の首都です。",
    "Late interaction models keep one vector per token and score with MaxSim. " * 30,
]
session = InferenceSession(str(INT8), providers=["CPUExecutionProvider"])
for is_query in (False, True):
    expected = reference.encode(texts, is_query=is_query, show_progress_bar=False)
    features = reference.tokenize(texts, is_query=is_query)
    output = session.run(
        None,
        {
            "input_ids": features["input_ids"].numpy(),
            "attention_mask": features["attention_mask"].numpy(),
        },
    )[0]
    # mLateOn's skiplist is empty, so PyLate only removes padding.
    mask = features["attention_mask"].numpy().astype(bool)
    actual = [output[i][mask[i]] for i in range(len(texts))]

    print("\nqueries:" if is_query else "\ndocuments:")
    all_cos = []
    for i, (e, a) in enumerate(zip(expected, actual)):
        cos = (e * a).sum(-1)  # both are unit vectors
        all_cos.append(cos)
        print(f"  text {i}: tokens={len(e):4d}  cos min={cos.min():.4f} mean={cos.mean():.4f}")
    mean_cos = np.concatenate(all_cos).mean()
    print(f"  all tokens: cos mean={mean_cos:.4f}")
    # A broken quantization (wrong graph, missing layers) lands far below this.
    assert mean_cos > 0.98, "INT8 output is too far from PyLate"

Something else to mention:

onnx_config.json names the wrong model.

It says "model_name": "lightonai/LateOn-multilingual" instead of lightonai/mLateOn.
That's probably the checkpoint the bad ONNX was exported from. It's only a label and doesn't change the output, but it's worth correcting when they replace model.onnx.
The other fields in that file are correct: query/document length 8192, prefix ids, mask and pad ids.

2. tokenizer_config.json caps the length at 299 tokens.

It sets "max_length": 299 and "model_max_length": 299.
The rest of the repo says 8192: config_sentence_transformers.json (query/document length 8192), sentence_bert_config.json (max_seq_length 8191) and tokenizer.json (truncation 8191).
PyLate isn't affected, because Sentence Transformers truncates using max_seq_length.
Tools that read model_max_length, like plain AutoTokenizer, would cut inputs at 299 tokens.

Thank you for all of this! Me and my own clunker reviewed it, and just merged https://huggingface.co/lightonai/mLateOn/discussions/3 :)
If it's all good on your side I'll close!

Thank you @ameliechatelain

Yes, feel free to close

ameliechatelain changed discussion status to closed

Sign up or log in to comment