Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions freevideo_engine/comfy_assets.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,6 +152,7 @@ def output_summary(report, relative_video):
gpu_total_bytes=report.get('profile', {}).get('policy', {}).get('hardware', {}).get('vram_total'),
geometry=report.get('geometry', {}),
sampling_plan=report.get('sampling_plan'),
reference_trims=report.get('encoding', {}).get('reference_trims') or [],
video=relative_video.as_posix(), report=relative_video.with_suffix(
'.debug.json' if report.get('diagnostic_file') == relative_video.with_suffix('.debug.json').name
else '.request.json').as_posix())
Expand Down
5 changes: 5 additions & 0 deletions freevideo_engine/comfy_bridge.py
Original file line number Diff line number Diff line change
Expand Up @@ -355,6 +355,8 @@ def progress_message(event):
'encoder_compute': 'Encoding text and images',
'encoder_oom': 'Releasing encoder weights after insufficient GPU memory',
'encoder_retry': 'Retrying text encoding with more GPU workspace',
'encoder_spill': 'Leaving more GPU memory for text encoding',
'encoder_low_memory': 'Encoding long references in smaller blocks',
'encoder_conditioning_pack': 'Preparing prompt data',
'keyframe_vae': 'Encoding reference media',
'media_vae': 'Encoding reference media',
Expand Down Expand Up @@ -414,6 +416,9 @@ def progress_message(event):
elapsed_seconds=event.get('elapsed_seconds'))
if name == 'media_encode_phase':
return {'label': str(event.get('phase', 'Encoding input media')), 'timing_phase': 'encoding'}
if name == 'reference_trimmed':
return {'label': 'Preparing reference media', 'timing_phase': 'encoding',
'reference_trimmed': {key: event.get(key) for key in ('kind', 'number', 'seconds', 'used_seconds')}}
if name == 'prepared_blocks':
return {'label': 'Loading cached video model · %s / 50 blocks' % event.get('blocks', '?'),
'timing_phase': 'load'}
Expand Down
5 changes: 1 addition & 4 deletions freevideo_engine/comfy_media.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,8 +26,6 @@ def _audio(value, path):
or waveform.shape[1] not in (1, 2) or type(rate) is not int or not 8000 <= rate <= 192000
or not bool(torch.isfinite(waveform).all())):
raise ValueError('Expected a finite mono/stereo Comfy AUDIO input')
if waveform.shape[-1] > rate * 15:
raise ValueError('Trim reference audio to 15 s or less before generation')
samples = waveform[0].detach().float().cpu().contiguous().numpy()
layout = 'mono' if len(samples) == 1 else 'stereo'
with av.open(str(path), 'w', format='wav') as container:
Expand Down Expand Up @@ -75,8 +73,7 @@ def export(run, canvas, *, first=None, last=None, references=None, loras=None, c
path = path.with_suffix('.png')
_image(value, path)
elif kind == 'video':
if value.get_duration() > 15.05:
raise ValueError('Trim reference video to 15 s or less before generation')
# A clip longer than the generated video is shortened, and reported, by media encoding.
path = path.with_suffix('.mp4')
# Native save_to honors crops/trims and remuxes compatible file
# inputs without expanding all frames into host float tensors.
Expand Down
2 changes: 1 addition & 1 deletion freevideo_engine/comfy_nodes.py
Original file line number Diff line number Diff line change
Expand Up @@ -219,7 +219,7 @@ class FreeVideoReference(io.ComfyNode):
def define_schema(cls):
return io.Schema(node_id='FreeVideoReference', display_name='FreeVideo · Reference (experimental)',
category='FreeVideo', is_experimental=True,
description='Append one image, video or audio reference in order. Clips must be trimmed to 15 s or less. '
description='Append one image, video or audio reference in order. A clip is cut to the generated length, at most 15 s, and the cut is reported. '
'Connect the stack to Generate. VDN reference quality is experimental.',
inputs=[io.Image.Input('image', optional=True), io.Video.Input('video', optional=True),
io.Audio.Input('audio', optional=True), References.Input('previous', optional=True)],
Expand Down
59 changes: 56 additions & 3 deletions freevideo_engine/encode_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -211,6 +211,8 @@ def phase(name, **metrics):
task = task_for(request['media'])
if task != 't2va':
normalized, media_kwargs = prepare(request['media'], request['geometry'], Path(request['output']).parent / 'media')
# Saved with the conditioning, so a reused input cache still reports them.
loading['reference_trims'] = [row['trimmed'] for row in normalized if row.get('trimmed')]
if keyframes is not None:
if not isinstance(keyframes, dict) or set(keyframes) != {'first', 'last'}:
raise ValueError('FL2VA encoding requires keyframes.first and keyframes.last')
Expand Down Expand Up @@ -281,6 +283,17 @@ def phase(name, **metrics):
from .encoder_memory import token_counts
loading['token_summary'] = token_counts(tokens)
tokenize_seconds = time.perf_counter() - encode_started
from contextlib import contextmanager, nullcontext
from .encoder_workspace import (LOW_MEMORY_CHUNK, WHOLE_BLOCK_TOKENS, SharedMemorySpill, SpillGuard, Workspace, adapter_reader,
dedicated_room, input_size, make_room, plan)
size = loading['encoder_input'] = input_size(tokens)
workspace = Workspace.for_encoder(torch, path)
reader = adapter_reader()
guard = SpillGuard(torch, reader)
current = {}
# What this process may hold under FreeVideo's plan, beside the physical room.
budget_bytes = int(float(request['gpu_budget_gb']) * 1e9) if request.get('gpu_budget_gb') else None

if resident is not None and encoder_cache_hit:
resident.encoder_room(clip, tokens, estimated)
loading['resident_admission'] = list(resident.decisions)
Expand All @@ -295,6 +308,10 @@ def timed_load(*values, **kwargs):
phase('encoder_device_load', load_seconds=load_seconds)
print(json.dumps({'event': 'encoder_device_reuse' if encoder_on_gpu(clip) else 'encoder_device_load_start'}), flush=True)
tick = time.perf_counter()
import comfy.model_management as memory
if current.get('need') is not None:
# The native planner keeps 0.8 GiB beside the reserve; leave this attempt's room.
memory.EXTRA_RESERVED_VRAM = max(memory.EXTRA_RESERVED_VRAM, current['need'] - int(.8 * 2**30))
result = native_load(*values, **kwargs)
torch.cuda.synchronize()
transfer_seconds[0] += time.perf_counter() - tick
Expand Down Expand Up @@ -328,19 +345,52 @@ def timed_load(*values, **kwargs):
released = release_mapped_pages(checkpoint_paths)
if released:
print(json.dumps(dict(event='encoder_pages_released', **released)), flush=True)
if current.get('need') is not None:
# cudaMemGetInfo can promise room the WDDM budget does not have; check what is
# really free and move weights back to host memory until the forward fits.
room = make_room(clip.patcher, torch, current['need'], reader, budget_bytes)
loading.setdefault('encoder_workspace', []).append(dict(current['plan'], mode=current['mode'], **room))
guard.arm()
from .encoder_memory import snapshot
import comfy.model_management as memory
phase('encoder_compute', load_seconds=load_seconds, device_load_seconds=transfer_seconds[0],
gpu=snapshot(torch, clip, memory))
print(json.dumps({'event': 'encoder_compute_start'}), flush=True)
return result
@contextmanager
def attempt(index, failure):
# Room is planned against what this device can give the encoder with none of its weights loaded.
capacity = dedicated_room(torch, reader, budget_bytes) + int(clip.patcher.loaded_size())
mode, need, details = plan(workspace, size, capacity, index, current.get('mode') if failure else None)
current.update(mode=mode, need=need, plan=details)
from .encoder_lowmem import low_memory
if mode == 'blocks':
phase('encoder_low_memory', **details)
context = low_memory(clip.cond_stage_model, LOW_MEMORY_CHUNK)
elif size['sequence'] >= WHOLE_BLOCK_TOKENS:
# One block: the whole sequence at once, without the native dense T x T mask.
context = low_memory(clip.cond_stage_model, 1 << 30)
else:
context = nullcontext()
from .encoder_precision import bf16_language_model
try:
# BF16, as the official pipeline runs this encoder (ComfyUI's base runs it in FP32).
with context, bf16_language_model(clip.cond_stage_model):
yield
workspace.learn(mode, size, guard.used(), spilled=guard.spill is not None)
except SharedMemorySpill as error:
workspace.learn(mode, size, error.used_bytes, spilled=True)
raise
finally:
guard.disarm()

clip.load_model = timed_load
try:
from .encoder_memory import encode_with_recovery
from .encoder_pinning import readonly_pinning
import comfy.model_management as memory
with readonly_pinning(torch, memory):
encoded, attempts = encode_with_recovery(clip, tokens, memory, torch, phase)
encoded, attempts = encode_with_recovery(clip, tokens, memory, torch, phase, max_attempts=4,
attempt=attempt)
loading['encoder_attempts'] = attempts
# Finish async forward work at the existing conditioning boundary, so
# it is not attributed to packing/saving on the host.
Expand All @@ -350,7 +400,10 @@ def timed_load(*values, **kwargs):
phase('encoder_conditioning_pack', gpu=snapshot(torch, clip, memory))
finally:
del clip.load_model
del timed_load, native_load
guard.disarm()
if reader is not None:
reader.close()
del timed_load, native_load, attempt
del tokens
task = 'fl2va' if images is not None else task
value = to_cache(encoded, prompt, task=task)
Expand Down
2 changes: 1 addition & 1 deletion freevideo_engine/encoder_diagnostics.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
'encoder_options', 'encoder_cuda_setup', 'encoder_path_config', 'encoder_native_import',
'encoder_media_prepare', 'encoder_lookup', 'encoder_load', 'encoder_checkpoint_map',
'encoder_construct', 'encoder_tokenize', 'encoder_device_load', 'encoder_page_release',
'encoder_compute', 'encoder_oom', 'encoder_retry', 'encoder_conditioning_pack',
'encoder_compute', 'encoder_oom', 'encoder_spill', 'encoder_retry', 'encoder_low_memory', 'encoder_conditioning_pack',
'keyframe_vae', 'media_vae', 'encoder_save', 'encoder_idle_preload')


Expand Down
180 changes: 180 additions & 0 deletions freevideo_engine/encoder_lowmem.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,180 @@
"""Run the H3 text encoder's language model in token blocks for very long inputs.

The native forward keeps FP32 activations, builds one dense T x T causal mask
and runs every projection and MLP over the whole sequence at once. Measured on
the 32B encoder, a 15 s reference video (15.7k tokens) needed 9.5 GiB beyond
the weights and two videos plus three images (34k tokens) 22 GiB, more than a
16 GB card has. The same math fits in a fraction of that:

- attention runs per block of queries, each against all earlier keys with the
causal limit as a lower-right bias, so no mask is ever materialized; the
memory-efficient kernel takes FP32 but not grouped K/V heads, so each K/V
head is broadcast (a view, no copy) to its group of query heads;
- the position-wise MLP runs per block of tokens;
- DeepStack features are dropped after the layers that use them.

Only the language-model path the H3 encoder uses is replaced. Calls with KV
caches, attention masks or intermediate layer lists go to the native code.
"""
from contextlib import contextmanager
import types

import torch
import torch.nn.functional as F


def grouped_causal_attention(q, k, v, start):
"""Queries at positions start..start+len(q) over keys 0..start+len(q), causal, grouped K/V heads."""
from torch.nn.attention.bias import causal_lower_right
batch, heads, rows, width = q.shape
groups = k.shape[1]
per_group = heads // groups
end = start + rows
bias = causal_lower_right(rows, end)
output = torch.empty((batch, heads, rows, width), device=q.device, dtype=q.dtype)
for group in range(groups):
first = group * per_group
keys = k[:, group:group + 1, :end].expand(batch, per_group, end, width)
values = v[:, group:group + 1, :end].expand(batch, per_group, end, width)
output[:, first:first + per_group] = F.scaled_dot_product_attention(
q[:, first:first + per_group], keys, values, attn_mask=bias)
return output


def rope(x, freqs_cis):
"""The native apply_rope, for one tensor (queries or keys) at a time."""
cos, sin, nsin = freqs_cis[0], freqs_cis[1], freqs_cis[2]
embedded = x * cos
half = embedded.shape[-1] // 2
embedded[..., :half].addcmul_(x[..., half:], nsin)
embedded[..., half:].addcmul_(x[..., :half], sin)
return embedded.to(x.dtype)


def chunked_attention(attention, chunk):
original = attention.forward

def forward(self, hidden_states, attention_mask=None, freqs_cis=None, optimized_attention=None,
past_key_value=None, sliding_window=None):
batch, length, _ = hidden_states.shape
if optimized_attention is not None: # a native caller: its mask and caches apply
return original(hidden_states, attention_mask=attention_mask, freqs_cis=freqs_cis,
optimized_attention=optimized_attention, past_key_value=past_key_value,
sliding_window=sliding_window)
if (attention_mask is not None or past_key_value is not None or sliding_window is not None
or freqs_cis is None or any(t.shape[-2] != length for t in freqs_cis[:3])):
raise ValueError('Block-wise encoder attention supports causal attention without caches only')
if self.merged_qkv:
fused = self.qkv_proj(hidden_states)
projected = dict(zip(('q', 'k', 'v'), fused.split((self.inner_size, self.kv_size, self.kv_size), dim=-1)))
else:
projected = {}
keys = (projected.get('k') if projected else self.k_proj(hidden_states))
values = (projected.get('v') if projected else self.v_proj(hidden_states))
keys = keys.view(batch, length, self.num_kv_heads, self.head_dim).transpose(1, 2)
values = values.view(batch, length, self.num_kv_heads, self.head_dim).transpose(1, 2)
if self.k_norm is not None:
keys = self.k_norm(keys)
keys = rope(keys, freqs_cis)
output = None
for start in range(0, length, chunk):
end = min(length, start + chunk)
queries = (projected['q'][:, start:end] if projected else self.q_proj(hidden_states[:, start:end]))
queries = queries.reshape(batch, end - start, self.num_heads, self.head_dim)
queries = queries.transpose(1, 2)
if self.q_norm is not None:
queries = self.q_norm(queries)
queries = rope(queries, tuple(t[..., start:end, :] for t in freqs_cis[:3]))
value = grouped_causal_attention(queries, keys, values, start)
del queries
value = self.o_proj(value.transpose(1, 2).reshape(batch, end - start, self.num_heads * self.head_dim))
if output is None:
output = torch.empty((batch, length, value.shape[-1]), device=value.device, dtype=value.dtype)
output[:, start:end] = value
del value
return output, None
return types.MethodType(forward, attention)


def chunked_mlp(mlp, chunk):
original = mlp.forward

def forward(self, x):
if x.shape[1] <= chunk:
return original(x)
output = torch.empty_like(x)
for start in range(0, x.shape[1], chunk):
output[:, start:start + chunk] = original(x[:, start:start + chunk])
return output
return types.MethodType(forward, mlp)


def language_model(model):
"""The Llama2_ module inside the native H3 text encoder."""
for _, module in model.named_modules():
if type(module).__name__ == 'Llama2_' and hasattr(module, 'layers'):
return module
raise ValueError('The native H3 text encoder has no Llama language model')


def chunked_forward(llama, original):
def forward(self, x, attention_mask=None, embeds=None, num_tokens=None, intermediate_output=None,
final_layer_norm_intermediate=True, dtype=None, position_ids=None, embeds_info=[],
past_key_values=None, input_ids=None, deepstack_embeds=None, visual_pos_masks=None):
if (attention_mask is not None or past_key_values is not None
or isinstance(intermediate_output, (list, str))):
return original(x, attention_mask=attention_mask, embeds=embeds, num_tokens=num_tokens,
intermediate_output=intermediate_output,
final_layer_norm_intermediate=final_layer_norm_intermediate, dtype=dtype,
position_ids=position_ids, embeds_info=embeds_info, past_key_values=past_key_values,
input_ids=input_ids, deepstack_embeds=deepstack_embeds, visual_pos_masks=visual_pos_masks)
x = embeds if embeds is not None else self.embed_tokens(x, out_dtype=dtype)
length = x.shape[1]
if position_ids is None:
position_ids = torch.arange(length, device=x.device).unsqueeze(0)
freqs_cis = self.compute_freqs_cis(position_ids, x.device)
if intermediate_output is not None and intermediate_output < 0:
intermediate_output = len(self.layers) + intermediate_output
deepstack = list(deepstack_embeds) if deepstack_embeds is not None else None
deepstack_embeds = None
intermediate = None
for index, layer in enumerate(self.layers):
# The block-wise attention installed on each layer needs no mask.
x, _ = layer(x=x, attention_mask=None, freqs_cis=freqs_cis, optimized_attention=None,
past_key_value=None)
# DeepStack: per-layer visual features at image positions (Qwen3-VL), as the native forward.
if deepstack is not None and index < len(deepstack):
x[visual_pos_masks] = x[visual_pos_masks] + deepstack[index].to(x)
deepstack[index] = None # used by this layer only
if index == intermediate_output:
intermediate = x.clone()
if self.norm is not None:
x = self.norm(x)
if intermediate is not None and final_layer_norm_intermediate and self.norm is not None:
intermediate = self.norm(intermediate)
return x, intermediate
return types.MethodType(forward, llama)


@contextmanager
def low_memory(model, chunk):
"""Within the block, the encoder's language model runs attention and MLPs in `chunk`-token blocks."""
llama = language_model(model)
patched = []

def patch(module, replacement):
patched.append((module, module.__dict__.get('forward')))
module.forward = replacement

try:
patch(llama, chunked_forward(llama, llama.forward))
for layer in llama.layers:
patch(layer.self_attn, chunked_attention(layer.self_attn, chunk))
patch(layer.mlp, chunked_mlp(layer.mlp, chunk))
yield
finally:
for module, previous in reversed(patched):
if previous is None:
del module.forward # back to the class's own forward
else:
module.forward = previous
Loading