Improve music activity reporting and preserve Comfy host fixes

This commit is contained in:
Towsty
2026-09-06 23:05:46 -05:00
parent 53bb58bbd5
commit 5f96807496
8 changed files with 246 additions and 10 deletions
+55 -8
View File
@@ -186,9 +186,13 @@
@click="inspectJobId = studioJobId || liveMusicId" @click="inspectJobId = studioJobId || liveMusicId"
>Inspect this job</button> >Inspect this job</button>
</div> </div>
<div v-if="busy" class="h-2 overflow-hidden rounded-full bg-zinc-800"> <div v-if="busy" class="h-2 overflow-hidden rounded-full bg-zinc-800" role="progressbar" aria-label="Music generation" :aria-valuenow="indeterminateMusic ? undefined : progress" :aria-valuetext="indeterminateMusic ? 'Completion percentage unavailable' : `${progress}%`">
<div class="progress-glow h-full rounded-full bg-amber-400 transition-[width]" :style="{ width: `${Math.max(progress, 4)}%` }" /> <div class="progress-glow h-full rounded-full bg-amber-400 transition-[width]" :class="{ 'music-activity': indeterminateMusic && recentMusicCheck }" :style="{ width: indeterminateMusic ? '30%' : `${Math.max(progress, 4)}%` }" />
</div> </div>
<p v-if="busy && !queued && jobId" class="text-xs text-zinc-400">
{{ musicElapsed }} elapsed · {{ musicActivityText }}
<span v-if="indeterminateMusic" class="mt-1 block">YuE does not report a completion percentage during this stage.</span>
</p>
<AudioPlayer <AudioPlayer
v-if="audioUrl" v-if="audioUrl"
:src="audioUrl" :src="audioUrl"
@@ -257,6 +261,23 @@ const forceClearing = ref(false)
const error = ref('') const error = ref('')
const status = ref('') const status = ref('')
const progress = ref(0) const progress = ref(0)
const musicClock = ref(Date.now())
const musicStartedAt = ref(0)
const activeMusicEngine = ref('')
const musicActivity = ref<{ checkedAt: number; running: boolean } | null>(null)
let musicClockTimer: ReturnType<typeof setInterval> | null = null
const indeterminateMusic = computed(() => !queued.value && activeMusicEngine.value === 'yue' && progress.value < 98)
const recentMusicCheck = computed(() => Boolean(musicActivity.value?.running && musicClock.value - musicActivity.value.checkedAt < 15000))
const musicElapsed = computed(() => {
const seconds = musicStartedAt.value ? Math.max(0, Math.floor((musicClock.value - musicStartedAt.value) / 1000)) : 0
return `${Math.floor(seconds / 60)}:${String(seconds % 60).padStart(2, '0')}`
})
const musicActivityText = computed(() => {
if (!musicActivity.value) return 'Waiting for a Comfy activity check…'
const age = Math.max(0, Math.floor((musicClock.value - musicActivity.value.checkedAt) / 1000))
if (age >= 15) return `Activity check delayed · last response ${age}s ago`
return musicActivity.value.running ? `Comfy confirms this job is running · checked ${age}s ago` : 'Comfy responded · checking for the result…'
})
const audioUrl = ref('') const audioUrl = ref('')
const trackId = ref('') const trackId = ref('')
const jobId = ref('') const jobId = ref('')
@@ -367,12 +388,23 @@ function stopRecoverPoll() {
async function tryRecover() { async function tryRecover() {
if (settled || !busy.value || queued.value) return if (settled || !busy.value || queued.value) return
const currentId = jobId.value
if (currentId) {
try {
const snapshot = await $fetch<Record<string, any>>(`/api/generate/${currentId}`, { timeout: 10000 })
if (jobId.value !== currentId || settled) return
applyEvent(snapshot)
return
} catch {
// Fall back to saved-output recovery when the live job is unavailable.
}
}
try { try {
const recovered = await $fetch<Record<string, any>>('/api/generate/recover', { const recovered = await $fetch<Record<string, any>>('/api/generate/recover', {
method: 'POST', method: 'POST',
body: { jobId: jobId.value || undefined } body: { jobId: jobId.value || undefined }
}) })
if (recovered?.trackId) applyEvent(recovered) if (jobId.value === currentId && !settled && recovered?.trackId) applyEvent(recovered)
} catch { } catch {
/* still running */ /* still running */
} }
@@ -393,14 +425,17 @@ async function attachLiveMusic(opts: {
settled = false settled = false
if (typeof opts.progress === 'number') progress.value = opts.progress if (typeof opts.progress === 'number') progress.value = opts.progress
if (opts.message) status.value = opts.message if (opts.message) status.value = opts.message
else if (!status.value || /waiting for a generate|waiting in the job queue|waiting for gpu/i.test(status.value)) { else if (!status.value || /queueing|waiting for a generate|waiting in the job queue|waiting for gpu/i.test(status.value)) {
status.value = selectedEngine.value === 'yue' status.value = selectedEngine.value === 'yue'
? 'YuE running on Comfy — Stage A can take 10–20+ minutes' ? 'YuE running on Comfy — Stage A can take 10–20+ minutes'
: 'Generating…' : 'Generating…'
} }
if (jobId.value !== liveId) { if (jobId.value !== liveId) {
musicActivity.value = null
musicStartedAt.value = Date.now()
jobId.value = liveId jobId.value = liveId
listen(liveId) listen(liveId)
void tryRecover()
} }
if (!recoverPoll) { if (!recoverPoll) {
recoverPoll = setInterval(() => { void tryRecover() }, 8000) recoverPoll = setInterval(() => { void tryRecover() }, 8000)
@@ -419,7 +454,9 @@ async function refreshStudioQueue() {
waitingCount?: number waitingCount?: number
}>('/api/studio-queue').catch(() => ({ jobs: [] as Array<{ id: string; status: string; kind?: string; liveJobId?: string; lastError?: string }>, waitingCount: 0 })) }>('/api/studio-queue').catch(() => ({ jobs: [] as Array<{ id: string; status: string; kind?: string; liveJobId?: string; lastError?: string }>, waitingCount: 0 }))
const rows = data.jobs || [] const rows = data.jobs || []
const musicRow = rows.find(job => job.kind === 'music' && (job.status === 'waiting' || job.status === 'running' || job.status === 'held')) const musicRow = rows.find(job => job.kind === 'music' && job.id === studioJobId.value)
|| rows.find(job => job.kind === 'music' && job.status === 'running')
|| rows.find(job => job.kind === 'music' && (job.status === 'waiting' || job.status === 'held'))
liveMusicId.value = musicRow?.id || '' liveMusicId.value = musicRow?.id || ''
queueCount.value = data.waitingCount || rows.filter(job => job.status === 'waiting').length queueCount.value = data.waitingCount || rows.filter(job => job.status === 'waiting').length
@@ -493,6 +530,9 @@ async function resumeActiveMusic() {
} }
function applyEvent(payload: Record<string, any>) { function applyEvent(payload: Record<string, any>) {
if (typeof payload.elapsedMs === 'number') musicStartedAt.value = Date.now() - payload.elapsedMs
if (payload.engine) activeMusicEngine.value = payload.engine
if (payload.musicActivity) musicActivity.value = payload.musicActivity
if (payload.message) status.value = payload.message if (payload.message) status.value = payload.message
if (typeof payload.progress === 'number') progress.value = payload.progress if (typeof payload.progress === 'number') progress.value = payload.progress
if (payload.type === 'complete' || payload.status === 'complete') { if (payload.type === 'complete' || payload.status === 'complete') {
@@ -628,13 +668,12 @@ async function generate() {
await refreshStudioQueue() await refreshStudioQueue()
stopQueuePoll() stopQueuePoll()
queuePoll = setInterval(() => { void refreshStudioQueue() }, 2500) queuePoll = setInterval(() => { void refreshStudioQueue() }, 2500)
if (started.queued) { if (started.queued && !jobId.value) {
queued.value = true queued.value = true
status.value = 'Waiting in the job queue…' status.value = 'Waiting in the job queue…'
return return
} }
jobId.value = started.jobId if (!jobId.value) await attachLiveMusic({ studioId: started.studioJobId, liveId: started.jobId })
listen(started.jobId)
stopRecoverPoll() stopRecoverPoll()
recoverPoll = setInterval(() => { void tryRecover() }, 8000) recoverPoll = setInterval(() => { void tryRecover() }, 8000)
window.setTimeout(() => { void tryRecover() }, 12000) window.setTimeout(() => { void tryRecover() }, 12000)
@@ -721,6 +760,7 @@ function applyPreset() {
} }
onMounted(async () => { onMounted(async () => {
musicClockTimer = setInterval(() => { musicClock.value = Date.now() }, 1000)
await loadLibrary() await loadLibrary()
await loadMusicPresets() await loadMusicPresets()
const route = useRoute() const route = useRoute()
@@ -749,8 +789,15 @@ onMounted(async () => {
}) })
onBeforeUnmount(() => { onBeforeUnmount(() => {
if (musicClockTimer) clearInterval(musicClockTimer)
stopListen() stopListen()
stopRecoverPoll() stopRecoverPoll()
stopQueuePoll() stopQueuePoll()
}) })
</script> </script>
<style scoped>
.music-activity { animation: music-travel 2s ease-in-out infinite alternate; }
@keyframes music-travel { from { transform: translateX(0); } to { transform: translateX(233%); } }
@media (prefers-reduced-motion: reduce) { .music-activity { animation: none; } }
</style>
+46
View File
@@ -0,0 +1,46 @@
"""Hide VideoHelperSuite's background command windows without changing its media commands."""
import ast
import pathlib
import shutil
import sys
def patch(source):
lines = source.splitlines(keepends=True)
offsets = [0]
for line in lines:
offsets.append(offsets[-1] + len(line.encode('utf-8')))
raw = source.encode('utf-8')
calls = []
for node in ast.walk(ast.parse(source)):
if not isinstance(node, ast.Call) or not isinstance(node.func, ast.Attribute):
continue
if not isinstance(node.func.value, ast.Name) or node.func.value.id != 'subprocess' or node.func.attr not in ('Popen', 'run', 'check_output', 'check_call', 'call'):
continue
if any(k.arg == 'creationflags' for k in node.keywords):
continue
calls.append(node)
for node in sorted(calls, key=lambda n: (n.end_lineno, n.end_col_offset), reverse=True):
end = offsets[node.end_lineno - 1] + node.end_col_offset - 1
prefix = b'' if raw[:end].rstrip().endswith(b',') else b','
raw = raw[:end] + prefix + b" creationflags=getattr(subprocess, 'CREATE_NO_WINDOW', 0)" + raw[end:]
result = raw.decode('utf-8')
ast.parse(result)
return result
if __name__ == '__main__':
root = pathlib.Path(sys.argv[1]).resolve(strict=True)
plans = []
for name in ('nodes.py', 'utils.py', 'load_video_nodes.py'):
path = root / 'videohelpersuite' / name
source = path.read_text(encoding='utf-8')
plans.append((path, source, patch(source)))
for path, source, result in plans:
if source == result:
print(f'Already patched: {path.name}')
continue
backup = path.with_suffix('.py.before-aigen-hidden')
if backup.exists():
raise RuntimeError(f'Backup already exists: {backup}')
shutil.copy2(path, backup)
path.write_text(result, encoding='utf-8')
print(f'Patched {path.name}; original backed up.')
+37
View File
@@ -0,0 +1,37 @@
"""Make YuE's HF token samplers honor Comfy's cancel flag. Run with the node directory."""
import ast
import pathlib
import shutil
import sys
CHECK = ' from comfy.model_management import throw_exception_if_processing_interrupted\n throw_exception_if_processing_interrupted()\n'
def patch(source):
tree = ast.parse(source)
cls = next(n for n in tree.body if isinstance(n, ast.ClassDef) and n.name == 'BlockTokenRangeProcessor')
method = next(n for n in cls.body if isinstance(n, ast.FunctionDef) and n.name == '__call__')
if 'throw_exception_if_processing_interrupted' in ast.get_source_segment(source, method):
return source
lines = source.splitlines(keepends=True)
lines.insert(method.body[0].lineno - 1, CHECK)
result = ''.join(lines)
ast.parse(result)
return result
if __name__ == '__main__':
root = pathlib.Path(sys.argv[1]).resolve(strict=True)
plans = []
for name in ('infer.py', 'common.py'):
path = root / 'inference' / name
source = path.read_text(encoding='utf-8')
plans.append((path, source, patch(source)))
for path, source, result in plans:
if source == result:
print(f'Already patched: {path.name}')
continue
backup = path.with_suffix('.py.before-aigen-cancel')
if backup.exists():
raise RuntimeError(f'Backup already exists: {backup}')
shutil.copy2(path, backup)
path.write_text(result, encoding='utf-8')
print(f'Patched {path.name}; original backed up.')
+3
View File
@@ -31,6 +31,7 @@ export interface JobEvent {
} }
export interface Job { export interface Job {
musicActivity?: { checkedAt: number; running: boolean }
id: string id: string
kind?: 'video' | 'edit' | 'music' kind?: 'video' | 'edit' | 'music'
promptId?: string promptId?: string
@@ -192,6 +193,8 @@ export function jobSnapshot(job: Job) {
return { return {
jobId: job.id, jobId: job.id,
kind: job.kind || 'video', kind: job.kind || 'video',
engine: job.library?.engine,
musicActivity: job.musicActivity,
type: 'snapshot' as const, type: 'snapshot' as const,
status: job.status, status: job.status,
message: job.message, message: job.message,
+14 -1
View File
@@ -1,6 +1,6 @@
import { createJob, emitJob, type Job } from '~/server/utils/jobs' import { createJob, emitJob, type Job } from '~/server/utils/jobs'
import { extractAudio, fetchHistory, fetchHistoryAll, findHistoryAudio, freeComfyVram, purgeComfyArtifacts, queuePrompt } from '~/server/utils/comfy' import { extractAudio, fetchHistory, fetchHistoryAll, findHistoryAudio, freeComfyVram, purgeComfyArtifacts, queuePrompt } from '~/server/utils/comfy'
import { comfyWsUrl } from '~/server/utils/comfy' import { comfyWsUrl, comfyFetch } from '~/server/utils/comfy'
import { ensureComfyReady } from '~/server/utils/comfyLifecycle' import { ensureComfyReady } from '~/server/utils/comfyLifecycle'
import { downloadComfyAudio, saveTrack } from '~/server/utils/library' import { downloadComfyAudio, saveTrack } from '~/server/utils/library'
import { buildMusicWorkflow, musicFilenamePrefix, assertMusicEngineNodes } from '~/server/utils/musicWorkflow' import { buildMusicWorkflow, musicFilenamePrefix, assertMusicEngineNodes } from '~/server/utils/musicWorkflow'
@@ -231,6 +231,19 @@ function watchMusicJob(job: Job): Promise<void> {
} }
const pollHistory = async () => { const pollHistory = async () => {
if (settled || finishing) return
if (job.promptId) {
try {
const response = await comfyFetch('/queue', { signal: AbortSignal.timeout(2500) })
if (response.ok) {
const queue = await response.json() as { queue_running?: unknown[][] }
job.musicActivity = {
checkedAt: Date.now(),
running: Boolean(queue.queue_running?.some(row => row[1] === job.promptId))
}
}
} catch { /* Keep the last confirmation timestamp so the UI can show stale checks. */ }
}
if (settled || finishing) return if (settled || finishing) return
try { try {
const history = await fetchHistoryAll() const history = await fetchHistoryAll()
+1 -1
View File
@@ -43,7 +43,7 @@ export function inspectStudioJob(job: StudioJob) {
createdAt: job.createdAt, createdAt: job.createdAt,
updatedAt: job.updatedAt, updatedAt: job.updatedAt,
status: job.status, status: job.status,
kind: job.kind === 'edit' ? 'edit' as const : 'video' as const, kind: job.kind === 'music' ? 'music' as const : job.kind === 'edit' ? 'edit' as const : 'video' as const,
name: job.name, name: job.name,
prompt: job.prompt, prompt: job.prompt,
shotCount: job.shotCount, shotCount: job.shotCount,
+55
View File
@@ -0,0 +1,55 @@
import test from 'node:test'
import assert from 'node:assert/strict'
import { readFileSync } from 'node:fs'
import ts from 'typescript'
const page = readFileSync(new URL('../pages/music.vue', import.meta.url), 'utf8')
const source = page.slice(page.indexOf('async function tryRecover()'), page.indexOf('async function attachLiveMusic('))
const code = ts.transpileModule(source, { compilerOptions: { target: ts.ScriptTarget.ES2022 } }).outputText
function recovery(fetch, apply, id = { value: 'music-1' }) {
return new Function('$fetch', 'applyEvent', 'jobId', 'busy', 'queued', 'settled', `${code}; return tryRecover`)(
fetch, apply, id, { value: true }, { value: false }, false
)
}
test('music status recovers without a working event stream, including completion and errors', async () => {
for (const status of ['running', 'complete', 'error']) {
const snapshot = { status, message: 'YuE Stage A', progress: 20 }
const events = []
await recovery(async url => {
assert.equal(url, '/api/generate/music-1')
return snapshot
}, event => events.push(event))()
assert.deepEqual(events, [snapshot])
}
})
test('late status response cannot overwrite a different music job', async () => {
const id = { value: 'music-1' }
const events = []
await recovery(async () => {
id.value = 'music-2'
return { status: 'complete' }
}, event => events.push(event), id)()
assert.deepEqual(events, [])
})
test('missing live job still falls back to saved audio recovery', async () => {
const events = []
await recovery(async url => {
if (url === '/api/generate/music-1') throw new Error('Job unavailable')
assert.equal(url, '/api/generate/recover')
return { trackId: 'saved-track' }
}, event => events.push(event))()
assert.deepEqual(events, [{ trackId: 'saved-track' }])
})
test('inspecting a queued YuE job uses music controls', () => {
const inspect = readFileSync(new URL('../server/utils/queuedInspect.ts', import.meta.url), 'utf8')
const fn = inspect.slice(inspect.indexOf('export function inspectStudioJob'), inspect.indexOf('export function inspectLiveMusicJob')).replace('export function', 'function')
const js = ts.transpileModule(fn, { compilerOptions: { target: ts.ScriptTarget.ES2022 } }).outputText
const result = new Function(`${js}; return inspectStudioJob({kind:'music',payload:{prompt:'rock',musicEngine:'yue'}})`)()
assert.equal(result.kind, 'music')
assert.equal(result.payload.musicEngine, 'yue')
assert.equal(result.payload.tags, 'rock')
})
+35
View File
@@ -0,0 +1,35 @@
import ast
import importlib.util
import pathlib
import sys
import types
import unittest
spec = importlib.util.spec_from_file_location('patcher', pathlib.Path(__file__).parents[1] / 'scripts/patch_yue_cancellation.py')
patcher = importlib.util.module_from_spec(spec)
spec.loader.exec_module(patcher)
class CancellationTests(unittest.TestCase):
def test_interrupt_propagates_before_sampling(self):
source = 'class BlockTokenRangeProcessor:\n def __call__(self, input_ids, scores):\n return scores\n'
result = patcher.patch(source)
self.assertEqual(result, patcher.patch(result))
cancelled = False
class Interrupted(Exception):
pass
def check():
if cancelled:
raise Interrupted()
module = types.ModuleType('comfy.model_management')
module.throw_exception_if_processing_interrupted = check
sys.modules['comfy.model_management'] = module
namespace = {}
exec(compile(ast.parse(result), '<patched>', 'exec'), namespace)
callback = namespace['BlockTokenRangeProcessor']()
self.assertEqual(callback(None, 42), 42)
cancelled = True
with self.assertRaises(Interrupted):
callback(None, 42)
if __name__ == '__main__':
unittest.main()