Implement the Encoder protocol for custom modalities
mainSemHash allows you to use any modality (like images) by providing a custom encoder that implements the Encoder protocol. An encoder must have an encode(inputs, **kwargs) method that accepts a batch of inputs and returns a numpy array of embeddings.
Example implementation for images using timm:
class VisionEncoder:
"""Custom encoder using timm models. Implements the Encoder protocol."""
def __init__(self, model_name: str = "mobilenetv3_small_100.lamb_in1k"):
self.model = timm.create_model(model_name, pretrained=True, num_classes=0).eval()
data_config = timm.data.resolve_model_data_config(self.model)
self.transform = timm.data.create_transform(**data_config, is_training=False)
def encode(self, inputs, batch_size: int = 128):
"""Encode a batch of PIL images into embeddings."""
import numpy as np
# Convert grayscale to RGB if needed
rgb_inputs = [img.convert("RGB") if img.mode != "RGB" else img for img in inputs]
# Process in batches to avoid memory issues
all_embeddings = []
with torch.no_grad():
for i in range(0, len(rgb_inputs), batch_size):
batch_inputs = rgb_inputs[i : i + batch_size]
batch = torch.stack([self.transform(img) for img in batch_inputs])
embeddings = self.model(batch).numpy()
all_embeddings.append(embeddings)
return np.vstack(all_embeddings)