| """High-level Forge2Vec API and vector arithmetic objects.""" |
|
|
| from __future__ import annotations |
|
|
| from dataclasses import dataclass, field, replace |
| from pathlib import Path |
| from typing import Any, Iterable, Iterator, Mapping, Optional, Sequence, Union |
|
|
| import numpy as np |
| import torch |
| from PIL import Image |
| from safetensors.torch import load_file |
| from transformers import AutoTokenizer |
|
|
| from .features import genre_indices, metadata_vector |
| from .modeling import UnifiedAttentionForge2Vec |
|
|
|
|
| Poster = Union[Image.Image, np.ndarray, torch.Tensor] |
|
|
|
|
| @dataclass(frozen=True) |
| class ForgeItem: |
| """An anime or manga profile accepted by Forge2Vec.""" |
|
|
| title: str = "" |
| native_title: str = "" |
| synonyms: Sequence[str] = field(default_factory=tuple) |
| synopsis: str = "" |
| synopsis_ua: str = "" |
| genres: Sequence[str] = field(default_factory=tuple) |
| content_type: str = "anime" |
| year: Optional[int] = None |
| score: Optional[float] = None |
| poster: Optional[Poster] = field(default=None, repr=False, compare=False) |
| id: Optional[Union[str, int]] = None |
|
|
| @classmethod |
| def from_value(cls, value: Union["ForgeItem", Mapping[str, Any]]) -> "ForgeItem": |
| if isinstance(value, cls): |
| return value |
| if not isinstance(value, Mapping): |
| raise TypeError("item must be a ForgeItem or a mapping") |
| aliases = { |
| "ua_title": "title", |
| "en_title": "title", |
| "original_title": "native_title", |
| "alternate_names": "synonyms", |
| "ua_description": "synopsis_ua", |
| "en_description": "synopsis", |
| "type": "content_type", |
| } |
| normalized = dict(value) |
| for old, new in aliases.items(): |
| if new not in normalized and old in normalized: |
| normalized[new] = normalized[old] |
| fields = cls.__dataclass_fields__ |
| return cls(**{key: value for key, value in normalized.items() if key in fields}) |
|
|
|
|
| class ForgeVector: |
| """A Forge2Vec embedding supporting ordinary vector arithmetic.""" |
|
|
| __array_priority__ = 1000 |
|
|
| def __init__( |
| self, |
| values: Union[np.ndarray, torch.Tensor, Sequence[float]], |
| *, |
| catalogue: Optional["ForgeCatalogue"] = None, |
| excluded_indices: Iterable[int] = (), |
| ) -> None: |
| array = np.asarray(values, dtype=np.float32).reshape(-1) |
| if array.shape != (256,): |
| raise ValueError(f"ForgeVector must have shape (256,), received {array.shape}") |
| self._values = array |
| self._catalogue = catalogue |
| self._excluded_indices = frozenset(excluded_indices) |
|
|
| @property |
| def values(self) -> np.ndarray: |
| return self._values.copy() |
|
|
| @property |
| def shape(self) -> tuple[int, ...]: |
| return self._values.shape |
|
|
| def numpy(self) -> np.ndarray: |
| return self.values |
|
|
| def tensor(self, device: Optional[Union[str, torch.device]] = None) -> torch.Tensor: |
| return torch.from_numpy(self._values.copy()).to(device=device) |
|
|
| def normalized(self) -> "ForgeVector": |
| norm = float(np.linalg.norm(self._values)) |
| if norm <= 1e-9: |
| raise ValueError("cannot normalize a zero vector") |
| return self._new(self._values / norm) |
|
|
| def find(self, limit: int = 10) -> "ForgeResults": |
| if self._catalogue is None: |
| raise ValueError("this vector is not attached to a catalogue") |
| return self._catalogue.find(self, limit=limit) |
|
|
| def _new(self, values, other: Optional["ForgeVector"] = None) -> "ForgeVector": |
| catalogue = self._catalogue |
| excluded = self._excluded_indices |
| if other is not None: |
| if catalogue is None: |
| catalogue = other._catalogue |
| elif other._catalogue is not None and other._catalogue is not catalogue: |
| catalogue = None |
| excluded = excluded | other._excluded_indices |
| return ForgeVector(values, catalogue=catalogue, excluded_indices=excluded) |
|
|
| def __add__(self, other: "ForgeVector") -> "ForgeVector": |
| if not isinstance(other, ForgeVector): |
| return NotImplemented |
| return self._new(self._values + other._values, other) |
|
|
| def __sub__(self, other: "ForgeVector") -> "ForgeVector": |
| if not isinstance(other, ForgeVector): |
| return NotImplemented |
| return self._new(self._values - other._values, other) |
|
|
| def __mul__(self, scalar: float) -> "ForgeVector": |
| return self._new(self._values * float(scalar)) |
|
|
| def __rmul__(self, scalar: float) -> "ForgeVector": |
| return self * scalar |
|
|
| def __truediv__(self, scalar: float) -> "ForgeVector": |
| if float(scalar) == 0.0: |
| raise ZeroDivisionError("cannot divide a ForgeVector by zero") |
| return self._new(self._values / float(scalar)) |
|
|
| def __neg__(self) -> "ForgeVector": |
| return self._new(-self._values) |
|
|
| def __array__(self, dtype=None) -> np.ndarray: |
| return np.asarray(self._values, dtype=dtype) |
|
|
| def __repr__(self) -> str: |
| return f"ForgeVector(shape={self.shape}, norm={np.linalg.norm(self._values):.4f})" |
|
|
|
|
| @dataclass(frozen=True) |
| class ForgeMatch: |
| rank: int |
| item: ForgeItem |
| vector: ForgeVector |
| similarity: float |
|
|
|
|
| class ForgeResults(Sequence[ForgeMatch]): |
| def __init__(self, matches: Sequence[ForgeMatch]) -> None: |
| self._matches = tuple(matches) |
|
|
| def __getitem__(self, index): |
| return self._matches[index] |
|
|
| def __len__(self) -> int: |
| return len(self._matches) |
|
|
| def __iter__(self) -> Iterator[ForgeMatch]: |
| return iter(self._matches) |
|
|
| def __repr__(self) -> str: |
| lines = ["ForgeResults("] |
| lines.extend( |
| f" {match.rank}. {match.item.title} ({match.similarity:.4f})" |
| for match in self._matches |
| ) |
| return "\n".join((*lines, ")")) |
|
|
|
|
| class ForgeCatalogue: |
| """An in-memory cosine index for anime and manga vectors.""" |
|
|
| def __init__(self, model: "Forge2Vec", items: Sequence[ForgeItem], embeddings: np.ndarray) -> None: |
| self.model = model |
| self.items = tuple(items) |
| values = np.asarray(embeddings, dtype=np.float32) |
| norms = np.linalg.norm(values, axis=1, keepdims=True) |
| self.embeddings = values / np.maximum(norms, 1e-9) |
| self._titles: dict[str, int] = {} |
| for index, item in enumerate(self.items): |
| for title in (item.title, item.native_title, *item.synonyms): |
| if title: |
| self._titles.setdefault(title.casefold().strip(), index) |
|
|
| def vec(self, title_or_index: Union[str, int]) -> ForgeVector: |
| if isinstance(title_or_index, str): |
| key = title_or_index.casefold().strip() |
| if key not in self._titles: |
| raise KeyError(f"title is not present in the catalogue: {title_or_index}") |
| index = self._titles[key] |
| else: |
| index = int(title_or_index) |
| return ForgeVector(self.embeddings[index], catalogue=self, excluded_indices=(index,)) |
|
|
| def find(self, query: Union[ForgeVector, ForgeItem, Mapping[str, Any]], limit: int = 10) -> ForgeResults: |
| vector = query if isinstance(query, ForgeVector) else self.model.vec(query) |
| normalized = vector.normalized()._values |
| scores = self.embeddings @ normalized |
| order = np.argsort(scores)[::-1] |
| selected = [index for index in order if index not in vector._excluded_indices][:limit] |
| matches = [ |
| ForgeMatch( |
| rank=rank, |
| item=self.items[index], |
| vector=ForgeVector(self.embeddings[index], catalogue=self, excluded_indices=(index,)), |
| similarity=float(scores[index]), |
| ) |
| for rank, index in enumerate(selected, 1) |
| ] |
| return ForgeResults(matches) |
|
|
|
|
| class Forge2Vec: |
| """Load and run hikka-forge2vec.""" |
|
|
| default_model_id = "Lorg0n/hikka-forge2vec" |
|
|
| def __init__( |
| self, |
| model_id_or_path: Union[str, Path] = default_model_id, |
| *, |
| device: Optional[Union[str, torch.device]] = None, |
| revision: Optional[str] = None, |
| cache_dir: Optional[Union[str, Path]] = None, |
| ) -> None: |
| self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu")) |
| root = Path(model_id_or_path) |
| if not root.is_dir(): |
| from huggingface_hub import snapshot_download |
|
|
| root = Path(snapshot_download( |
| repo_id=str(model_id_or_path), revision=revision, cache_dir=cache_dir, |
| allow_patterns=("config.json", "model.safetensors", "assets/tokenizer/*"), |
| )) |
| import json |
|
|
| self.config = json.loads((root / "config.json").read_text(encoding="utf-8")) |
| tokenizer_path = root / "assets" / "tokenizer" |
| self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_path, local_files_only=True) |
| self.model = UnifiedAttentionForge2Vec( |
| str(tokenizer_path), max_style_weight=float(self.config["max_style_weight"]) |
| ) |
| incompatible = self.model.load_state_dict( |
| load_file(str(root / "model.safetensors"), device="cpu"), strict=False |
| ) |
| if incompatible.missing_keys or incompatible.unexpected_keys: |
| raise RuntimeError( |
| f"incompatible model artifact: missing={incompatible.missing_keys}, " |
| f"unexpected={incompatible.unexpected_keys}" |
| ) |
| self.model.to(self.device).eval() |
|
|
| @staticmethod |
| def _poster_tensor(poster: Poster) -> torch.Tensor: |
| if isinstance(poster, Image.Image): |
| image = poster.convert("RGB").resize((224, 224), Image.Resampling.BICUBIC) |
| tensor = torch.from_numpy(np.asarray(image, dtype=np.float32).copy()).permute(2, 0, 1) |
| else: |
| tensor = torch.as_tensor(poster).detach().to(dtype=torch.float32, device="cpu") |
| if tensor.ndim != 3: |
| raise ValueError("poster tensor or array must have three dimensions") |
| if tensor.shape[0] not in (1, 3, 4) and tensor.shape[-1] in (1, 3, 4): |
| tensor = tensor.permute(2, 0, 1) |
| if tensor.shape[0] == 1: |
| tensor = tensor.expand(3, -1, -1) |
| elif tensor.shape[0] == 4: |
| tensor = tensor[:3] |
| tensor = torch.nn.functional.interpolate( |
| tensor.unsqueeze(0), size=(224, 224), mode="bicubic", align_corners=False |
| ).squeeze(0) |
| if tensor.max() > 1.0: |
| tensor = tensor / 255.0 |
| if tensor.min() >= 0.0: |
| tensor = tensor * 2.0 - 1.0 |
| return tensor.clamp(-1.0, 1.0) |
|
|
| def _tokenize(self, texts: Sequence[str]) -> tuple[torch.Tensor, torch.Tensor]: |
| tokens = self.tokenizer( |
| list(texts), padding=True, truncation=True, |
| max_length=int(self.config["text_max_length"]), return_tensors="pt", |
| ) |
| return tokens["input_ids"].to(self.device), tokens["attention_mask"].to(self.device) |
|
|
| @torch.inference_mode() |
| def vecs( |
| self, |
| items: Iterable[Union[ForgeItem, Mapping[str, Any]]], |
| *, |
| batch_size: int = 32, |
| ) -> list[ForgeVector]: |
| profiles = [ForgeItem.from_value(item) for item in items] |
| output: list[ForgeVector] = [] |
| for start in range(0, len(profiles), batch_size): |
| batch = profiles[start:start + batch_size] |
| descriptions_ua = [item.synopsis_ua or item.synopsis for item in batch] |
| descriptions_en = [item.synopsis or item.synopsis_ua for item in batch] |
| titles = [", ".join(filter(None, (item.title, item.native_title, *item.synonyms))) for item in batch] |
| ua_ids, ua_mask = self._tokenize(descriptions_ua) |
| en_ids, en_mask = self._tokenize(descriptions_en) |
| title_ids, title_mask = self._tokenize(titles) |
| ua = self.model.encode_text(ua_ids, ua_mask) |
| en = self.model.encode_text(en_ids, en_mask) |
| title_vectors = self.model.encode_text(title_ids, title_mask) |
| genres = torch.tensor([genre_indices(list(item.genres)) for item in batch], device=self.device) |
| metadata = torch.tensor([ |
| metadata_vector(item.content_type, item.year, item.score) for item in batch |
| ], dtype=torch.float32, device=self.device) |
| mask = torch.tensor([ |
| [float(bool(item.synopsis or item.synopsis_ua)), float(bool(item.title or item.native_title or item.synonyms)), |
| float(bool(item.genres)), 1.0, float(item.poster is not None)] |
| for item in batch |
| ], dtype=torch.float32, device=self.device) |
| posters = torch.zeros((len(batch), 3, 224, 224), device=self.device) |
| for index, item in enumerate(batch): |
| if item.poster is not None: |
| posters[index] = self._poster_tensor(item.poster).to(self.device) |
| embeddings = self.model(ua, en, title_vectors, genres, metadata, mask, posters) |
| output.extend(ForgeVector(row) for row in embeddings.cpu().numpy()) |
| return output |
|
|
| def vec(self, item: Optional[Union[ForgeItem, Mapping[str, Any]]] = None, **fields: Any) -> ForgeVector: |
| if item is not None and fields: |
| profile = replace(ForgeItem.from_value(item), **fields) |
| elif item is not None: |
| profile = ForgeItem.from_value(item) |
| else: |
| profile = ForgeItem(**fields) |
| return self.vecs([profile])[0] |
|
|
| def catalogue( |
| self, |
| items: Iterable[Union[ForgeItem, Mapping[str, Any]]], |
| *, |
| batch_size: int = 32, |
| ) -> ForgeCatalogue: |
| profiles = [ForgeItem.from_value(item) for item in items] |
| vectors = self.vecs(profiles, batch_size=batch_size) |
| embeddings = np.stack([vector._values for vector in vectors]) |
| return ForgeCatalogue(self, profiles, embeddings) |
|
|