Mettle / run_test.py
kykatro's picture
Add Mettle RC2 weights, loader, and model card
0e83a2b verified
Raw
History Blame Contribute Delete
1.64 kB
"""Smoke-test a packaged local or downloaded Mettle-MX repository."""
import argparse
import torch
from PIL import Image
from transformers import AutoImageProcessor, AutoModel
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("model", help="local directory or Hugging Face repository ID")
parser.add_argument("--image", help="optional RGB image to encode")
parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
return parser.parse_args()
def main():
args = parse_args()
processor = AutoImageProcessor.from_pretrained(args.model)
model = AutoModel.from_pretrained(
args.model,
trust_remote_code=True,
).to(args.device).eval()
if args.image:
image = Image.open(args.image).convert("RGB")
else:
image = Image.new("RGB", (224, 224), color=(128, 96, 128))
inputs = processor(images=image, return_tensors="pt")
pixel_values = inputs["pixel_values"].to(args.device)
with torch.inference_mode():
cls = model.encode(pixel_values, feature_view="cls")
cls_mean = model.encode(pixel_values, feature_view="cls_mean")
if tuple(cls.shape) != (1, 1536):
raise RuntimeError(f"unexpected CLS shape: {tuple(cls.shape)}")
if tuple(cls_mean.shape) != (1, 3072):
raise RuntimeError(f"unexpected CLS+mean shape: {tuple(cls_mean.shape)}")
if not torch.equal(cls, cls_mean[:, :1536]):
raise RuntimeError("CLS differs between the two feature views")
print(f"PASS cls={tuple(cls.shape)} cls_mean={tuple(cls_mean.shape)}")
if __name__ == "__main__":
main()