LFM2.5-Encoder-350M-GGUF / fill-mask.py
mlabonne's picture v4zhong's picture
GGUF release: F16/Q8_0/Q4_0, card + fill-mask.py (stock llama.cpp usage) (#2)
9d64d2e
Raw History Blame Contribute Delete
2.15 kB
# /// script
# requires-python = ">=3.10"
# dependencies = ["numpy", "requests", "gguf"]
# ///
# fill-mask.py — masked-token prediction against a stock llama-server.
#
# The encoder's MLM head is tied to the token embeddings, so the logits at the
# mask position are just `hidden @ token_embd^T`: fetch the per-token hidden
# states from `llama-server --embeddings --pooling none`, read the embedding
# matrix straight out of the GGUF, and take the top-K at the mask position.
#
# llama-server -m LFM2.5-Encoder-230M-F16.gguf --embeddings --pooling none
# uv run fill-mask.py LFM2.5-Encoder-230M-F16.gguf "The capital of France is [MASK]."
import sys
import numpy as np
import requests
from gguf import GGUFReader
gguf_path, prompt = sys.argv[1], sys.argv[2]
topk = int(sys.argv[3]) if len(sys.argv) > 3 else 5
# token_embd from the GGUF (memory-mapped; fp16/fp32 tensors read directly)
reader = GGUFReader(gguf_path)
embd = next(t for t in reader.tensors if t.name == "token_embd.weight")
W = np.array(embd.data).astype(np.float32) # [n_vocab, n_embd]
# tokenize server-side, replacing [MASK] with the model's mask token id
def tokenize(text: str, special: bool) -> list[int]:
r = requests.post("http://localhost:8080/tokenize",
json={"content": text, "add_special": special, "parse_special": True})
return r.json()["tokens"]
meta = {f.name: f for f in reader.fields.values()}
mask_id = int(meta["tokenizer.ggml.mask_token_id"].parts[-1][0])
pre, _, post = prompt.partition("[MASK]")
toks = tokenize(pre, True) + [mask_id] + tokenize(post, False)
pos = toks.index(mask_id)
# one non-causal forward; per-token hidden states
r = requests.post("http://localhost:8080/embedding",
json={"content": toks})
hidden = np.array(r.json()[0]["embedding"], dtype=np.float32) # [n_tok, n_embd]
logits = hidden[pos] @ W.T
top = np.argsort(logits)[::-1][:topk]
detok = lambda t: requests.post("http://localhost:8080/detokenize", json={"tokens": [int(t)]}).json()["content"]
print(f"top-{topk} at [MASK]:")
for i, t in enumerate(top, 1):
print(f" {i:>2} {logits[t]:9.4f} '{detok(t)}'")