Download src/hikka_forge/api.py from Lorg0n/hikka-forge2vec: direct link, hf CLI and curl.
- Browser
- Download file 14.1 kB
-
https://huggingface.co/Lorg0n/hikka-forge2vec/resolve/main/src/hikka_forge/api.py
- Command line
-
hf download hf://Lorg0n/hikka-forge2vec/src/hikka_forge/api.py
-
curl -L -o api.py https://huggingface.co/Lorg0n/hikka-forge2vec/resolve/main/src/hikka_forge/api.py
14.1 kB
| """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] | |
| 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 | |
| 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) | |
| def values(self) -> np.ndarray: | |
| return self._values.copy() | |
| 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})" | |
| 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() | |
| 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) | |
| 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) | |