Migrate latent preview code, embed in KSampler

Latent video previews are now properly supported on the frontend without
the need for making manual changes.

Further testing/polish is still needed.
This commit is contained in:
Austin Mroz
2024-12-04 20:00:00 -06:00
parent 2b15c4f03c
commit a9d08964ad
4 changed files with 190 additions and 145 deletions
+1
View File
@@ -2,6 +2,7 @@ from .videohelpersuite.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPI
import folder_paths
from .videohelpersuite.server import server
from .videohelpersuite import documentation
from .videohelpersuite import latent_preview
WEB_DIRECTORY = "./web"
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
+104
View File
@@ -0,0 +1,104 @@
import asyncio
import subprocess
from PIL import Image
import latent_preview
import server
serv = server.PromptServer.instance
web = server.web
class PreviewInstance:
def __init__(self):
self.has_data = asyncio.Event()
self.preview_at = 0
self.data = []
self.closed = False
#TODO sid?
preview_instances = {}
def get_instance(node_id):
if node_id not in preview_instances:
preview_instances[node_id] = PreviewInstance()
return preview_instances[node_id]
@serv.routes.get("/vhs/latentvideopreview")
async def latent_video_preview(request):
query = request.rel_url.query
if 'node_id' in query:
instance = get_instance(query['node_id'])
elif len(preview_instances) > 0:
instance = next(preview_instances.values())
else:
return web.response(status=400)
rate = 8.0
args = ['ffmpeg','-v', 'error', '-f', 'rawvideo', '-pix_fmt', 'rgb24', '-s', '512x512', '-r', str(rate), '-i', '-']
args += ['-c:v', 'libvpx-vp9','-deadline', 'realtime', '-cpu-used', '8', '-f', 'webm', '-']
try:
proc = await asyncio.create_subprocess_exec(*args, stdout=subprocess.PIPE, stdin=subprocess.PIPE)
try:
resp = web.StreamResponse()
resp.content_type = 'video/webm'
await resp.prepare(request)
async def read_loop():
delay = asyncio.sleep(.1)
while data := await proc.stdout.read(2**20):
await asyncio.gather(delay, resp.write(data))
delay = asyncio.sleep(.1)
async def write_loop():
await instance.has_data.wait()
preview_at = 0
delay = asyncio.create_task(asyncio.sleep(1/rate))
frame_at = 0
while True:
proc.stdin.write(instance.data[frame_at])
frame_at = (frame_at + 1) % len(instance.data)
if frame_at == instance.preview_at:
if not instance.has_data.is_set() and instance.closed:
proc.stdin.close()
delay.close()
return
await instance.has_data.wait()
instance.has_data.clear()
await asyncio.gather(proc.stdin.drain(), delay)
delay = asyncio.create_task(asyncio.sleep(1/rate))
await asyncio.gather(read_loop(), write_loop())
await proc.wait()
except (ConnectionResetError, ConnectionError) as e:
pass
finally:
print('ded')
#Kill ffmpeg before the pipe is closed
proc.kill()
except BrokenPipeError as e:
pass
return resp
orig_get_previewer = latent_preview.get_previewer
def get_latent_video_previewer(device, latent_format):
node_id = serv.last_node_id
serv.send_sync('VHS_latentpreview', node_id)
previewer = orig_get_previewer(device, latent_format)
if not hasattr(previewer, "decode_latent_to_preview"):
return None
original_decode = previewer.decode_latent_to_preview
def wrapped_decode(_, x0):
inst = get_instance(node_id)
num_images = x0.size(0)
if len(inst.data) != num_images:
inst.data = [b''] * num_images
for i in range(num_images):
sub_image = original_decode(x0[i:i+1])
if sub_image.size[0] != 512 or sub_image.size[1] != 512:
sub_image = sub_image.resize((512, 512),
Image.Resampling.NEAREST)
inst.data[i] = sub_image.tobytes()
inst.has_data.set()
return None
previewer.decode_latent_to_preview_image = wrapped_decode
return previewer
latent_preview.get_previewer = get_latent_video_previewer
-71
View File
@@ -5,12 +5,10 @@ import subprocess
import re
import asyncio
from PIL import Image
from .utils import is_url, get_sorted_dir_files_from_directory, ffmpeg_path, \
validate_sequence, is_safe_path, strip_path, try_download_video, ENCODE_ARGS
from comfy.k_diffusion.utils import FolderOfImages
import latent_preview
web = server.web
@@ -194,72 +192,3 @@ async def get_path(request):
pass
return web.json_response(valid_items)
has_preview_data = asyncio.Event()
preview_at = 0
preview_data = []
@server.PromptServer.instance.routes.get("/vhs/latentvideopreview")
async def latent_video_preview(request):
rate = 8.0
args = ['ffmpeg','-v', 'error', '-f', 'rawvideo', '-pix_fmt', 'rgb24', '-s', '512x512', '-r', str(rate), '-i', '-']
args += ['-c:v', 'libvpx-vp9','-deadline', 'realtime', '-cpu-used', '8', '-f', 'webm', '-']
try:
proc = await asyncio.create_subprocess_exec(*args, stdout=subprocess.PIPE, stdin=subprocess.PIPE)
try:
resp = web.StreamResponse()
resp.content_type = 'video/webm'
await resp.prepare(request)
async def read_loop():
delay = asyncio.sleep(.1)
while data := await proc.stdout.read(2**20):
await asyncio.gather(delay, resp.write(data))
delay = asyncio.sleep(.1)
async def write_loop():
await has_preview_data.wait()
preview_at = 0
delay = asyncio.create_task(asyncio.sleep(1/rate))
frame_at = 0
while True:
proc.stdin.write(preview_data[frame_at])
frame_at = (frame_at + 1) % len(preview_data)
if frame_at == preview_at:
await has_preview_data.wait()
has_preview_data.clear()
await asyncio.gather(proc.stdin.drain(), delay)
delay = asyncio.create_task(asyncio.sleep(1/rate))
await asyncio.gather(read_loop(), write_loop())
await proc.wait()
except (ConnectionResetError, ConnectionError) as e:
pass
finally:
print('ded')
#Kill ffmpeg before the pipe is closed
proc.kill()
except BrokenPipeError as e:
pass
return resp
orig_get_previewer = latent_preview.get_previewer
def get_latent_video_previewer(device, latent_format):
previewer = orig_get_previewer(device, latent_format)
original_decode = previewer.decode_latent_to_preview
def wrapped_decode(x0):
global preview_data
num_images = x0.size(0)
if len(preview_data) != num_images:
preview_data = [b''] * num_images
for i in range(num_images):
sub_image = original_decode(x0[i:i+1])
if sub_image.size[0] != 512 or sub_image.size[1] != 512:
sub_image = sub_image.resize((512, 512),
Image.Resampling.NEAREST)
preview_data[i] = sub_image.tobytes()
has_preview_data.set()
return sub_image
previewer.decode_latent_to_preview = wrapped_decode
return previewer
latent_preview.get_previewer = get_latent_video_previewer
+85 -74
View File
@@ -748,80 +748,85 @@ function addUploadWidget(nodeType, nodeData, widgetName, type="video") {
uploadWidget.options.serialize = false;
});
}
function _addVideoPreview(node) {
var element = document.createElement("div");
const previewNode = node;
var previewWidget = node.addDOMWidget("videopreview", "preview", element, {
serialize: false,
hideOnZoom: false,
getValue() {
return element.value;
},
setValue(v) {
element.value = v;
},
});
previewWidget.computeSize = function(width) {
if (this.aspectRatio && !this.parentEl.hidden) {
let height = (previewNode.size[0]-20)/ this.aspectRatio + 10;
if (!(height > 0)) {
height = 0;
}
this.computedHeight = height + 10;
return [width, height];
}
return [width, -4];//no loaded src, widget should not display
}
element.addEventListener('contextmenu', (e) => {
e.preventDefault()
return app.canvas._mousedown_callback(e)
}, true);
element.addEventListener('pointerdown', (e) => {
e.preventDefault()
return app.canvas._mousedown_callback(e)
}, true);
element.addEventListener('mousewheel', (e) => {
e.preventDefault()
return app.canvas._mousewheel_callback(e)
}, true);
previewWidget.value = {hidden: false, paused: false, params: {},
muted: app.ui.settings.getSettingValue("VHS.DefaultMute", false)}
previewWidget.parentEl = document.createElement("div");
previewWidget.parentEl.className = "vhs_preview";
previewWidget.parentEl.style['width'] = "100%"
element.appendChild(previewWidget.parentEl);
previewWidget.videoEl = document.createElement("video");
previewWidget.videoEl.controls = false;
previewWidget.videoEl.loop = true;
previewWidget.videoEl.muted = true;
previewWidget.videoEl.style['width'] = "100%"
previewWidget.videoEl.addEventListener("loadedmetadata", () => {
previewWidget.aspectRatio = previewWidget.videoEl.videoWidth / previewWidget.videoEl.videoHeight;
fitHeight(node);
});
previewWidget.videoEl.addEventListener("error", () => {
//TODO: consider a way to properly notify the user why a preview isn't shown.
previewWidget.parentEl.hidden = true;
fitHeight(node);
});
previewWidget.videoEl.onmouseenter = () => {
previewWidget.videoEl.muted = previewWidget.value.muted
};
previewWidget.videoEl.onmouseleave = () => {
previewWidget.videoEl.muted = true;
};
previewWidget.imgEl = document.createElement("img");
previewWidget.imgEl.style['width'] = "100%"
previewWidget.imgEl.hidden = true;
previewWidget.imgEl.onload = () => {
previewWidget.aspectRatio = previewWidget.imgEl.naturalWidth / previewWidget.imgEl.naturalHeight;
fitHeight(node);
};
previewWidget.parentEl.appendChild(previewWidget.videoEl)
previewWidget.parentEl.appendChild(previewWidget.imgEl)
return previewWidget
}
function addVideoPreview(nodeType) {
chainCallback(nodeType.prototype, "onNodeCreated", function() {
var element = document.createElement("div");
const previewNode = this;
var previewWidget = this.addDOMWidget("videopreview", "preview", element, {
serialize: false,
hideOnZoom: false,
getValue() {
return element.value;
},
setValue(v) {
element.value = v;
},
});
previewWidget.computeSize = function(width) {
if (this.aspectRatio && !this.parentEl.hidden) {
let height = (previewNode.size[0]-20)/ this.aspectRatio + 10;
if (!(height > 0)) {
height = 0;
}
this.computedHeight = height + 10;
return [width, height];
}
return [width, -4];//no loaded src, widget should not display
}
element.addEventListener('contextmenu', (e) => {
e.preventDefault()
return app.canvas._mousedown_callback(e)
}, true);
element.addEventListener('pointerdown', (e) => {
e.preventDefault()
return app.canvas._mousedown_callback(e)
}, true);
element.addEventListener('mousewheel', (e) => {
e.preventDefault()
return app.canvas._mousewheel_callback(e)
}, true);
previewWidget.value = {hidden: false, paused: false, params: {},
muted: app.ui.settings.getSettingValue("VHS.DefaultMute", false)}
previewWidget.parentEl = document.createElement("div");
previewWidget.parentEl.className = "vhs_preview";
previewWidget.parentEl.style['width'] = "100%"
element.appendChild(previewWidget.parentEl);
previewWidget.videoEl = document.createElement("video");
previewWidget.videoEl.controls = false;
previewWidget.videoEl.loop = true;
previewWidget.videoEl.muted = true;
previewWidget.videoEl.style['width'] = "100%"
previewWidget.videoEl.addEventListener("loadedmetadata", () => {
previewWidget.aspectRatio = previewWidget.videoEl.videoWidth / previewWidget.videoEl.videoHeight;
fitHeight(this);
});
previewWidget.videoEl.addEventListener("error", () => {
//TODO: consider a way to properly notify the user why a preview isn't shown.
previewWidget.parentEl.hidden = true;
fitHeight(this);
});
previewWidget.videoEl.onmouseenter = () => {
previewWidget.videoEl.muted = previewWidget.value.muted
};
previewWidget.videoEl.onmouseleave = () => {
previewWidget.videoEl.muted = true;
};
previewWidget.imgEl = document.createElement("img");
previewWidget.imgEl.style['width'] = "100%"
previewWidget.imgEl.hidden = true;
previewWidget.imgEl.onload = () => {
previewWidget.aspectRatio = previewWidget.imgEl.naturalWidth / previewWidget.imgEl.naturalHeight;
fitHeight(this);
};
let previewWidget = _addVideoPreview(this)
var timeout = null;
this.updateParameters = (params, force_update) => {
if (!previewWidget.value.params) {
@@ -857,9 +862,9 @@ function addVideoPreview(nodeType) {
(params.format?.split('/')[1] == 'gif') || params.format == 'folder') {
this.videoEl.autoplay = !this.value.paused && !this.value.hidden;
let target_width = 256
if (element.style?.width) {
if (previewWidget.element?.style?.width) {
//overscale to allow scrolling. Endpoint won't return higher than native
target_width = element.style.width.slice(0,-2)*2;
target_width = previewWidget.element.style.width.slice(0,-2)*2;
}
if (!params.force_size || params.force_size.includes("?") || params.force_size == "Disabled") {
params.force_size = target_width+"x?"
@@ -1609,6 +1614,12 @@ app.registerExtension({
}
}
}
api.addEventListener('VHS_latentpreview', ({detail}) => {
let node = app.graph.getNodeById(detail)
let previewWidget = node.widgets.find((w) => w.name == 'videopreview') ??
_addVideoPreview(node)
previewWidget.videoEl.src = api.apiURL('/vhs/latentvideopreview?node_id=' + detail)
previewWidget.videoEl.autoplay = true
});
},
});