mo-ocr / utils /visual_features.py
mnhkahn's picture
feat(seal): 扩充说文字形并融合碑帖候选
f2487c7
Raw
History Blame Contribute Delete
6.62 kB
"""Pretrained vision-model embeddings for small-seal glyph retrieval."""
from __future__ import annotations
from pathlib import Path
from typing import Iterable
import numpy as np
import torch
import torch.nn as nn
from PIL import Image
from transformers import AutoImageProcessor, AutoModel
from torchvision.models import resnet50
DEFAULT_VISUAL_MODEL = "facebook/dinov2-small"
def load_visual_model(model_id: str, device: str = "cpu"):
"""Load a Hugging Face vision backbone in inference-only mode."""
processor = AutoImageProcessor.from_pretrained(model_id)
model = AutoModel.from_pretrained(model_id)
model.to(device)
model.eval()
return processor, model
def embed_images(
image_paths: Iterable[Path], processor, model, batch_size: int = 16, device: str = "cpu"
) -> np.ndarray:
"""Create L2-normalized CLS embeddings for a sequence of glyph images."""
paths = list(image_paths)
chunks: list[np.ndarray] = []
with torch.inference_mode():
for start in range(0, len(paths), batch_size):
batch_paths = paths[start : start + batch_size]
images = []
for path in batch_paths:
with Image.open(path) as image:
images.append(image.convert("RGB"))
inputs = processor(images=images, return_tensors="pt")
inputs = {key: value.to(device) for key, value in inputs.items()}
outputs = model(**inputs)
embedding = getattr(outputs, "pooler_output", None)
if embedding is None:
embedding = outputs.last_hidden_state[:, 0]
embedding = torch.nn.functional.normalize(embedding, dim=1)
chunks.append(embedding.cpu().numpy().astype(np.float32))
return np.concatenate(chunks, axis=0)
def load_local_resnet50_feature_model(
checkpoint_path: Path, num_classes: int, device: str = "cpu"
) -> nn.Module:
"""Load this project's classifier and expose its 2048-D pre-classifier features."""
model = resnet50(weights=None)
model.fc = nn.Sequential(nn.Dropout(0.3), nn.Linear(model.fc.in_features, num_classes))
try:
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=True)
except Exception:
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
if isinstance(checkpoint, nn.Module):
state_dict = checkpoint.state_dict()
elif isinstance(checkpoint, dict):
state_dict = checkpoint.get("state_dict", checkpoint.get("model_state", checkpoint))
else:
raise ValueError("不支持的模型文件格式")
model.load_state_dict(state_dict, strict=True)
feature_model = nn.Sequential(*list(model.children())[:-1]).to(device)
feature_model.eval()
return feature_model
def load_seal_retriever_feature_model(checkpoint_path: Path, device: str = "cpu") -> nn.Module:
"""Load a small-seal fine-tuned backbone and expose 2048-D features.
The training checkpoint intentionally stores no classification head, so it
can be used for retrieval across both supervised and single-sample chars.
"""
model = resnet50(weights=None)
try:
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=True)
except Exception:
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
state_dict = checkpoint.get("backbone_state_dict", checkpoint) if isinstance(checkpoint, dict) else checkpoint
if not isinstance(state_dict, dict):
raise ValueError("不支持的小篆检索训练权重格式")
incompatible = model.load_state_dict(state_dict, strict=False)
unexpected = [key for key in incompatible.unexpected_keys if not key.startswith("fc.")]
if unexpected:
raise ValueError(f"小篆检索权重包含未知参数:{unexpected[0]}")
feature_model = nn.Sequential(*list(model.children())[:-1]).to(device)
feature_model.eval()
return feature_model
def embed_images_local_resnet50(
image_paths: Iterable[Path],
model: nn.Module,
batch_size: int = 16,
device: str = "cpu",
image_size: int = 128,
) -> np.ndarray:
"""Create L2-normalized embeddings with the project's existing classifier backbone."""
paths = list(image_paths)
chunks: list[np.ndarray] = []
mean = torch.tensor([0.485, 0.456, 0.406], dtype=torch.float32).view(3, 1, 1)
std = torch.tensor([0.229, 0.224, 0.225], dtype=torch.float32).view(3, 1, 1)
with torch.inference_mode():
for start in range(0, len(paths), batch_size):
images = []
for path in paths[start : start + batch_size]:
with Image.open(path) as image:
rgb = image.convert("RGB").resize(
(image_size, image_size), Image.Resampling.BILINEAR
)
pixels = np.asarray(rgb, dtype=np.float32) / 255.0
tensor = torch.from_numpy(pixels).permute(2, 0, 1)
images.append((tensor - mean) / std)
batch = torch.stack(images).to(device)
embedding = model(batch).flatten(1)
embedding = torch.nn.functional.normalize(embedding, dim=1)
chunks.append(embedding.cpu().numpy().astype(np.float32))
return np.concatenate(chunks, axis=0)
def embed_pil_images_local_resnet50(
images: Iterable[Image.Image],
model: nn.Module,
batch_size: int = 16,
device: str = "cpu",
image_size: int = 128,
) -> np.ndarray:
"""Embed in-memory images using the same preprocessing as the index builder."""
image_list = list(images)
chunks: list[np.ndarray] = []
mean = torch.tensor([0.485, 0.456, 0.406], dtype=torch.float32).view(3, 1, 1)
std = torch.tensor([0.229, 0.224, 0.225], dtype=torch.float32).view(3, 1, 1)
with torch.inference_mode():
for start in range(0, len(image_list), batch_size):
tensors = []
for image in image_list[start : start + batch_size]:
rgb = image.convert("RGB").resize(
(image_size, image_size), Image.Resampling.BILINEAR
)
pixels = np.asarray(rgb, dtype=np.float32) / 255.0
tensor = torch.from_numpy(pixels).permute(2, 0, 1)
tensors.append((tensor - mean) / std)
batch = torch.stack(tensors).to(device)
embedding = torch.nn.functional.normalize(model(batch).flatten(1), dim=1)
chunks.append(embedding.cpu().numpy().astype(np.float32))
return np.concatenate(chunks, axis=0)