perf(tagger): chunked, yielding inference so tagging stops freezing the UI
The tagging worker claimed a batch from the DB but ran the model one image at a time, so `tagger_batch_size` had no effect on inference. Batching the whole claim into a single DirectML forward pass fixed that but introduced a worse problem: on a shared GPU each wide dispatch locks the device (and the WebView2 compositor with it) for 1-3.7s, freezing the entire app while tagging runs — worst of all on first launch with a cold graph compile. The WD model is compute-bound here (~50-230ms/image of actual GPU work), so a wide batch buys almost no throughput; it only lengthens each uninterruptible GPU lock. So decouple DB claim size from GPU granularity: claim `tagger_batch_size` for DB efficiency, but feed the GPU in TAGGER_INFER_CHUNK (4) images per forward pass with a brief yield between chunks. Each dispatch now lasts ~200-900ms, the UI gets windows to paint, peak decode memory is bounded, and the cold compile is for a smaller shape. As a bonus the wide-batch throughput spikes disappear — steady ~1.6-1.8s/16 vs the old 0.8-3.7s swings. - run_batch: parallel (rayon) decode + single forward pass per chunk, with per-image fallback and decode failures kept attached to their input slot. - Worker iterates source_paths.chunks(TAGGER_INFER_CHUNK), yielding 40ms between chunks; outputs zip back to jobs 1:1, write tx unchanged. - Remove now-dead WdTagger::run() (fallback uses infer_one directly).
This commit is contained in:
@@ -93,6 +93,12 @@ aims to follow [Semantic Versioning](https://semver.org/spec/v2.0.0.html)
|
|||||||
- **Explore polish** — the Tag Cloud now uses the in-app tooltip instead of the
|
- **Explore polish** — the Tag Cloud now uses the in-app tooltip instead of the
|
||||||
native browser one (and reads "1 image", not "1 images"), and the folder-scope
|
native browser one (and reads "1 image", not "1 images"), and the folder-scope
|
||||||
dropdown no longer opens behind the cluster cards.
|
dropdown no longer opens behind the cluster cards.
|
||||||
|
- **AI tagging no longer freezes the app** — tagging now runs inference in small
|
||||||
|
GPU micro-batches with a brief yield between them, instead of one wide batch
|
||||||
|
that monopolised the GPU (and with it the whole UI) for seconds at a time. The
|
||||||
|
app stays responsive while tagging runs, throughput is steadier (the old wide
|
||||||
|
batches caused periodic slowdowns), and the first batch after launch starts
|
||||||
|
faster.
|
||||||
|
|
||||||
## [0.1.1] — 2026-06-23
|
## [0.1.1] — 2026-06-23
|
||||||
|
|
||||||
|
|||||||
+43
-10
@@ -1248,13 +1248,16 @@ fn process_tagging_batch(
|
|||||||
|
|
||||||
// Exclude actively-indexing folders for the same reason as the other
|
// Exclude actively-indexing folders for the same reason as the other
|
||||||
// workers: don't compete with a running scan.
|
// workers: don't compete with a running scan.
|
||||||
|
let batch_started_at = Instant::now();
|
||||||
let mut excluded_folders = paused_folder_ids("tagging");
|
let mut excluded_folders = paused_folder_ids("tagging");
|
||||||
excluded_folders.extend(active_indexing_folders());
|
excluded_folders.extend(active_indexing_folders());
|
||||||
let batch_size = crate::tagger::tagger_batch_size(app_data_dir);
|
let batch_size = crate::tagger::tagger_batch_size(app_data_dir);
|
||||||
|
let claim_started_at = Instant::now();
|
||||||
let jobs = with_db_write_lock(|| {
|
let jobs = with_db_write_lock(|| {
|
||||||
let mut conn = pool.get()?;
|
let mut conn = pool.get()?;
|
||||||
db::claim_tagging_jobs(&mut conn, &excluded_folders, batch_size)
|
db::claim_tagging_jobs(&mut conn, &excluded_folders, batch_size)
|
||||||
})?;
|
})?;
|
||||||
|
let claim_elapsed = claim_started_at.elapsed();
|
||||||
|
|
||||||
if jobs.is_empty() {
|
if jobs.is_empty() {
|
||||||
return Ok(false);
|
return Ok(false);
|
||||||
@@ -1288,21 +1291,41 @@ fn process_tagging_batch(
|
|||||||
.as_mut()
|
.as_mut()
|
||||||
.expect("tagger should be initialized before tagging batch processing");
|
.expect("tagger should be initialized before tagging batch processing");
|
||||||
|
|
||||||
let tag_results = jobs
|
// Resolve each job's source image (AVIF can't be decoded directly, so it
|
||||||
|
// falls back to its thumbnail), then tag the batch in small chunks (below).
|
||||||
|
let source_paths: Vec<PathBuf> = jobs
|
||||||
.iter()
|
.iter()
|
||||||
.map(|job| {
|
.map(|job| {
|
||||||
let source_path = if is_avif_path(Path::new(&job.path)) {
|
if is_avif_path(Path::new(&job.path)) {
|
||||||
job.thumbnail_path.as_deref().unwrap_or(&job.path)
|
PathBuf::from(job.thumbnail_path.as_deref().unwrap_or(&job.path))
|
||||||
} else {
|
} else {
|
||||||
&job.path
|
PathBuf::from(&job.path)
|
||||||
};
|
}
|
||||||
(
|
|
||||||
job.clone(),
|
|
||||||
tagger_ref.run(Path::new(source_path), tagger::DEFAULT_MAX_TAGS),
|
|
||||||
)
|
|
||||||
})
|
})
|
||||||
.collect::<Vec<_>>();
|
.collect();
|
||||||
|
|
||||||
|
// Tag in small micro-batches instead of one wide forward pass. On a shared
|
||||||
|
// GPU every DirectML dispatch blocks the WebView2 compositor for its whole
|
||||||
|
// duration, so a 16-wide batch freezes the UI for seconds; small chunks keep
|
||||||
|
// each GPU lock short (and bound peak decode memory) while a brief yield
|
||||||
|
// between them lets the UI and other workers grab the GPU/CPU. The model is
|
||||||
|
// compute-bound here, so this costs almost no throughput.
|
||||||
|
let infer_started_at = Instant::now();
|
||||||
|
let mut outputs = Vec::with_capacity(source_paths.len());
|
||||||
|
let mut chunks = source_paths.chunks(tagger::TAGGER_INFER_CHUNK).peekable();
|
||||||
|
while let Some(chunk) = chunks.next() {
|
||||||
|
outputs.extend(tagger_ref.run_batch(chunk, tagger::DEFAULT_MAX_TAGS));
|
||||||
|
if chunks.peek().is_some() {
|
||||||
|
std::thread::sleep(std::time::Duration::from_millis(
|
||||||
|
tagger::TAGGER_INFER_YIELD_MS,
|
||||||
|
));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let infer_elapsed = infer_started_at.elapsed();
|
||||||
|
|
||||||
|
let tag_results = jobs.iter().cloned().zip(outputs).collect::<Vec<_>>();
|
||||||
|
|
||||||
|
let write_started_at = Instant::now();
|
||||||
let updated_images = with_db_write_lock(|| {
|
let updated_images = with_db_write_lock(|| {
|
||||||
let mut conn = pool.get()?;
|
let mut conn = pool.get()?;
|
||||||
let tx = conn.transaction()?;
|
let tx = conn.transaction()?;
|
||||||
@@ -1354,6 +1377,7 @@ fn process_tagging_batch(
|
|||||||
db::requeue_tagging_jobs(&conn, &image_ids)
|
db::requeue_tagging_jobs(&conn, &image_ids)
|
||||||
});
|
});
|
||||||
})?;
|
})?;
|
||||||
|
let write_elapsed = write_started_at.elapsed();
|
||||||
|
|
||||||
if !updated_images.is_empty() {
|
if !updated_images.is_empty() {
|
||||||
let folder_ids = updated_images
|
let folder_ids = updated_images
|
||||||
@@ -1369,6 +1393,15 @@ fn process_tagging_batch(
|
|||||||
emit_folder_job_progress(app, pool, &folder_ids.into_iter().collect::<Vec<_>>(), true);
|
emit_folder_job_progress(app, pool, &folder_ids.into_iter().collect::<Vec<_>>(), true);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log::info!(
|
||||||
|
"Tagging batch timing: {} items, claim {:?}, tag {:?}, write {:?}, total {:?}",
|
||||||
|
jobs.len(),
|
||||||
|
claim_elapsed,
|
||||||
|
infer_elapsed,
|
||||||
|
write_elapsed,
|
||||||
|
batch_started_at.elapsed()
|
||||||
|
);
|
||||||
|
|
||||||
Ok(true)
|
Ok(true)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+151
-39
@@ -4,6 +4,7 @@ use image::{imageops::FilterType, DynamicImage, ImageReader};
|
|||||||
use ort::ep;
|
use ort::ep;
|
||||||
use ort::session::{builder::GraphOptimizationLevel, Session};
|
use ort::session::{builder::GraphOptimizationLevel, Session};
|
||||||
use ort::value::Tensor;
|
use ort::value::Tensor;
|
||||||
|
use rayon::prelude::*;
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use std::path::{Path, PathBuf};
|
use std::path::{Path, PathBuf};
|
||||||
use std::sync::atomic::{AtomicBool, Ordering};
|
use std::sync::atomic::{AtomicBool, Ordering};
|
||||||
@@ -37,6 +38,19 @@ const RATING_CATEGORY: u8 = 9;
|
|||||||
pub const DEFAULT_THRESHOLD: f32 = 0.35;
|
pub const DEFAULT_THRESHOLD: f32 = 0.35;
|
||||||
pub const DEFAULT_MAX_TAGS: usize = 30;
|
pub const DEFAULT_MAX_TAGS: usize = 30;
|
||||||
|
|
||||||
|
/// How many images are fed to the GPU in a single forward pass, regardless of
|
||||||
|
/// how many the worker claims from the DB at once. The WD model is compute-
|
||||||
|
/// bound on a shared GPU, so a wide batch buys little throughput but holds the
|
||||||
|
/// GPU (and therefore the WebView2 compositor) hostage for its whole duration,
|
||||||
|
/// freezing the UI. Small chunks keep each DirectML dispatch short — lower this
|
||||||
|
/// for a smoother UI under heavy tagging, raise it for marginally higher
|
||||||
|
/// throughput on a dedicated GPU.
|
||||||
|
pub const TAGGER_INFER_CHUNK: usize = 4;
|
||||||
|
|
||||||
|
/// Brief pause between inference chunks so the compositor and other workers can
|
||||||
|
/// claim the GPU/CPU between dispatches. Negligible against per-chunk inference.
|
||||||
|
pub const TAGGER_INFER_YIELD_MS: u64 = 40;
|
||||||
|
|
||||||
/// Set to `true` by `set_tagger_acceleration` so the tagging worker loop
|
/// Set to `true` by `set_tagger_acceleration` so the tagging worker loop
|
||||||
/// knows to drop its cached `WdTagger` and reload with the new EP.
|
/// knows to drop its cached `WdTagger` and reload with the new EP.
|
||||||
pub static TAGGER_SESSION_DIRTY: AtomicBool = AtomicBool::new(false);
|
pub static TAGGER_SESSION_DIRTY: AtomicBool = AtomicBool::new(false);
|
||||||
@@ -456,33 +470,34 @@ impl WdTagger {
|
|||||||
|
|
||||||
let session = create_tagger_session(&model_path, acceleration)?;
|
let session = create_tagger_session(&model_path, acceleration)?;
|
||||||
|
|
||||||
// Determine the input spatial size from the ONNX model graph.
|
// Determine the input spatial size (and batch axis, for diagnostics) from
|
||||||
// WD v3 models use (1, H, W, 3) where H == W, typically 448.
|
// the ONNX model graph. WD v3 models use (N, H, W, 3) with H == W,
|
||||||
let input_size = {
|
// typically 448, and a dynamic batch dimension (reported as -1/0) which
|
||||||
|
// batched inference relies on.
|
||||||
|
let (input_size, batch_axis) = {
|
||||||
let inputs = session.inputs();
|
let inputs = session.inputs();
|
||||||
let dim = inputs
|
inputs
|
||||||
.first()
|
.first()
|
||||||
.and_then(|inp| {
|
.and_then(|inp| {
|
||||||
if let ort::value::ValueType::Tensor { shape, .. } = inp.dtype() {
|
if let ort::value::ValueType::Tensor { shape, .. } = inp.dtype() {
|
||||||
shape
|
let size = shape.get(1).copied().filter(|&d| d > 0).map(|d| d as usize);
|
||||||
.get(1)
|
Some((size.unwrap_or(448), shape.first().copied()))
|
||||||
.and_then(|&d| if d > 0 { Some(d as usize) } else { None })
|
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
.unwrap_or(448);
|
.unwrap_or((448, None))
|
||||||
dim
|
|
||||||
};
|
};
|
||||||
|
|
||||||
let labels = load_labels(&labels_path)?;
|
let labels = load_labels(&labels_path)?;
|
||||||
|
|
||||||
log::info!(
|
log::info!(
|
||||||
"WD tagger loaded in {:?} ({} labels, input {}x{}, {:?} acceleration)",
|
"WD tagger loaded in {:?} ({} labels, input {}x{}, batch axis {:?}, {:?} acceleration)",
|
||||||
started_at.elapsed(),
|
started_at.elapsed(),
|
||||||
labels.len(),
|
labels.len(),
|
||||||
input_size,
|
input_size,
|
||||||
input_size,
|
input_size,
|
||||||
|
batch_axis,
|
||||||
acceleration,
|
acceleration,
|
||||||
);
|
);
|
||||||
|
|
||||||
@@ -494,14 +509,96 @@ impl WdTagger {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn run(&mut self, image_path: &Path, max_tags: usize) -> Result<TaggerOutput> {
|
/// Tag a batch of images in a single forward pass. Returns one result per
|
||||||
let started_at = Instant::now();
|
/// input path, in the same order. Per-image decode failures — and, if the
|
||||||
|
/// batched inference itself fails (e.g. a model pinned to batch=1), a
|
||||||
|
/// per-image inference fallback — are reflected in the individual `Result`s.
|
||||||
|
pub fn run_batch(
|
||||||
|
&mut self,
|
||||||
|
image_paths: &[PathBuf],
|
||||||
|
max_tags: usize,
|
||||||
|
) -> Vec<Result<TaggerOutput>> {
|
||||||
|
if image_paths.is_empty() {
|
||||||
|
return Vec::new();
|
||||||
|
}
|
||||||
|
|
||||||
let image_array = preprocess_image(image_path, self.input_size)?;
|
let started_at = Instant::now();
|
||||||
let batch_size = 1usize;
|
let input_size = self.input_size;
|
||||||
|
let stride = input_size * input_size * 3;
|
||||||
|
|
||||||
|
// Decode + preprocess every image in parallel, keeping results aligned to
|
||||||
|
// the inputs so a failure stays attached to its own path.
|
||||||
|
let preprocessed: Vec<Result<Vec<f32>>> = image_paths
|
||||||
|
.par_iter()
|
||||||
|
.map(|path| preprocess_image(path, input_size))
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
// Pack the successfully-decoded images into one contiguous buffer for a
|
||||||
|
// single batched inference, remembering which input slot each came from.
|
||||||
|
let mut batch_slots: Vec<usize> = Vec::new();
|
||||||
|
let mut batch_pixels: Vec<f32> = Vec::with_capacity(image_paths.len() * stride);
|
||||||
|
for (i, result) in preprocessed.iter().enumerate() {
|
||||||
|
if let Ok(pixels) = result {
|
||||||
|
batch_slots.push(i);
|
||||||
|
batch_pixels.extend_from_slice(pixels);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Seed each slot with its decode error; inference overwrites the ones
|
||||||
|
// that decoded. The placeholder error never survives a decoded slot.
|
||||||
|
let mut results: Vec<Result<TaggerOutput>> = preprocessed
|
||||||
|
.into_iter()
|
||||||
|
.map(|r| match r {
|
||||||
|
Ok(_) => Err(anyhow::anyhow!(
|
||||||
|
"tagging inference did not run for this image"
|
||||||
|
)),
|
||||||
|
Err(error) => Err(error),
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
if batch_slots.is_empty() {
|
||||||
|
return results; // nothing decoded
|
||||||
|
}
|
||||||
|
|
||||||
|
match self.infer_batch(&batch_pixels, batch_slots.len(), max_tags) {
|
||||||
|
Ok(outputs) => {
|
||||||
|
for (&slot, output) in batch_slots.iter().zip(outputs) {
|
||||||
|
results[slot] = Ok(output);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Err(batch_error) => {
|
||||||
|
log::warn!(
|
||||||
|
"WD tagger batch inference failed for {} images, falling back to per-image: {batch_error}",
|
||||||
|
batch_slots.len()
|
||||||
|
);
|
||||||
|
for (k, &slot) in batch_slots.iter().enumerate() {
|
||||||
|
let pixels = batch_pixels[k * stride..(k + 1) * stride].to_vec();
|
||||||
|
results[slot] = self.infer_one(pixels, max_tags);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
log::debug!(
|
||||||
|
"WD tagger batch: {} images ({} decoded) in {:?}",
|
||||||
|
image_paths.len(),
|
||||||
|
batch_slots.len(),
|
||||||
|
started_at.elapsed(),
|
||||||
|
);
|
||||||
|
|
||||||
|
results
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Run `count` already-preprocessed images (contiguous `[count, H, W, 3]`
|
||||||
|
/// pixels) through the model in one forward pass.
|
||||||
|
fn infer_batch(
|
||||||
|
&mut self,
|
||||||
|
pixels: &[f32],
|
||||||
|
count: usize,
|
||||||
|
max_tags: usize,
|
||||||
|
) -> Result<Vec<TaggerOutput>> {
|
||||||
let input = Tensor::from_array((
|
let input = Tensor::from_array((
|
||||||
[batch_size, self.input_size, self.input_size, 3usize],
|
[count, self.input_size, self.input_size, 3usize],
|
||||||
image_array.into_boxed_slice(),
|
pixels.to_vec().into_boxed_slice(),
|
||||||
))
|
))
|
||||||
.map_err(|error| anyhow::anyhow!("{error}"))?;
|
.map_err(|error| anyhow::anyhow!("{error}"))?;
|
||||||
|
|
||||||
@@ -511,23 +608,48 @@ impl WdTagger {
|
|||||||
.run(ort::inputs! { input_name.as_str() => input })
|
.run(ort::inputs! { input_name.as_str() => input })
|
||||||
.map_err(|error| anyhow::anyhow!("{error}"))?;
|
.map_err(|error| anyhow::anyhow!("{error}"))?;
|
||||||
|
|
||||||
let (_, probabilities) = outputs[0]
|
let (_, probs) = outputs[0]
|
||||||
.try_extract_tensor::<f32>()
|
.try_extract_tensor::<f32>()
|
||||||
.map_err(|error| anyhow::anyhow!("{error}"))?;
|
.map_err(|error| anyhow::anyhow!("{error}"))?;
|
||||||
|
|
||||||
let probs: &[f32] = probabilities;
|
let n_labels = self.labels.len();
|
||||||
|
if probs.len() != count * n_labels {
|
||||||
if probs.len() != self.labels.len() {
|
|
||||||
anyhow::bail!(
|
anyhow::bail!(
|
||||||
"Model output length {} does not match label count {}",
|
"Model output length {} does not match {} images x {} labels",
|
||||||
probs.len(),
|
probs.len(),
|
||||||
self.labels.len()
|
count,
|
||||||
|
n_labels
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Collect rating scores (category 9) - pick the argmax as the rating.
|
// Borrow disjoint fields (not all of `self`) so this doesn't conflict
|
||||||
let rating = self
|
// with the mutable session borrow `outputs` still holds.
|
||||||
.labels
|
let labels = &self.labels;
|
||||||
|
let threshold = self.threshold;
|
||||||
|
Ok(probs
|
||||||
|
.chunks_exact(n_labels)
|
||||||
|
.map(|row| Self::tags_from_probs(labels, threshold, row, max_tags))
|
||||||
|
.collect())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Run a single preprocessed image through the model.
|
||||||
|
fn infer_one(&mut self, image_array: Vec<f32>, max_tags: usize) -> Result<TaggerOutput> {
|
||||||
|
self.infer_batch(&image_array, 1, max_tags)?
|
||||||
|
.into_iter()
|
||||||
|
.next()
|
||||||
|
.ok_or_else(|| anyhow::anyhow!("tagger produced no output for image"))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Convert one image's class probabilities into its rating + sorted tags.
|
||||||
|
/// Associated (not `&self`) so callers can hold a disjoint session borrow.
|
||||||
|
fn tags_from_probs(
|
||||||
|
labels: &[TagEntry],
|
||||||
|
threshold: f32,
|
||||||
|
probs: &[f32],
|
||||||
|
max_tags: usize,
|
||||||
|
) -> TaggerOutput {
|
||||||
|
// Rating (category 9): pick the argmax label as the rating.
|
||||||
|
let rating = labels
|
||||||
.iter()
|
.iter()
|
||||||
.zip(probs.iter())
|
.zip(probs.iter())
|
||||||
.filter(|(entry, _)| entry.category == RATING_CATEGORY)
|
.filter(|(entry, _)| entry.category == RATING_CATEGORY)
|
||||||
@@ -535,14 +657,13 @@ impl WdTagger {
|
|||||||
.map(|(entry, _)| entry.name.clone())
|
.map(|(entry, _)| entry.name.clone())
|
||||||
.unwrap_or_else(|| "general".to_string());
|
.unwrap_or_else(|| "general".to_string());
|
||||||
|
|
||||||
// Collect general + character tags above threshold, sorted by confidence.
|
// General + character tags above threshold, sorted by confidence.
|
||||||
let mut tags: Vec<TagResult> = self
|
let mut tags: Vec<TagResult> = labels
|
||||||
.labels
|
|
||||||
.iter()
|
.iter()
|
||||||
.zip(probs.iter())
|
.zip(probs.iter())
|
||||||
.filter(|(entry, prob)| {
|
.filter(|(entry, prob)| {
|
||||||
(entry.category == GENERAL_CATEGORY || entry.category == CHARACTER_CATEGORY)
|
(entry.category == GENERAL_CATEGORY || entry.category == CHARACTER_CATEGORY)
|
||||||
&& **prob >= self.threshold
|
&& **prob >= threshold
|
||||||
})
|
})
|
||||||
.map(|(entry, prob)| TagResult {
|
.map(|(entry, prob)| TagResult {
|
||||||
tag: entry.name.clone(),
|
tag: entry.name.clone(),
|
||||||
@@ -553,16 +674,7 @@ impl WdTagger {
|
|||||||
tags.sort_by(|a, b| b.confidence.total_cmp(&a.confidence));
|
tags.sort_by(|a, b| b.confidence.total_cmp(&a.confidence));
|
||||||
tags.truncate(max_tags);
|
tags.truncate(max_tags);
|
||||||
|
|
||||||
log::info!(
|
TaggerOutput { tags, rating }
|
||||||
"WD tagger: {} tags (threshold={}, rating={}) in {:?} for {}",
|
|
||||||
tags.len(),
|
|
||||||
self.threshold,
|
|
||||||
rating,
|
|
||||||
started_at.elapsed(),
|
|
||||||
image_path.display(),
|
|
||||||
);
|
|
||||||
|
|
||||||
Ok(TaggerOutput { tags, rating })
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user