jonpark0's picture
Add EmbeddingGemma 2 converted for the AX650 / AX8850 NPU
9037c05 verified
Raw History Blame Contribute Delete
2.76 kB
#!/usr/bin/env python3
"""Embed text / images / audio with EmbeddingGemma 2 on an LLM-8850 (AXCL) card and compare them.
python3 eg2_cli.py "text:์˜ค๋Š˜ ์„œ์šธ ๋‚ ์”จ" "image:photos/cat.jpg" "audio:clip.wav"
python3 eg2_cli.py --prompt query --dim 256 "text:๊ณ ์–‘์ด ์‚ฌ์ง„" "image:a.jpg" "image:b.jpg"
python3 eg2_cli.py "mix:๋นจ๊ฐ„ ์šด๋™ํ™” <|image|>|shoe.jpg" (text with placeholders | media files)
Items: text:<string>, image:<path>, audio:<path>, mix:<text with <|image|>/<|audio|>>|<file>|<file>...
--prompt applies to text and mix items (query, document, sts, classification, clustering, code).
"""
import argparse
import time
from pathlib import Path
import numpy as np
from eg2_host import Eg2
def main():
ap = argparse.ArgumentParser(description=__doc__.splitlines()[0])
ap.add_argument("items", nargs="+")
ap.add_argument("--prompt", default=None)
ap.add_argument("--dim", type=int, default=768, choices=[768, 512, 256, 128])
ap.add_argument("--assets", default=str(Path(__file__).resolve().parent / "assets"))
args = ap.parse_args()
kinds = {i.split(":", 1)[0] for i in args.items}
load = ["text"]
if kinds & {"image", "mix"}:
load.append("vision")
if kinds & {"audio", "mix"}:
load.append("audio")
eg = Eg2(args.assets, load=tuple(load))
try:
vecs, labels = [], []
for item in args.items:
kind, value = item.split(":", 1)
t0 = time.perf_counter()
if kind == "text":
v = eg.encode(text=value, prompt=args.prompt, dim=args.dim)
elif kind == "image":
v = eg.encode(images=[value], dim=args.dim)
elif kind == "audio":
v = eg.encode(audios=[value], dim=args.dim)
elif kind == "mix":
text, *files = value.split("|")
imgs = [f for f in files if Path(f).suffix.lower() in (".jpg", ".jpeg", ".png", ".bmp", ".webp")]
auds = [f for f in files if f not in imgs]
v = eg.encode(text=text, images=imgs, audios=auds, prompt=args.prompt, dim=args.dim)
else:
raise SystemExit(f"unknown item kind: {kind}")
print(f"{(time.perf_counter() - t0) * 1000:7.1f} ms {item[:70]}")
vecs.append(v)
labels.append(item[:24])
if len(vecs) > 1:
sims = np.stack(vecs) @ np.stack(vecs).T
print("\ncosine similarity")
for lab, row in zip(labels, sims):
print(f"{lab:26s}" + " ".join(f"{x:6.3f}" for x in row))
else:
print(np.array2string(vecs[0][:8], precision=4), "...")
finally:
eg.close()
if __name__ == "__main__":
main()