Added a Sample Next Step in the job gear dropdown to force a sample on the next step.
This commit is contained in:
parent
7a3d94ed03
commit
988d891102
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in New Issue