From 988d8911021649742405eeff42c4ef4550172a97 Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Fri, 17 Jul 2026 07:59:30 -0600 Subject: [PATCH] Added a Sample Next Step in the job gear dropdown to force a sample on the next step. --- .../sd_trainer/DiffusionTrainer.py | 35 +++++++++++++++++++ ui/prisma/schema.prisma | 1 + .../app/api/jobs/[jobID]/sample_now/route.ts | 19 ++++++++++ ui/src/components/JobActionBar.tsx | 20 +++++++++-- ui/src/utils/jobs.ts | 16 +++++++++ 5 files changed, 88 insertions(+), 3 deletions(-) create mode 100644 ui/src/app/api/jobs/[jobID]/sample_now/route.ts diff --git a/extensions_built_in/sd_trainer/DiffusionTrainer.py b/extensions_built_in/sd_trainer/DiffusionTrainer.py index 59a264c6..524db3cd 100644 --- a/extensions_built_in/sd_trainer/DiffusionTrainer.py +++ b/extensions_built_in/sd_trainer/DiffusionTrainer.py @@ -203,6 +203,40 @@ class DiffusionTrainer(SDTrainer): if self.progress_bar is not None: self.progress_bar.unpause() + def should_sample(self): + if not self.is_ui_trainer: + return False + def _check_sample(): + with self._db_connect() as conn: + cursor = conn.cursor() + cursor.execute( + "SELECT sample_now FROM Job WHERE id = ?", (self.job_id,)) + sample_now = cursor.fetchone() + return False if sample_now is None else sample_now[0] == 1 + + return self._retry_db_operation(_check_sample) + + def maybe_sample(self): + if not self.is_ui_trainer: + return + if self.should_sample(): + self.update_db_key("sample_now", 0) + if self.progress_bar is not None: + self.progress_bar.pause() + print_acc(f"\nSampling at step {self.step_num}") + # clear any grads + self.optimizer.zero_grad() + if self.train_config.free_u: + self.sd.pipeline.disable_freeu() + self.sample(self.step_num) + if self.train_config.unload_text_encoder: + # make sure the text encoder is unloaded + self.sd.text_encoder_to('cpu') + self.ensure_params_requires_grad() + flush() + if self.progress_bar is not None: + self.progress_bar.unpause() + async def _update_key(self, key, value): if not self.accelerator.is_main_process: return @@ -319,6 +353,7 @@ class DiffusionTrainer(SDTrainer): self.update_step() self.maybe_stop() self.maybe_save() + self.maybe_sample() def hook_before_model_load(self): super().hook_before_model_load() diff --git a/ui/prisma/schema.prisma b/ui/prisma/schema.prisma index 79c90fbf..a74c425f 100644 --- a/ui/prisma/schema.prisma +++ b/ui/prisma/schema.prisma @@ -40,6 +40,7 @@ model Job { job_type String @default("train") // 'train', 'caption' job_ref String? // can be used for anything for special jobs, like dataset path for caption jobs save_now Boolean @default(false) // if true, the job will be saved on the next step + sample_now Boolean @default(false) // if true, the job will be sampled on the next step @@index([status]) @@index([gpu_ids]) diff --git a/ui/src/app/api/jobs/[jobID]/sample_now/route.ts b/ui/src/app/api/jobs/[jobID]/sample_now/route.ts new file mode 100644 index 00000000..7d785be5 --- /dev/null +++ b/ui/src/app/api/jobs/[jobID]/sample_now/route.ts @@ -0,0 +1,19 @@ +import { NextRequest, NextResponse } from 'next/server'; +import { PrismaClient } from '@prisma/client'; + +const prisma = new PrismaClient(); + +export async function GET(request: NextRequest, { params }: { params: { jobID: string } }) { + const { jobID } = await params; + + const job = await prisma.job.update({ + where: { id: jobID }, + data: { + sample_now: true, + }, + }); + + console.log(`Job ${jobID} marked to sample on next step`); + + return NextResponse.json(job); +} diff --git a/ui/src/components/JobActionBar.tsx b/ui/src/components/JobActionBar.tsx index 6ae7d781..b63d60ee 100644 --- a/ui/src/components/JobActionBar.tsx +++ b/ui/src/components/JobActionBar.tsx @@ -1,9 +1,9 @@ import Link from 'next/link'; -import { Eye, Trash2, Pen, Play, Pause, Cog, X, Copy, Save, OctagonX } from 'lucide-react'; +import { Eye, Trash2, Pen, Play, Pause, Cog, X, Copy, Save, OctagonX, Image } from 'lucide-react'; import { Button } from '@headlessui/react'; import { openConfirm } from '@/components/ConfirmModal'; import { Job } from '@prisma/client'; -import { startJob, stopJob, deleteJob, getAvaliableJobActions, markJobAsStopped, saveJobNow } from '@/utils/jobs'; +import { startJob, stopJob, deleteJob, getAvaliableJobActions, markJobAsStopped, saveJobNow, sampleJobNow } from '@/utils/jobs'; import { startQueue } from '@/utils/queue'; import { Menu, MenuButton, MenuItem, MenuItems } from '@headlessui/react'; import { redirect } from 'next/navigation'; @@ -146,7 +146,7 @@ export default function JobActionBar({ {job.job_type === 'train' && ( @@ -173,6 +173,20 @@ export default function JobActionBar({ )} + {job.job_type === 'train' && canStop && ( + +
{ + await sampleJobNow(job.id); + if (onRefresh) onRefresh(); + }} + > + + Sample Next Step +
+
+ )}
{ }); }; +export const sampleJobNow = (jobID: string) => { + return new Promise((resolve, reject) => { + apiClient + .get(`/api/jobs/${jobID}/sample_now`) + .then(res => res.data) + .then(data => { + console.log('Job set to sample on next step:', data); + resolve(); + }) + .catch(error => { + console.error('Error setting job to sample on next step:', error); + reject(error); + }); + }); +}; + export const markJobAsStopped = (jobID: string) => { return new Promise((resolve, reject) => { apiClient