feat(tagger): Settings picker to switch tagging model (WD / JoyTag)
Adds a "Tagging model" control to Settings → AI Workspace that calls get/set_tagger_model, alongside the existing acceleration/threshold/batch controls. Switching refreshes the model status so the download/ready row and the model name/description reflect the selected model, and the threshold hint shows the model-appropriate default (0.4 for JoyTag, 0.35 for WD). Frontend only; mirrors the existing acceleration toggle pattern. End-to-end JoyTag tagging still needs the model files on disk to verify.
This commit is contained in:
@@ -1,5 +1,5 @@
|
|||||||
import { useEffect, useMemo, useRef, useState } from "react";
|
import { useEffect, useMemo, useRef, useState } from "react";
|
||||||
import { AppTheme, CleanupOrphanedThumbnailsResult, DatabaseInfo, OrphanedThumbnailsInfo, TaggerAcceleration, TaggingQueueScope, VacuumResult, useGalleryStore } from "../store";
|
import { AppTheme, CleanupOrphanedThumbnailsResult, DatabaseInfo, OrphanedThumbnailsInfo, TaggerAcceleration, TaggerModel, TaggingQueueScope, VacuumResult, useGalleryStore } from "../store";
|
||||||
import { FfmpegStatusRow } from "./onboarding/StepWelcome";
|
import { FfmpegStatusRow } from "./onboarding/StepWelcome";
|
||||||
import { ThemedDropdown } from "./ThemedDropdown";
|
import { ThemedDropdown } from "./ThemedDropdown";
|
||||||
import { getChangelogForVersion } from "../changelog";
|
import { getChangelogForVersion } from "../changelog";
|
||||||
@@ -100,6 +100,44 @@ function ScopeButton({ scope, current, onSelect, children }: {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Display metadata for each selectable tagging model.
|
||||||
|
const TAGGER_MODELS: Record<TaggerModel, { name: string; tab: string; description: string }> = {
|
||||||
|
wd: {
|
||||||
|
name: "WD SwinV2 Tagger v3",
|
||||||
|
tab: "WD (anime)",
|
||||||
|
description:
|
||||||
|
"Anime-focused vision model by SmilingWolf. Generates booru-style tags with configurable confidence thresholds.",
|
||||||
|
},
|
||||||
|
joytag: {
|
||||||
|
name: "JoyTag",
|
||||||
|
tab: "JoyTag (general)",
|
||||||
|
description:
|
||||||
|
"Booru-schema tagger that also handles photographic content and is strong on NSFW concepts. The explicitness rating is derived from its tags.",
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
function TaggerModelButton({ model, current, onSelect, children }: {
|
||||||
|
model: TaggerModel;
|
||||||
|
current: TaggerModel;
|
||||||
|
onSelect: (model: TaggerModel) => void;
|
||||||
|
children: React.ReactNode;
|
||||||
|
}) {
|
||||||
|
const active = model === current;
|
||||||
|
return (
|
||||||
|
<button
|
||||||
|
type="button"
|
||||||
|
className={`rounded-md border px-3 py-1.5 text-xs transition-colors ${
|
||||||
|
active
|
||||||
|
? "border-emerald-400/35 bg-emerald-500/15 text-emerald-200 light-theme:border-emerald-600/50 light-theme:bg-emerald-100 light-theme:text-emerald-700"
|
||||||
|
: "border-transparent text-gray-500 hover:bg-white/[0.06] hover:text-gray-200 light-theme:text-gray-600 light-theme:hover:bg-gray-900 light-theme:hover:text-white"
|
||||||
|
}`}
|
||||||
|
onClick={() => onSelect(model)}
|
||||||
|
>
|
||||||
|
{children}
|
||||||
|
</button>
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
function TaggerAccelerationButton({ acceleration, current, onSelect, children }: {
|
function TaggerAccelerationButton({ acceleration, current, onSelect, children }: {
|
||||||
acceleration: TaggerAcceleration;
|
acceleration: TaggerAcceleration;
|
||||||
current: TaggerAcceleration;
|
current: TaggerAcceleration;
|
||||||
@@ -129,6 +167,8 @@ export function SettingsModal() {
|
|||||||
const [taggerClearing, setTaggerClearing] = useState(false);
|
const [taggerClearing, setTaggerClearing] = useState(false);
|
||||||
const [taggerAccelerationSaving, setTaggerAccelerationSaving] = useState(false);
|
const [taggerAccelerationSaving, setTaggerAccelerationSaving] = useState(false);
|
||||||
const [taggerAccelerationError, setTaggerAccelerationError] = useState<string | null>(null);
|
const [taggerAccelerationError, setTaggerAccelerationError] = useState<string | null>(null);
|
||||||
|
const [taggerModelSwitching, setTaggerModelSwitching] = useState(false);
|
||||||
|
const [taggerModelSwitchError, setTaggerModelSwitchError] = useState<string | null>(null);
|
||||||
const [taggerThresholdDraft, setTaggerThresholdDraft] = useState<string | null>(null);
|
const [taggerThresholdDraft, setTaggerThresholdDraft] = useState<string | null>(null);
|
||||||
const [taggerThresholdSaving, setTaggerThresholdSaving] = useState(false);
|
const [taggerThresholdSaving, setTaggerThresholdSaving] = useState(false);
|
||||||
const [taggerThresholdError, setTaggerThresholdError] = useState<string | null>(null);
|
const [taggerThresholdError, setTaggerThresholdError] = useState<string | null>(null);
|
||||||
@@ -173,6 +213,9 @@ export function SettingsModal() {
|
|||||||
const deleteTaggerModel = useGalleryStore((state) => state.deleteTaggerModel);
|
const deleteTaggerModel = useGalleryStore((state) => state.deleteTaggerModel);
|
||||||
const loadTaggerAcceleration = useGalleryStore((state) => state.loadTaggerAcceleration);
|
const loadTaggerAcceleration = useGalleryStore((state) => state.loadTaggerAcceleration);
|
||||||
const setTaggerAcceleration = useGalleryStore((state) => state.setTaggerAcceleration);
|
const setTaggerAcceleration = useGalleryStore((state) => state.setTaggerAcceleration);
|
||||||
|
const taggerModel = useGalleryStore((state) => state.taggerModel);
|
||||||
|
const loadTaggerModel = useGalleryStore((state) => state.loadTaggerModel);
|
||||||
|
const setTaggerModel = useGalleryStore((state) => state.setTaggerModel);
|
||||||
const loadTaggerThreshold = useGalleryStore((state) => state.loadTaggerThreshold);
|
const loadTaggerThreshold = useGalleryStore((state) => state.loadTaggerThreshold);
|
||||||
const setTaggerThreshold = useGalleryStore((state) => state.setTaggerThreshold);
|
const setTaggerThreshold = useGalleryStore((state) => state.setTaggerThreshold);
|
||||||
const loadTaggerBatchSize = useGalleryStore((state) => state.loadTaggerBatchSize);
|
const loadTaggerBatchSize = useGalleryStore((state) => state.loadTaggerBatchSize);
|
||||||
@@ -210,6 +253,7 @@ export function SettingsModal() {
|
|||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (!settingsOpen) return;
|
if (!settingsOpen) return;
|
||||||
void loadTaggerModelStatus();
|
void loadTaggerModelStatus();
|
||||||
|
void loadTaggerModel();
|
||||||
void loadTaggerAcceleration();
|
void loadTaggerAcceleration();
|
||||||
void loadTaggerThreshold();
|
void loadTaggerThreshold();
|
||||||
void loadTaggerBatchSize();
|
void loadTaggerBatchSize();
|
||||||
@@ -221,7 +265,7 @@ export function SettingsModal() {
|
|||||||
};
|
};
|
||||||
window.addEventListener("keydown", handleKeyDown);
|
window.addEventListener("keydown", handleKeyDown);
|
||||||
return () => window.removeEventListener("keydown", handleKeyDown);
|
return () => window.removeEventListener("keydown", handleKeyDown);
|
||||||
}, [settingsOpen, loadTaggerModelStatus, loadTaggerAcceleration, loadTaggerThreshold, loadTaggerBatchSize, loadTaggingQueueScope, loadTaggingQueueFolderIds, setSettingsOpen]);
|
}, [settingsOpen, loadTaggerModelStatus, loadTaggerModel, loadTaggerAcceleration, loadTaggerThreshold, loadTaggerBatchSize, loadTaggingQueueScope, loadTaggingQueueFolderIds, setSettingsOpen]);
|
||||||
|
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (!settingsOpen || activeSection !== "general") return;
|
if (!settingsOpen || activeSection !== "general") return;
|
||||||
@@ -365,10 +409,39 @@ export function SettingsModal() {
|
|||||||
{activeSection === "workspace" ? (
|
{activeSection === "workspace" ? (
|
||||||
<div className="mt-8 space-y-9">
|
<div className="mt-8 space-y-9">
|
||||||
<SettingsGroup title="Model">
|
<SettingsGroup title="Model">
|
||||||
|
<SettingsItem label="Tagging model" description="Choose which model generates tags. JoyTag suits photos and NSFW; WD is anime-focused. Switching may require a one-time download.">
|
||||||
|
<div className="flex flex-col items-end gap-1.5">
|
||||||
|
<div className="flex rounded-lg border border-white/[0.07] p-0.5 light-theme:border-gray-300/80">
|
||||||
|
{(["wd", "joytag"] as const).map((model) => (
|
||||||
|
<TaggerModelButton
|
||||||
|
key={model}
|
||||||
|
model={model}
|
||||||
|
current={taggerModel}
|
||||||
|
onSelect={(nextModel) => {
|
||||||
|
if (nextModel === taggerModel) return;
|
||||||
|
setTaggerModelSwitching(true);
|
||||||
|
setTaggerModelSwitchError(null);
|
||||||
|
void setTaggerModel(nextModel)
|
||||||
|
.catch((error: unknown) => setTaggerModelSwitchError(String(error)))
|
||||||
|
.finally(() => setTaggerModelSwitching(false));
|
||||||
|
}}
|
||||||
|
>
|
||||||
|
{TAGGER_MODELS[model].tab}
|
||||||
|
</TaggerModelButton>
|
||||||
|
))}
|
||||||
|
</div>
|
||||||
|
{taggerModelSwitchError ? (
|
||||||
|
<p className="text-[11px] text-amber-300">{taggerModelSwitchError}</p>
|
||||||
|
) : (
|
||||||
|
<p className="text-[11px] text-gray-600">{taggerModelSwitching ? "Switching..." : `Current: ${TAGGER_MODELS[taggerModel].name}`}</p>
|
||||||
|
)}
|
||||||
|
</div>
|
||||||
|
</SettingsItem>
|
||||||
|
|
||||||
<SettingsItem
|
<SettingsItem
|
||||||
label={
|
label={
|
||||||
<>
|
<>
|
||||||
WD SwinV2 Tagger v3{" "}
|
{TAGGER_MODELS[taggerModel].name}{" "}
|
||||||
<span className="ml-1.5 align-middle">
|
<span className="ml-1.5 align-middle">
|
||||||
<StatusPill tone={taggerReady ? "ready" : taggerModelPreparing ? "busy" : "muted"}>
|
<StatusPill tone={taggerReady ? "ready" : taggerModelPreparing ? "busy" : "muted"}>
|
||||||
{taggerReady ? "Installed" : taggerModelPreparing ? "Preparing" : "Not installed"}
|
{taggerReady ? "Installed" : taggerModelPreparing ? "Preparing" : "Not installed"}
|
||||||
@@ -376,7 +449,7 @@ export function SettingsModal() {
|
|||||||
</span>
|
</span>
|
||||||
</>
|
</>
|
||||||
}
|
}
|
||||||
description="Anime-focused vision model by SmilingWolf. Generates booru-style tags with configurable confidence thresholds."
|
description={TAGGER_MODELS[taggerModel].description}
|
||||||
>
|
>
|
||||||
<div className="flex items-center gap-2">
|
<div className="flex items-center gap-2">
|
||||||
{taggerReady ? (
|
{taggerReady ? (
|
||||||
@@ -469,7 +542,7 @@ export function SettingsModal() {
|
|||||||
{taggerThresholdError ? (
|
{taggerThresholdError ? (
|
||||||
<p className="text-[11px] text-amber-300">{taggerThresholdError}</p>
|
<p className="text-[11px] text-amber-300">{taggerThresholdError}</p>
|
||||||
) : (
|
) : (
|
||||||
<p className="text-[11px] text-gray-600">{taggerThresholdSaving ? "Saving..." : "Default: 0.35"}</p>
|
<p className="text-[11px] text-gray-600">{taggerThresholdSaving ? "Saving..." : `Default: ${taggerModel === "joytag" ? "0.4" : "0.35"}`}</p>
|
||||||
)}
|
)}
|
||||||
</div>
|
</div>
|
||||||
</SettingsItem>
|
</SettingsItem>
|
||||||
|
|||||||
@@ -51,6 +51,7 @@ export type SearchCommand = "filename" | "semantic" | "tag";
|
|||||||
export type CaptionAcceleration = "auto" | "cpu" | "directml";
|
export type CaptionAcceleration = "auto" | "cpu" | "directml";
|
||||||
export type CaptionDetail = "short" | "detailed" | "paragraph";
|
export type CaptionDetail = "short" | "detailed" | "paragraph";
|
||||||
export type TaggerAcceleration = "auto" | "cpu" | "directml";
|
export type TaggerAcceleration = "auto" | "cpu" | "directml";
|
||||||
|
export type TaggerModel = "wd" | "joytag";
|
||||||
export type AiRating = "general" | "sensitive" | "questionable" | "explicit";
|
export type AiRating = "general" | "sensitive" | "questionable" | "explicit";
|
||||||
export type TaggingQueueScope = "all" | "selected";
|
export type TaggingQueueScope = "all" | "selected";
|
||||||
export type SimilarScope = "all_media" | "current_folder" | "current_album";
|
export type SimilarScope = "all_media" | "current_folder" | "current_album";
|
||||||
@@ -424,6 +425,7 @@ interface GalleryState {
|
|||||||
taggerModelPreparing: boolean;
|
taggerModelPreparing: boolean;
|
||||||
taggerModelError: string | null;
|
taggerModelError: string | null;
|
||||||
taggerModelProgress: TaggerModelProgress | null;
|
taggerModelProgress: TaggerModelProgress | null;
|
||||||
|
taggerModel: TaggerModel;
|
||||||
taggerAcceleration: TaggerAcceleration;
|
taggerAcceleration: TaggerAcceleration;
|
||||||
taggerThreshold: number;
|
taggerThreshold: number;
|
||||||
taggerBatchSize: number;
|
taggerBatchSize: number;
|
||||||
@@ -551,6 +553,8 @@ interface GalleryState {
|
|||||||
deleteTaggerModel: () => Promise<void>;
|
deleteTaggerModel: () => Promise<void>;
|
||||||
loadTaggerAcceleration: () => Promise<void>;
|
loadTaggerAcceleration: () => Promise<void>;
|
||||||
setTaggerAcceleration: (acceleration: TaggerAcceleration) => Promise<void>;
|
setTaggerAcceleration: (acceleration: TaggerAcceleration) => Promise<void>;
|
||||||
|
loadTaggerModel: () => Promise<void>;
|
||||||
|
setTaggerModel: (model: TaggerModel) => Promise<void>;
|
||||||
loadTaggerThreshold: () => Promise<void>;
|
loadTaggerThreshold: () => Promise<void>;
|
||||||
setTaggerThreshold: (threshold: number) => Promise<void>;
|
setTaggerThreshold: (threshold: number) => Promise<void>;
|
||||||
loadTaggerBatchSize: () => Promise<void>;
|
loadTaggerBatchSize: () => Promise<void>;
|
||||||
@@ -902,6 +906,7 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
|
|||||||
taggerModelPreparing: false,
|
taggerModelPreparing: false,
|
||||||
taggerModelError: null,
|
taggerModelError: null,
|
||||||
taggerModelProgress: null,
|
taggerModelProgress: null,
|
||||||
|
taggerModel: "wd",
|
||||||
taggerAcceleration: "auto",
|
taggerAcceleration: "auto",
|
||||||
taggerThreshold: 0.35,
|
taggerThreshold: 0.35,
|
||||||
taggerBatchSize: 8,
|
taggerBatchSize: 8,
|
||||||
@@ -2158,6 +2163,29 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
|
|||||||
set({ taggerAcceleration, taggerRuntimeProbe: null });
|
set({ taggerAcceleration, taggerRuntimeProbe: null });
|
||||||
},
|
},
|
||||||
|
|
||||||
|
loadTaggerModel: async () => {
|
||||||
|
try {
|
||||||
|
const taggerModel = await invoke<TaggerModel>("get_tagger_model");
|
||||||
|
set({ taggerModel });
|
||||||
|
} catch (error) {
|
||||||
|
set({ taggerModelError: String(error) });
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
setTaggerModel: async (model) => {
|
||||||
|
const taggerModel = await invoke<TaggerModel>("set_tagger_model", {
|
||||||
|
params: { model },
|
||||||
|
});
|
||||||
|
// Switching models changes which files are required on disk, so refresh the
|
||||||
|
// status too — the download/ready UI must reflect the newly-selected model.
|
||||||
|
try {
|
||||||
|
const taggerModelStatus = await invoke<TaggerModelStatus>("get_tagger_model_status");
|
||||||
|
set({ taggerModel, taggerModelStatus, taggerModelError: null, taggerRuntimeProbe: null });
|
||||||
|
} catch (error) {
|
||||||
|
set({ taggerModel, taggerRuntimeProbe: null, taggerModelError: String(error) });
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
loadTaggerThreshold: async () => {
|
loadTaggerThreshold: async () => {
|
||||||
try {
|
try {
|
||||||
const taggerThreshold = await invoke<number>("get_tagger_threshold");
|
const taggerThreshold = await invoke<number>("get_tagger_threshold");
|
||||||
|
|||||||
Reference in New Issue
Block a user