DageBjorne
Package project as pip-installable augmenator.
7025ca1
Raw
History Blame Contribute Delete
3.72 kB
import numpy as np
from sentence_transformers import SentenceTransformer
from augmenator.keyword_catalog import AUGMENT_KEYWORDS, AugmentKeyword
MODEL_ID = "sentence-transformers/all-MiniLM-L6-v2"
SIMILARITY_THRESHOLD = 0.38
RELATIVE_MARGIN = 0.25
MAX_SELECTIONS = 5
MUTUAL_EXCLUSION_GROUPS: tuple[tuple[str, ...], ...] = (
("brighten", "darken"),
("contrast_up", "contrast_down"),
("warmer", "cooler"),
("saturate", "desaturate"),
("tint_red", "tint_green", "tint_blue"),
("cover_random", "cover_avoid_text", "cover_and_cutout", "cover_and_cutout_avoid_text"),
("add_cutout_transparent", "add_cutout_solid", "add_cutout", "add_cutout_avoid_text"),
(
"rotate",
"rotate_90_random",
"rotate_left",
"rotate_right",
"rotate_180",
"flip",
"flip_vertical",
),
("gamma_up", "gamma_down"),
(
"style_candy",
"style_mosaic",
"style_rain_princess",
"style_udnie",
"style_pointilism",
"style_starry_night",
"style_sketch",
"style_random",
),
)
_model: SentenceTransformer | None = None
_phrase_embeddings: dict[str, np.ndarray] = {}
def _load() -> SentenceTransformer:
global _model
if _model is None:
_model = SentenceTransformer(MODEL_ID)
return _model
def _build_phrase_cache(model: SentenceTransformer) -> None:
global _phrase_embeddings
if _phrase_embeddings:
return
for keyword in AUGMENT_KEYWORDS:
embeddings = model.encode(
list(keyword.phrases),
convert_to_numpy=True,
normalize_embeddings=True,
show_progress_bar=False,
)
_phrase_embeddings[keyword.id] = embeddings
def warmup() -> None:
model = _load()
_build_phrase_cache(model)
def _cosine_scores(instruction_embedding: np.ndarray) -> list[tuple[AugmentKeyword, float]]:
scored: list[tuple[AugmentKeyword, float]] = []
for keyword in AUGMENT_KEYWORDS:
phrase_embs = _phrase_embeddings[keyword.id]
similarities = phrase_embs @ instruction_embedding
score = float(np.max(similarities))
scored.append((keyword, score))
scored.sort(key=lambda item: item[1], reverse=True)
return scored
def _apply_mutual_exclusion(
matches: list[tuple[AugmentKeyword, float]],
) -> list[tuple[AugmentKeyword, float]]:
selected_ids: set[str] = set()
result: list[tuple[AugmentKeyword, float]] = []
for keyword, score in matches:
skip = False
for group in MUTUAL_EXCLUSION_GROUPS:
if keyword.id not in group:
continue
if any(member in selected_ids for member in group):
skip = True
break
if skip:
continue
selected_ids.add(keyword.id)
result.append((keyword, score))
return result
def score_keywords(instruction: str) -> list[tuple[AugmentKeyword, float]]:
warmup()
model = _load()
instruction_embedding = model.encode(
instruction,
convert_to_numpy=True,
normalize_embeddings=True,
show_progress_bar=False,
)
return _cosine_scores(instruction_embedding)
def select_keywords(instruction: str) -> list[tuple[AugmentKeyword, float]]:
scored = score_keywords(instruction)
if not scored:
return []
top_score = scored[0][1]
min_relative = top_score - RELATIVE_MARGIN
above_threshold = [
(keyword, score)
for keyword, score in scored
if score >= SIMILARITY_THRESHOLD and score >= min_relative
]
deduped = _apply_mutual_exclusion(above_threshold)
return deduped[:MAX_SELECTIONS]