Added a Sample Next Step in the job gear dropdown to force a sample on the next step.

This commit is contained in:
Jaret Burkett 2026-07-17 07:59:30 -06:00
parent 7a3d94ed03
commit 988d891102
5 changed files with 88 additions and 3 deletions

View File

@ -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()

View File

@ -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])

View File

@ -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);
}

View File

@ -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({
</MenuButton>
<MenuItems
anchor={{ to: menuAnchor, gap: 16 }}
className="bg-gray-900 border border-gray-700 rounded shadow-lg w-52 px-2 py-2 z-50"
className="bg-gray-900 border border-gray-700 rounded shadow-lg w-60 px-2 py-2 z-50"
>
{job.job_type === 'train' && (
<MenuItem>
@ -173,6 +173,20 @@ export default function JobActionBar({
</div>
</MenuItem>
)}
{job.job_type === 'train' && canStop && (
<MenuItem>
<div
className="cursor-pointer px-4 py-1 hover:bg-gray-800 rounded flex items-center gap-2"
onClick={async () => {
await sampleJobNow(job.id);
if (onRefresh) onRefresh();
}}
>
<Image className="w-4 h-4" />
Sample Next Step
</div>
</MenuItem>
)}
<MenuItem>
<div
className="cursor-pointer px-4 py-1 hover:bg-gray-800 rounded flex items-center gap-2"

View File

@ -66,6 +66,22 @@ export const saveJobNow = (jobID: string) => {
});
};
export const sampleJobNow = (jobID: string) => {
return new Promise<void>((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<void>((resolve, reject) => {
apiClient